Python类型提示装饰器:实现DBConnection注入与手动传参兼容问题
实现支持手动传入与自动注入的DBConnection装饰器
你需要让装饰器同时支持两种调用方式:自动注入DBConnection,或手动传入该参数。之前的实现类型提示不兼容手动传参,通过@overload可以解决这个问题,以下是修正后的代码:
from functools import wraps import inspect from typing import Awaitable, Callable, Concatenate, ParamSpec, TypeVar, overload from typing_extensions import reveal_type class DBConnection: ... T = TypeVar("T") P = ParamSpec("P") @overload def inject_db_connection( f: Callable[Concatenate[DBConnection, P], Awaitable[T]] ) -> Callable[P, Awaitable[T]] | Callable[Concatenate[DBConnection, P], Awaitable[T]]: ... @overload def inject_db_connection( f: None = None, ) -> Callable[[Callable[Concatenate[DBConnection, P], Awaitable[T]]], Callable[P, Awaitable[T]] | Callable[Concatenate[DBConnection, P], Awaitable[T]]]: ... def inject_db_connection( f: Callable[Concatenate[DBConnection, P], Awaitable[T]] | None = None, ) -> Callable[[Callable[Concatenate[DBConnection, P], Awaitable[T]]], Callable[P, Awaitable[T]] | Callable[Concatenate[DBConnection, P], Awaitable[T]]] | Callable[P, Awaitable[T]] | Callable[Concatenate[DBConnection, P], Awaitable[T]]: def decorator(func: Callable[Concatenate[DBConnection, P], Awaitable[T]]) -> Callable[P, Awaitable[T]] | Callable[Concatenate[DBConnection, P], Awaitable[T]]: @wraps(func) async def inner(*args: P.args | tuple[DBConnection, *P.args], **kwargs: P.kwargs) -> T: # 检查是否手动传入了db_connection参数 signature = inspect.signature(func).parameters passed_args = dict(zip(signature, args)) if "db_connection" in kwargs or "db_connection" in passed_args: return await func(*args, **kwargs) # 未传入则自动注入DBConnection实例 return await func(DBConnection(), *args, **kwargs) return inner # 处理装饰器带括号和不带括号的两种调用方式 if f is None: return decorator return decorator(f) @inject_db_connection async def get_user(db_connection: DBConnection, user_id: int) -> dict: assert db_connection return {"user_id": user_id} async def main() -> None: # 自动注入DBConnection,类型检查无问题 user1 = await get_user(user_id=1) # 手动传入db_connection关键字参数,现在类型检查不再报错 db_connection = DBConnection() user2 = await get_user(db_connection=db_connection, user_id=1) # 也支持位置参数手动传入db_connection user3 = await get_user(db_connection, 2) # 类型推断均正常为dict[Any, Any] reveal_type(user1) reveal_type(user2) reveal_type(user3)
关键修改说明
- 重载装饰器签名:通过
@overload定义两种情况,覆盖装饰器直接调用和带括号调用的场景,同时声明装饰后的函数支持两种参数列表(带/不带DBConnection)。 - 放宽内部函数参数类型:将
inner的args类型设为P.args | tuple[DBConnection, *P.args],兼容自动注入时的参数列表,以及手动传入DBConnection的参数列表。 - 保留原有逻辑:依然通过检查参数是否存在来决定是直接调用原函数还是自动注入实例,确保运行时行为符合预期。
这样修改后,类型检查工具(如mypy)会认可两种调用方式,不会再抛出Unexpected keyword argument错误,同时返回值的类型推断保持正常。
内容的提问来源于stack exchange,提问作者Michal K
相关产品推荐
相关产品推荐

