如何正确为TypeVar绑定Protocol类型约束?
我需要给一个简单的快速排序函数添加类型注解,满足以下要求:
- 接收的列表元素可以是任意值,但必须实现
__lt__、__gt__中的一种或两种方法 - 列表内所有元素必须使用相同的比较逻辑(不能混合仅支持
__lt__和仅支持__gt__的元素) - 返回排序后的同元素列表
第一次尝试与问题
from typing import Any, Protocol class SupportsLT(Protocol): def __lt__(self, other: Any, /) -> bool: ... class SupportsGT(Protocol): def __gt__(self, other: Any, /) -> bool: ... def quick_sort[T: SupportsLT | SupportsGT](seq: list[T]) -> list[T]: if len(seq) <= 1: return list(seq) pivot = seq[0] smaller: list[T] = [] larger_or_equal: list[T] = [] for i in range(1, len(seq)): item = seq[i] if item < pivot: smaller.append(item) else: larger_or_equal.append(item) return [*quick_sort(smaller), pivot, *quick_sort(larger_or_equal)]
运行mypy --strict时报错:
error: Unsupported left operand type for < (some union) [operator]
问题出在类型约束是SupportsLT | SupportsGT,类型检查器无法保证列表中的元素不会混合两种协议类型,因此无法确认<操作符对所有元素都合法。
第二次尝试与问题
改用约束而非绑定的写法:
def quick_sort[T: (SupportsLT, SupportsGT)](seq: list[T]) -> list[T]: if len(seq) <= 1: return list(seq) pivot = seq[0] smaller: list[T] = [] larger_or_equal: list[T] = [] for i in range(1, len(seq)): item = seq[i] if item < pivot: smaller.append(item) else: larger_or_equal.append(item) return [*quick_sort(smaller), pivot, *quick_sort(larger_or_equal)]
运行mypy --strict时出现列表元素类型不兼容错误:
error: List item 0 has incompatible type "list[T]"; expected "SupportsLT" [list-item]
error: List item 0 has incompatible type "list[T]"; expected "SupportsGT" [list-item]
error: List item 2 has incompatible type "list[T]"; expected "SupportsLT" [list-item]
error: List item 2 has incompatible type "list[T]"; expected "SupportsGT" [list-item]
更严重的是,这种写法会导致类型推导错误:
测试代码:
from typing import reveal_type x = quick_sort([12, 3]) reveal_type(x)
输出结果为:
Revealed type is "builtins.list[SupportsLT]"
而正确的推导结果应该是builtins.list[builtins.int]。
更新说明
我同时使用SupportsLT和SupportsGT作为约束的原因是:Python的比较机制中,当左操作数没有实现__lt__方法时,会自动调用右操作数的__gt__方法并传入左操作数。因此仅实现__gt__方法的对象也应该被函数接受,示例如下:
from typing import Self class HasGT: def __init__(self, value: int) -> None: self.value = value def __gt__(self, other: Self) -> bool: return self.value > other.value print(HasGT(5) < HasGT(6)) # 输出True
内容的提问来源于stack exchange,提问作者CleverBoy

