能否对Python数据类的模式匹配进行静态穷尽性检查?
我们有多个继承自同一个ABC类的数据类,编写了一个通过模式匹配分别处理每种情况的函数:
from abc import ABC from dataclasses import dataclass @dataclass class Token(ABC): pass @dataclass class PlusToken(Token): pass @dataclass class DigitToken(Token): digit: int def pretty_print(token: Token) -> str: match token: case PlusToken(): return 'PLUS <+>' case DigitToken(digit=digit): return f'DIGIT <{digit}>'
希望通过mypy静态检查上述pretty_print函数中的模式匹配是否穷尽(即没有遗漏任何子类情况)。
针对Enum或Union类型,通常会在case _:分支中使用typing模块的assert_never来检查穷尽性,但这对基于ABC的数据类无效:
def pretty_print(token: Token) -> str: match token: ... # PlusToken、DigitToken的处理分支 case _: assert_never(token) # 报错:Argument 1 to "assert_never" has incompatible type "Token"; expected "NoReturn" [arg-type]
推测原因是Python中可以直接实例化无抽象方法的ABC类,但想知道是否有其他方法实现模式匹配的穷尽性静态检查?
环境信息
- Python 3.11.2
- mypy 1.3.0
解决方法
1. 用@sealed标记密封基类
mypy支持通过@sealed装饰器标记密封类,限制基类只能在当前模块内被继承,同时mypy会追踪所有子类,配合assert_never就能实现穷尽性检查。
首先需要在mypy配置中启用strict模式,或者单独开启sealed_classes选项。修改代码如下:
from abc import ABC from dataclasses import dataclass from typing import assert_never, sealed @sealed @dataclass class Token(ABC): pass @dataclass class PlusToken(Token): pass @dataclass class DigitToken(Token): digit: int def pretty_print(token: Token) -> str: match token: case PlusToken(): return 'PLUS <+>' case DigitToken(digit=digit): return f'DIGIT <{digit}>' case _: assert_never(token)
当新增Token的子类但未在模式匹配中处理时,mypy会触发arg-type错误,提示传入assert_never的类型不是NoReturn,以此检查穷尽性。
2. 用Union类型替代ABC基类
如果可以调整类型结构,将Token定义为所有子类的Union类型,就能直接复用assert_never的方式:
from dataclasses import dataclass from typing import assert_never, Union @dataclass class PlusToken: pass @dataclass class DigitToken: digit: int Token = Union[PlusToken, DigitToken] def pretty_print(token: Token) -> str: match token: case PlusToken(): return 'PLUS <+>' case DigitToken(digit=digit): return f'DIGIT <{digit}>' case _: assert_never(token)
这种方式下,mypy会严格检查Union的所有成员是否都被模式匹配覆盖,遗漏时直接报错。
3. 给ABC基类添加抽象方法+配合类型检查
默认无抽象方法的ABC类可以被实例化,给Token添加抽象方法后,就能强制子类实现方法,同时阻止基类被实例化。结合assert_never可以辅助验证:
from abc import ABC, abstractmethod from dataclasses import dataclass from typing import assert_never @dataclass class Token(ABC): @abstractmethod def accept(self, visitor) -> str: pass @dataclass class PlusToken(Token): def accept(self, visitor) -> str: return visitor.visit_plus(self) @dataclass class DigitToken(Token): digit: int def accept(self, visitor) -> str: return visitor.visit_digit(self) def pretty_print(token: Token) -> str: match token: case PlusToken(): return 'PLUS <+>' case DigitToken(digit=digit): return f'DIGIT <{digit}>' case _: assert_never(token)
添加抽象方法后,Token无法被实例化,case _:分支只会匹配未处理的子类。不过这种方式mypy不会主动检查所有子类是否被覆盖,建议搭配密封类一起使用。
内容的提问来源于stack exchange,提问作者Evgeniy Slobodkin

