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

如何正确注解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]],符合预期

原理说明

  1. 类型变量约束:通过T、U、R三个类型变量分别标记当前容器元素类型、另一个容器元素类型以及相加后的结果类型,类型检查器会根据元素的__add__逻辑自动推断R的具体类型。
  2. Self类型的泛型支持:Self[R]表示返回的是当前容器子类的实例,且其元素类型为相加后的R,既保留了容器的子类类型,又正确推导了元素类型。
  3. 运行时兼容性: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 16:20:10