如何为接受基类任意子类的方法添加类型注解?
我在为基类的子类添加类型注解时遇到了问题:定义了基类Table(所有表子类共享其行为)和管理表集合的Database类,代码运行正常,但为Database的方法使用Table做类型注解时,传入子类(如Ducks)会导致MyPy类型检查失败。尝试过有界泛型但未成功,希望让类型检查器理解该方法可接受Table的任意子类,同时支持特定操作的更具体类型。
原代码
#!/usr/bin/env python3 from collections.abc import Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from typing import Any, Self, TypeVar # _Table = TypeVar("_Table", bound="Table") <- Tried this but did not work @dataclass class Table: @classmethod def get_name(cls) -> str: return cls.__name__.lower() @dataclass class Database: records: Mapping[str, Sequence[Table]] def filter( self, table: type[Table], _filter: Callable[[Table], bool] ) -> Iterable[Table]: yield from (i for i in self.records[table.get_name()] if _filter(i)) # USING cast(T) here kind of works but i think is kind a cheat. @dataclass class Ducks(Table): name: str age: int if __name__ == "__main__": records = { "ducks": [ Ducks(name="Patinhas", age=104), Ducks(name="Donald", age=30), Ducks(name="Huguinho", age=12), Ducks(name="Zezinho", age=12), Ducks(name="Luizinho", age=12), ] } db = Database(records) f: Callable[[Ducks], bool] = lambda t: t.age < 100 # <- BONUS QUESTION: is there a way to not use this long typing lambda and just use the untyped notation ? print(*db.filter(Ducks, _filter=f), sep="\n") # <- PROBLEM HERE
运行结果
$ python3 ex.py Ducks(name='Donald', age=30) Ducks(name='Huguinho', age=12) Ducks(name='Zezinho', age=12) Ducks(name='Luizinho', age=12)
MyPy报错信息
Mypy lint: ex.py 46 37 error arg-type Argument "_filter" to "filter" of "Database" has incompatible type "Callable[[Ducks], bool]"; expected "Callable[[Table], bool]" (lsp)
一、修复主问题:实现泛型兼容
你之前尝试的有界泛型思路是正确的,只是未在filter方法中正确应用。通过绑定到Table的TypeVar,可以让方法的table参数、过滤函数和返回值保持类型一致,让类型检查器认可子类的兼容性:
修改后的完整代码
#!/usr/bin/env python3 from collections.abc import Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from typing import Self, TypeVar, cast # 定义绑定到Table的TypeVar,约束为任意Table子类 _Table = TypeVar("_Table", bound="Table") @dataclass class Table: @classmethod def get_name(cls) -> str: return cls.__name__.lower() @dataclass class Database: records: Mapping[str, Sequence[Table]] def filter( self, table: type[_Table], _filter: Callable[[_Table], bool] ) -> Iterable[_Table]: # 用cast告诉类型检查器:当前表名对应的记录是_Table类型 table_records = cast(Sequence[_Table], self.records[table.get_name()]) yield from (item for item in table_records if _filter(item)) @dataclass class Ducks(Table): name: str age: int if __name__ == "__main__": records = { "ducks": [ Ducks(name="Patinhas", age=104), Ducks(name="Donald", age=30), Ducks(name="Huguinho", age=12), Ducks(name="Zezinho", age=12), Ducks(name="Luizinho", age=12), ] } db = Database(records) # 直接传入lambda,无需额外类型注解 print(*db.filter(Ducks, lambda t: t.age < 100), sep="\n")
关键修改点
- 定义有界
TypeVar:_Table = TypeVar("_Table", bound="Table"),指定该泛型变量只能是Table或其子类 - 泛型化
filter方法:将table参数标注为type[_Table],_filter标注为Callable[[_Table], bool],返回值标注为Iterable[_Table],确保三者类型统一 - 类型转换:用
cast(Sequence[_Table], self.records[table.get_name()])明确告知类型检查器,当前表对应的记录是目标子类类型,避免误报
修改后,MyPy会自动识别:当传入Ducks作为table参数时,过滤函数必须接受Ducks类型,返回的结果也是Iterable[Ducks],完全兼容子类操作。
二、解决额外问题:简化Lambda类型注解
不需要给lambda单独添加冗长的Callable[[Ducks], bool]类型注解。因为泛化后的filter方法已经通过_Table约束了过滤函数的参数类型,当你传入Ducks作为table参数时,类型检查器会自动推断lambda的参数t是Ducks类型,直接写lambda t: t.age < 100即可,MyPy会正确识别且不会报错。
内容的提问来源于stack exchange,提问作者hugo.avlia

