如何正确注解Python子容器类以实现加法的正确类型推断
背景
假设你实现了一个基础容器类MyList及其子类MyList2:
from collections import UserList from typing import Generic, TypeVar, Self T = TypeVar("T") class MyList(UserList[T]): def __add__(self, other: MyList) -> Self: return type(self)(val + other[idx] for idx, val in enumerate(self)) class MyList2(MyList[T], Generic[T]): pass
定义以下变量:
a = MyList([1, 2, 3]) # MyList[int] b = MyList([0.5, 1.5, 2.5]) # MyList[float] c = MyList2([1, 2, 3]) # MyList2[int]
静态类型检查器对d和e的类型推断结果如下:
d = a + b # 推断为MyList[int],期望为MyList[float] e = c + b # 推断为MyList2[int],期望为MyList2[float]
这是因为实现中使用了type(self),导致类型检查器认为结果与加号左侧变量的容器类型和元素类型(int)一致。运行时无问题,但每次调用__add__都需手动辅助类型检查器,不够理想。
若按以下方式实现可避免该问题:
class MyList3(UserList[T]): def __add__(self, other: MyList3): return MyList3(val + other[idx] for idx, val in enumerate(self)) class MyList4(MyList3[T], Generic[T]): pass
此时:
f = MyList3([1, 2, 3]) # MyList3[int] g = MyList3([0.5, 1.5, 2.5]) # MyList3[float] h = f + g # MyList3[float]
但这种方式无法生成与参数相同的容器类型(若f和g是MyList4实例,h仍为MyList3类型),只能使用父类。若要解决需为所有子类实现__add__,会导致代码冗余。
问题
请问如何实现,使得对于MyList的任意子容器类L,都能满足:
__add__(L[int], L[float]) -> L[float] __add__(L[float], L[npt.NDArray[np.floating]]) -> L[npt.NDArray[np.floating]] ... __add__(L[A], L[B]) -> L[C] # 当__add__(A, B) -> C,即A类型与B类型相加得到C类型
解决方案
可以通过泛型类型注解结合Self类型实现需求,同时确保子类无需重写__add__方法。核心思路是让类型检查器自动推断元素类型相加后的结果,并保留容器的具体子类类型:
实现代码
from collections import UserList from typing import Generic, TypeVar, Self # 定义类型变量:T表示当前容器的元素类型,U表示另一个容器的元素类型,R表示相加后的结果元素类型 T = TypeVar("T") U = TypeVar("U") R = TypeVar("R") class MyList(UserList[T], Generic[T]): def __add__(self, other: "MyList[U]") -> "Self[R]": # 运行时通过type(self)确保返回当前子类的实例 return type(self)(val + other[idx] for idx, val in enumerate(self)) # 子类无需重写__add__,自动继承类型注解和实现 class MyList2(MyList[T], Generic[T]): pass
效果验证
import numpy as np from numpy.typing import NDArray # 基础类型测试 a = MyList([1, 2, 3]) # MyList[int] b = MyList([0.5, 1.5, 2.5]) # MyList[float] d = a + b # 类型检查器推断为MyList[float],符合预期 # 子类测试 c = MyList2([1, 2, 3]) # MyList2[int] e = c + b # 类型检查器推断为MyList2[float],符合预期 # 复杂元素类型测试 f = MyList2([1.0, 2.0, 3.0]) # MyList2[float] g = MyList2([np.array([1.5]), np.array([2.5]), np.array([3.5])]) # MyList2[NDArray[np.floating]] h = f + g # 类型检查器推断为MyList2[NDArray[np.floating]],符合预期
原理说明
- 类型变量约束:通过
T、U、R三个类型变量分别标记当前容器元素类型、另一个容器元素类型以及相加后的结果类型,类型检查器会根据元素的__add__逻辑自动推断R的具体类型。 Self类型的泛型支持:Self[R]表示返回的是当前容器子类的实例,且其元素类型为相加后的R,既保留了容器的子类类型,又正确推导了元素类型。- 运行时兼容性:
type(self)确保运行时返回的是当前调用者的子类实例,和原实现的运行逻辑一致,无额外开销。
如果使用Python 3.12+,可以利用PEP 695的简化泛型语法,代码更简洁:
from collections import UserList from typing import Self class MyList[T](UserList[T]): def __add__(self, other: MyList[U]) -> Self[V]: return type(self)(val + other[idx] for idx, val in enumerate(self)) class MyList2[T](MyList[T]): pass
内容的提问来源于stack exchange,提问作者N.D
相关产品推荐
相关产品推荐

