You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让类型检查器支持多类型赋值,且Pydantic始终转为指定Enum类型?

解决Pydantic自定义Enum与类型检查器的兼容问题

首先修正你代码中的两处运行时错误:

  1. CustomEnumMeta里未定义names变量,需替换为list(self.__members__.keys());
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 06:17:33