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

Python类型提示:如何指定相互依赖的合法泛型类型组合?

问题:为Python泛型类添加精准的类型提示

我有一个小型但功能灵活(或许过于灵活)的Python类,它实际上包含两个相互约束的泛型类型。以下是我意图的实现代码(目前类型提示不完善):

from typing import Callable, Any, Sequence, Literal, MutableSequence, MutableMapping

class ApplyTo:
    def __init__(
        self,
        *transforms: Callable[..., Any],
        to: Any | Sequence[Any],
        dispatch: Literal['separate', 'joint'] = 'separate',
    ):
        self._transform = TransformsPipeline(*transforms)
        self._to = to if isinstance(to, Sequence) else [to]
        self._dispatch = dispatch

    def __call__(self, data: MutableSequence | MutableMapping):
        if self._dispatch == 'separate':
            for key in self._to:
                data[key] = self._transform(data[key])
            return data

        if self._dispatch == 'joint':
            args = [data[key] for key in self._to]
            transformed = self._transform(*args)
            for output, key in zip(transformed, self._to):
                data[key] = output
            return data

        assert False

我已确认该代码在运行时可正常工作,但类型提示非常糟糕。

我的需求是:

  • 当to的类型为int时,data的类型应为MutableSequence | MutableMapping[int, Any]
  • 当to的类型为Hashable(非int的情况)时,data的类型应为MutableMapping[对应to的类型, Any]

我知道int属于Hashable,这增加了实现难度。

我尝试的类型标注代码如下:

from typing import TypeVar, Generic, Callable, Any, Sequence, Literal, MutableSequence, MutableMapping, Hashable

T = TypeVar('T', bound=Hashable | int)
C = TypeVar('C', bound=MutableMapping[T, Any] | MutableSequence)


class ApplyTo(Generic[C, T]):
    def __init__(
        self,
        *transforms: Callable[..., Any],
        to: T | Sequence[T],
        dispatch: Literal['separate', 'joint'] = 'separate',
    ):
        self._transform = TransformsPipeline(*transforms)
        self._to = to if isinstance(to, Sequence) else [to]
        self._dispatch = dispatch

    def __call__(self, data: C):
        if self._dispatch == 'separate':
            for key in self._to:
                data[key] = self._transform(data[key])
            return data

        if self._dispatch == 'joint':
            args = [data[key] for key in self._to]
            transformed = self._transform(*args)
            for output, key in zip(transformed, self._to):
                data[key] = output
            return data

        assert False

这导致mypy报错:

error: Type variable "task_driven_sr.transforms.generic.T" is unbound  [valid-type]
note: (Hint: Use "Generic[T]" or "Protocol[T]" base class to bind "T" inside a class)
note: (Hint: Use "T" in function signature to bind "T" inside a function)

请问是否有正确的方式进行类型提示,将to和data的类型绑定在一起?或者我的实现思路存在问题?

编辑:已修复dispatch的joint分支中的代码,该部分与类型提示无关,但已修正为正确逻辑。


解决方案

核心问题分析

你遇到的unbound错误是因为泛型参数顺序错误:C依赖T,但你把C放在了Generic的参数首位,导致类型系统无法先绑定T来推导C。另外,直接用单个TypeVar无法区分int的特殊场景,需要结合@overload来实现精准的类型约束。

修正后的类型标注代码

from typing import (
    TypeVar, Generic, Callable, Any, Sequence, Literal,
    MutableSequence, MutableMapping, Hashable, overload, Tuple
)

# 定义泛型变量:TKey对应to的类型,TData对应data中元素的类型
TKey = TypeVar('TKey', bound=Hashable)
TData = TypeVar('TData')

class ApplyTo(Generic[TKey]):
    # 重载__init__:区分to为int和其他Hashable的场景
    @overload
    def __init__(
        self,
        *transforms: Callable[[TData], TData],
        to: int | Sequence[int],
        dispatch: Literal['separate', 'joint'] = 'separate',
    ): ...

    @overload
    def __init__(
        self,
        *transforms: Callable[[TData], TData],
        to: TKey | Sequence[TKey],
        dispatch: Literal['separate', 'joint'] = 'separate',
    ): ...

    def __init__(
        self,
        *transforms: Callable[..., Any],
        to: TKey | Sequence[TKey],
        dispatch: Literal['separate', 'joint'] = 'separate',
    ):
        self._transform = TransformsPipeline(*transforms)
        self._to = to if isinstance(to, Sequence) else [to]
        self._dispatch = dispatch

    # 重载__call__:区分data是序列还是映射的场景
    @overload
    def __call__(self, data: MutableSequence[TData]) -> MutableSequence[TData]: ...

    @overload
    def __call__(self, data: MutableMapping[TKey, TData]) -> MutableMapping[TKey, TData]: ...

    def __call__(self, data: MutableSequence[TData] | MutableMapping[TKey, TData]):
        if self._dispatch == 'separate':
            for key in self._to:
                data[key] = self._transform(data[key])
            return data

        if self._dispatch == 'joint':
            args = [data[key] for key in self._to]
            transformed = self._transform(*args)
            for output, key in zip(transformed, self._to):
                data[key] = output
            return data

        assert False

关键修改说明

  1. 调整泛型参数逻辑:移除冗余的C类型变量,直接用TKey约束MutableMapping的键类型,同时用TData跟踪数据元素类型,让类型推导更清晰。
  2. 用@overload处理特殊场景:
    • 重载__init__明确:当to是int时,data可以是序列或整数键的映射;当to是其他Hashable类型时,data只能是对应键类型的映射。
    • 重载__call__让类型系统能正确推断输入输出的类型一致性。
  3. 修复unbound错误:将TKey作为类的唯一泛型参数,确保类型系统能先绑定键类型,再推导数据结构类型。

进阶优化(针对joint模式)

如果需要约束joint模式下变换函数的输入输出数量与to的长度匹配,可以添加额外重载:

# 针对joint模式的多参数变换重载
@overload
def __init__(
    self,
    *transforms: Callable[[Tuple[TData, ...]], Tuple[TData, ...]],
    to: Sequence[TKey],
    dispatch: Literal['joint'],
): ...

内容的提问来源于stack exchange,提问作者Maciej Ziaja

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 10:02:11