You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用ParamSpec和Concatenate实现Python装饰器的任意参数位置类型标注?

类型安全的psycopg AsyncConnection装饰器实现问题

我有一个Python装饰器,用于确保方法接收psycopg的AsyncConnection实例,但当前实现不具备类型安全性。尝试使用ParamSpec和Concatenate更新类型标注时遇到了困难。

当前实现代码

from typing import Callable, Coroutine, Any, TypeVar
from psycopg import AsyncConnection

R = TypeVar('R')

def ensure_conn(func: Callable[..., Coroutine[Any, Any, R]]) -> Callable[..., Coroutine[Any, Any, R]]:
    """确保函数接收conn参数。如果未提供conn,则生成新连接并传入函数"""

    async def wrapper(*args: Any, **kwargs: Any) -> R:
        # 获取命名关键字参数conn,或在位置参数中查找AsyncConnection实例
        kwargs_conn = kwargs.get("conn")
        conn_arg: AsyncConnection[Any] | None = None
        if isinstance(kwargs_conn, AsyncConnection):
            conn_arg = kwargs_conn
        elif not conn_arg:
            for arg in args:
                if isinstance(arg, AsyncConnection):
                    conn_arg = arg
                    break
        if conn_arg:
            # 如果已提供conn,直接调用原方法
            return await func(*args, **kwargs)
        else:
            # 如果未提供conn,生成新连接并传入方法
            db_driver = DbDriver()
            async with db_driver.connection() as conn:
                return await func(*args, **kwargs, conn=conn)

    return wrapper

当前使用方式

@ensure_conn
async def get_user(user_id: UUID, conn: AsyncConnection):
    async with conn.cursor() as cursor:
        # 业务逻辑

现有问题

调用时传入错误参数不会触发类型检查,例如:

get_user('519766c5-af86-47ea-9fa9-cee0c0de66b1', conn, arg_that_should_fail_typing)

尝试的改进实现(使用ParamSpec和Concatenate)

from typing import Callable, Coroutine, Any, ParamSpec, Concatenate, TypeVar
from psycopg import AsyncConnection

P = ParamSpec('P')
R = TypeVar('R')

def ensure_conn_decorator[**P, R](func: Callable[Concatenate[AsyncConnection[Any], P], R]) -> Coroutine[Any, Any, R]:
    """确保函数接收conn参数。如果未提供conn,则生成新连接并传入函数"""
    async def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
        # 获取命名关键字参数conn,或在位置参数中查找AsyncConnection实例
        kwargs_conn = kwargs.get("conn")
        conn_arg: AsyncConnection[Any] | None = None
        if isinstance(kwargs_conn, AsyncConnection):
            conn_arg = kwargs_conn
        elif not conn_arg:
            for arg in args:
                if isinstance(arg, AsyncConnection):
                    conn_arg = arg
                    break
        if conn_arg:
            # 如果已提供conn,直接调用原方法
            return await func(*args, **kwargs)
        else:
            # 如果未提供conn,生成新连接并传入方法
            db_driver = DbDriver()
            async with db_driver.connection() as conn:
                return await func(*args, **kwargs, conn=conn)

    return wrapper

改进版本的问题

  • conn必须作为方法的第一个参数,无法放在任意位置(通常我们会把它作为最后一个参数)
  • 返回类型报错:
Expression of type "(**P@ensure_conn_decorator) -> Coroutine[Any, Any, R@ensure_conn_decorator]" is incompatible with return type "Coroutine[Any, Any, R@ensure_conn_decorator]"
  "function" is incompatible with "Coroutine[Any, Any, R@ensure_conn_decorator]"

内容的提问来源于stack exchange,提问作者Alice Yu

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.13 11:50:09