如何在Python中为特定连接结构的图提供类型提示?
我正在开发一个可归结为遍历不同图结构的应用程序,有一组共享函数要求图必须符合特定“结构”(即顶点间存在预期的转换关系)。我希望通过类型提示来表达这一要求,以下是最简图示例:
class NodeA: def go_to_B(self) -> NodeB: return NodeB() def go_to_C(self) -> NodeC: return NodeC() class NodeB: def go_to_A(self) -> NodeA: return NodeA() class NodeC: def go_to_A(self) -> NodeA: return NodeA()
假设我有一个函数,要求从给定节点X可以前往节点Y并返回X,我尝试用Protocols来表达该约束:
class BGoer(Protocol): def go_to_B(self) -> AGoer: pass class AGoer(Protocol): def go_to_A(self) -> BGoer: pass class Traverser: def traverse(self, start: BGoer) -> BGoer: b: AGoer = start.go_to_B() a: BGoer = b.go_to_A() return a
但这种方式会丢失原类型信息,无法调用原类型的其他方法,例如:
def main() -> None: a = NodeA() t = Traverser() a_after_traversing:BGoer = t.traverse(a) a_after_traversing.go_to_C()
mypy会报错:graph.py:37: error: "BGoer" has no attribute "go_to_C"; maybe "go_to_B"? [attr-defined]。
若将返回值指定为原类型:
def main() -> None: a = NodeA() t = Traverser() a_after_traversing:NodeA = t.traverse(a) a_after_traversing.go_to_C()
则会出现错误:graph.py:35: error: Incompatible types in assignment (expression has type "BGoer", variable has type "NodeA") [assignment]。
我尝试使用泛型Protocol解决,但未能成功。请问如何在表达图的预期转换关系的同时,保留原类型信息?
要同时表达图结构的转换约束并保留原类型信息,需使用泛型Protocol,通过泛型参数让协议间的类型引用形成闭环,从而保留输入节点的原始类型。具体实现步骤如下:
1. 定义泛型协议
首先声明类型变量,再定义关联的泛型协议,明确方法的返回类型与原节点类型的绑定关系:
from typing import Protocol, TypeVar # 定义绑定到对应协议的类型变量 T_BGoer = TypeVar('T_BGoer', bound='BGoer') T_AGoer = TypeVar('T_AGoer', bound='AGoer') class AGoer(Protocol[T_BGoer]): def go_to_A(self) -> T_BGoer: ... class BGoer(Protocol[T_AGoer]): def go_to_B(self) -> T_AGoer: ...
核心逻辑:
AGoer的go_to_A方法返回的泛型参数T_BGoer指向原始节点类型(如NodeA)BGoer的go_to_B方法返回的泛型参数T_AGoer指向目标节点类型(如NodeB)
2. 实现泛型遍历方法
修改Traverser的traverse方法,通过泛型参数确保输入与输出类型一致:
class Traverser: def traverse(self, start: T_BGoer) -> T_BGoer: b: AGoer[T_BGoer] = start.go_to_B() a: T_BGoer = b.go_to_A() return a
此时traverse方法会完全保留输入节点的原始类型,返回值类型与输入严格匹配。
3. 验证类型有效性
在main函数中,无需强制指定返回类型,mypy会自动推导原始节点类型,且可正常调用原类型的所有方法:
def main() -> None: a = NodeA() t = Traverser() # mypy自动推导a_after_traversing类型为NodeA a_after_traversing = t.traverse(a) # 可正常调用NodeA的go_to_C方法,无类型报错 a_after_traversing.go_to_C()
4. 节点类的自动兼容性
你的NodeA、NodeB类会自动满足泛型协议约束:
NodeA实现的go_to_B返回NodeB,而NodeB的go_to_A返回NodeA,完全符合BGoer[NodeB]与AGoer[NodeA]的协议要求- mypy会自动识别
NodeA是BGoer[NodeB]的实例,NodeB是AGoer[NodeA]的实例
这种方案既清晰表达了图节点间的转换规则,又完整保留了原始类型的方法信息,彻底解决了之前的类型报错问题。
内容的提问来源于stack exchange,提问作者Filip

