Python类型提示:如何实现枚举参数对应的函数重载强制参数校验
问题描述
定义了一个继承自StrEnum的指令枚举类ServerCommand,其中仅SERVER_CONFIRMATION取值为空字符串,其余枚举项对应具体指令值:
from enum import StrEnum class ServerCommand(StrEnum): SERVER_CONFIRMATION = "" SERVER_MOVE = "102 MOVE" SERVER_TURN_LEFT = "103 TURN LEFT" SERVER_TURN_RIGHT = "104 TURN RIGHT" SERVER_PICK_UP = "105 GET MESSAGE" SERVER_LOGOUT = "106 LOGOUT" SERVER_KEY_REQUEST = "107 KEY REQUEST" SERVER_OK = "200 OK" SERVER_LOGIN_FAILED = "300 LOGIN FAILED" SERVER_SYNTAX_ERROR = "301 SYNTAX ERROR" SERVER_LOGIC_ERROR = "302 LOGIC ERROR" SERVER_KEY_OUT_OF_RANGE_ERROR = "303 KEY OUT OF RANGE"
随后编写抽象类CommandCreator,其中create_message方法要求:传入SERVER_CONFIRMATION时必须附带int类型的confirmation_number参数;传入其他枚举项时无需该参数。
尝试用typing.overload实现重载,但当前写法仍允许不带confirmation_number就传入SERVER_CONFIRMATION的调用:
from abc import ABC, abstractmethod from typing import overload, Literal class CommandCreator(ABC): @overload @abstractmethod def create_message( self, cmd: Literal[ServerCommand.SERVER_CONFIRMATION], confirmation_number: int ) -> bytes: pass @overload @abstractmethod def create_message(self, cmd: ServerCommand) -> bytes: pass @abstractmethod def create_message(self, cmd: ServerCommand, confirmation_number: int | None = None) -> bytes: pass
比如以下调用不会被类型检查器拦截:
command_creator.create_message(ServerCommand.SERVER_CONFIRMATION)
手动定义排除SERVER_CONFIRMATION的Literal类型不够优雅,需要更优的解决方案。
解决方案
问题根源是第二个重载的cmd: ServerCommand包含了SERVER_CONFIRMATION,导致类型检查器认为无参数调用合法。我们可以通过精准限定第二个重载的指令类型解决,利用Python 3.12+(或typing_extensions库)的Exclude类型优雅排除目标枚举项:
步骤1:定义排除后的类型别名
from typing import StrEnum, ABC, abstractmethod, overload, Literal from typing_extensions import Exclude # Python 3.12+可直接用typing.Exclude class ServerCommand(StrEnum): # 枚举定义不变 SERVER_CONFIRMATION = "" SERVER_MOVE = "102 MOVE" # ... 其他枚举项 ... # 定义排除SERVER_CONFIRMATION的枚举类型 NonConfirmationServerCommand = Exclude[ServerCommand, Literal[ServerCommand.SERVER_CONFIRMATION]]
步骤2:修正重载方法
调整重载的类型限定,让第二个重载仅匹配非SERVER_CONFIRMATION的指令:
class CommandCreator(ABC): @overload @abstractmethod def create_message( self, cmd: Literal[ServerCommand.SERVER_CONFIRMATION], confirmation_number: int ) -> bytes: pass @overload @abstractmethod def create_message(self, cmd: NonConfirmationServerCommand) -> bytes: pass @abstractmethod def create_message(self, cmd: ServerCommand, confirmation_number: int | None = None) -> bytes: # 运行时校验,防止类型检查未覆盖的情况 if cmd is ServerCommand.SERVER_CONFIRMATION: if confirmation_number is None: raise ValueError("confirmation_number is required for SERVER_CONFIRMATION") return f"{confirmation_number}".encode() else: return cmd.value.encode()
效果
现在类型检查器会拦截非法调用:
# 类型检查报错:缺少必填参数confirmation_number command_creator.create_message(ServerCommand.SERVER_CONFIRMATION)
合法调用会被正确识别:
# 合法:传入SERVER_CONFIRMATION时带参数 command_creator.create_message(ServerCommand.SERVER_CONFIRMATION, 123) # 合法:其他指令无需额外参数 command_creator.create_message(ServerCommand.SERVER_MOVE)
兼容旧Python版本
如果无法使用Exclude,可以通过Union结合Literal手动生成非确认指令的类型(枚举更新时需同步修改):
NonConfirmationServerCommand = Literal[ ServerCommand.SERVER_MOVE, ServerCommand.SERVER_TURN_LEFT, ServerCommand.SERVER_TURN_RIGHT, # ... 其他非SERVER_CONFIRMATION的枚举项 ... ]
内容的提问来源于stack exchange,提问作者Hryhorii Biloshenko
相关产品推荐
相关产品推荐

