适配SQLAlchemy场景:如何为双输入类型函数添加Mypy类型注解?
问题背景
SQLAlchemy具备一项实用特性:编写的同一表达式可根据输入类型生成SQL条件或直接求值。我尝试在代码中用函数复刻该特性,但无法为函数添加正确的类型注解以通过Mypy的类型检查。
代码可正常运行,两个断言均通过,但Mypy抛出如下错误:
example.py:13: error: Argument 1 has incompatible type "SQLColumnExpression[Any]"; expected "str | int | bool | date | datetime" [arg-type]
example.py:14: error: Argument 1 has incompatible type "str | int | bool | date | datetime"; expected "SQLColumnExpression[Any]" [arg-type]
问题核心在于:Operator类型被声明为两个Callable的Union,但代码要求该变量同时兼容两种输入类型(而非二选一)。需要明确如何声明此类类型,或是否存在SQLAlchemy内置的现成类型可用。
原始代码
from sqlalchemy import SQLColumnExpression, Column, Integer, SQLColumnExpression from sqlalchemy.orm import declarative_base from datetime import date, datetime from typing import Union, Callable, Any, overload Value = Union[str, int, bool, date, datetime] Operator = Union[ Callable[[SQLColumnExpression[Any], Value], SQLColumnExpression[Any]], Callable[[Value, Value], bool], ] def some_function(var: Operator, column: SQLColumnExpression[Any], value: Value) -> None: assert isinstance(var(column, value), SQLColumnExpression) assert isinstance(var(value, value), bool) @overload def op(column: SQLColumnExpression[Any], value: Value) -> SQLColumnExpression[Any]: ... @overload def op(column: Value, value: Value) -> bool: ... def op(column, value): # type: ignore print("OP", type(column), type(value)) return column == 0 Base = declarative_base() class Model(Base): # type: ignore __tablename__ = "example" id = Column(Integer, primary_key=True) some_function(op, Model.id, 0)
解决方案
1. 用Protocol定义多态调用类型
Union类型的Callable表示函数只能是两种签名中的一种,但我们需要的是同时支持两种调用签名的函数。此时应使用typing.Protocol自定义一个协议,明确要求实现类/函数必须兼容两种调用方式:
from typing import Protocol, Any from sqlalchemy import SQLColumnExpression from datetime import date, datetime Value = Union[str, int, bool, date, datetime] class Operator(Protocol): def __call__(self, column: SQLColumnExpression[Any], value: Value) -> SQLColumnExpression[Any]: ... def __call__(self, column: Value, value: Value) -> bool: ...
2. 修正代码的类型注解
将some_function的var参数类型改为上述自定义的Operator协议,同时可以移除op函数的# type: ignore注释,Mypy现在能正确识别重载的合法性。
修改后的完整代码:
from sqlalchemy import SQLColumnExpression, Column, Integer from sqlalchemy.orm import declarative_base from datetime import date, datetime from typing import Union, Protocol, Any, overload Value = Union[str, int, bool, date, datetime] class Operator(Protocol): def __call__(self, column: SQLColumnExpression[Any], value: Value) -> SQLColumnExpression[Any]: ... def __call__(self, column: Value, value: Value) -> bool: ... def some_function(var: Operator, column: SQLColumnExpression[Any], value: Value) -> None: assert isinstance(var(column, value), SQLColumnExpression) assert isinstance(var(value, value), bool) @overload def op(column: SQLColumnExpression[Any], value: Value) -> SQLColumnExpression[Any]: ... @overload def op(column: Value, value: Value) -> bool: ... def op(column, value): print("OP", type(column), type(value)) return column == 0 Base = declarative_base() class Model(Base): __tablename__ = "example" id = Column(Integer, primary_key=True) some_function(op, Model.id, 0)
3. SQLAlchemy内置类型参考
SQLAlchemy内部存在类似的多态表达式类型(如_ColumnExpressionArgument),但这类多为内部私有类型,不建议直接依赖。自定义Protocol是最稳妥、最贴合需求的方案。
内容的提问来源于stack exchange,提问作者Pul Ess

