如何让类型检查器支持多类型赋值,且Pydantic始终转为指定Enum类型?
解决Pydantic自定义Enum与类型检查器的兼容问题
首先修正你代码中的两处运行时错误:
CustomEnumMeta里未定义names变量,需替换为list(self.__members__.keys());CalculationOptions的默认值Solvers.runge_kutta不存在,应改为Solvers.runge_kutta34。
针对类型检查器的问题,以下是两种可行解决方案:
方案1:Annotated结合类型断言(轻量改造)
在CoercedEnum的注解中明确允许输入str/int/Solvers,保留BeforeValidator确保转换为Solvers实例,访问属性时通过类型断言消除检查器警告。
from typing import Any, Annotated, TypeVar, List, Type, Union from enum import Enum, EnumMeta from pydantic import BaseModel, Field, model_validator from pydantic.functional_validators import BeforeValidator from functools import partial class CustomEnumMeta(EnumMeta): """Matches an Enum member based on string value or integer value""" def __getitem__(self, name: Any) -> Any: members = list(self.__members__.values()) values: List[Any] = [a.value for a in members] flags: List[int] = [a.flag for a in members] names = list(self.__members__.keys()) # 修复未定义的names变量 try: name = int(name) except (ValueError, TypeError): pass if name in values: name = names[values.index(name)] elif name in flags: name = names[flags.index(name)] if not isinstance(name, str): raise ValueError(f"{name!s} is not an enumerated value of {type(self)!s}") return super().__getitem__(name) class FlagDataEnum(Enum, metaclass=CustomEnumMeta): """Adds data storage and an integer flag as well as a string name to Enum class""" def __init__(self, desc: Any, flag: int, *args: Any) -> None: self._value_ = desc self.flag = flag self.data = args[0] if args else None # 简化data赋值逻辑 # Generic "Enum" subclass type E = TypeVar("E", bound=Enum) def _coerce_value_to_data_enum(value: Any, enum_type: Type[E]) -> E: """Tries to convert a value to an instance of the given enum_type""" if isinstance(value, enum_type): return value else: return enum_type[value] class Solvers(FlagDataEnum): runge_kutta34 = "Runge-Kutta 3/4", 1, {'predictor':3,'corrector':4} runge_kutta78 = "Runge-Kutta 7/8", 2, {'predictor':7,'corrector':8} adams_bashforth = "Adams-Bashforth", 3 central_difference = "Central Difference", 4 # 明确标注允许的输入类型,保留验证器确保转换为Solvers CoercedEnum = Annotated[ Union[Solvers, str, int], BeforeValidator(partial(_coerce_value_to_data_enum, enum_type=Solvers)) ] class CalculationOptions(BaseModel): solver: CoercedEnum = Field(default=Solvers.runge_kutta34) init_conditions: List[int] = Field(default_factory=list) # 可选:添加后置验证器,确保solver最终为Solvers实例 @model_validator(mode='after') def ensure_solver_type(self) -> 'CalculationOptions': assert isinstance(self.solver, Solvers), "solver must be a Solvers instance" return self
使用时通过类型断言消除属性访问警告:
opts = CalculationOptions(solver=1) assert isinstance(opts.solver, Solvers) print(opts.solver.flag) # 类型检查器不再报错
方案2:自定义Pydantic类型(优雅的类型检查支持)
创建自定义Pydantic类型,让类型检查器直接识别输入允许str/int,但验证后的值一定是Solvers实例,无需额外断言。
from typing import Any, TypeVar, List, Type, Union from enum import Enum, EnumMeta from pydantic import BaseModel, Field, GetCoreSchemaHandler from pydantic_core import core_schema E = TypeVar("E", bound=Enum) class CoercedEnumType: def __class_getitem__(cls, enum_type: Type[E]) -> Any: class _CoercedEnum: @classmethod def __get_pydantic_core_schema__(cls, source_type: Any, handler: GetCoreSchemaHandler) -> core_schema.CoreSchema: # 定义核心schema:接受str/int/Solvers,转换为Solvers实例 return core_schema.no_info_wrap_validator_function( lambda v: enum_type[v] if not isinstance(v, enum_type) else v, core_schema.union_schema([ core_schema.is_instance_schema(enum_type), core_schema.str_schema(), core_schema.int_schema() ]) ) return _CoercedEnum # 修正后的CustomEnumMeta和FlagDataEnum同方案1 class CustomEnumMeta(EnumMeta): def __getitem__(self, name: Any) -> Any: members = list(self.__members__.values()) values: List[Any] = [a.value for a in members] flags: List[int] = [a.flag for a in members] names = list(self.__members__.keys()) try: name = int(name) except (ValueError, TypeError): pass if name in values: name = names[values.index(name)] elif name in flags: name = names[flags.index(name)] if not isinstance(name, str): raise ValueError(f"{name!s} is not an enumerated value of {type(self)!s}") return super().__getitem__(name) class FlagDataEnum(Enum, metaclass=CustomEnumMeta): def __init__(self, desc: Any, flag: int, *args: Any) -> None: self._value_ = desc self.flag = flag self.data = args[0] if args else None class Solvers(FlagDataEnum): runge_kutta34 = "Runge-Kutta 3/4", 1, {'predictor':3,'corrector':4} runge_kutta78 = "Runge-Kutta 7/8", 2, {'predictor':7,'corrector':8} adams_bashforth = "Adams-Bashforth", 3 central_difference = "Central Difference", 4 class CalculationOptions(BaseModel): solver: CoercedEnumType[Solvers] = Field(default=Solvers.runge_kutta34) init_conditions: List[int] = Field(default_factory=list)
这种方式下,类型检查器会自动识别:
- 赋值时允许传入
str/int/Solvers; - 访问
solver.flag或solver.value时,确认值为Solvers实例,无属性不存在警告。
额外优化建议
为提升枚举查找效率,可在CustomEnumMeta中预构建值和标志到成员的映射,避免每次查找都遍历所有成员:
class CustomEnumMeta(EnumMeta): def __new__(cls, name, bases, namespace): enum_cls = super().__new__(cls, name, bases, namespace) # 预构建值和标志到成员的映射 enum_cls._value_map = {member.value: member for member in enum_cls} enum_cls._flag_map = {member.flag: member for member in enum_cls} return enum_cls def __getitem__(self, name: Any) -> Any: # 优先通过名称查找 if isinstance(name, str) and name in self.__members__: return super().__getitem__(name) # 通过标志查找 try: name_int = int(name) if name_int in self._flag_map: return self._flag_map[name_int] except (ValueError, TypeError): pass # 通过值查找 if name in self._value_map: return self._value_map[name] raise ValueError(f"{name!s} is not an enumerated value of {type(self)!s}")
内容的提问来源于stack exchange,提问作者Beezum
相关产品推荐
相关产品推荐

