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

如何在Pydantic Spark模型中指定枚举字段的Schema?

解决Pydantic Spark模型枚举字段Schema生成错误并转为字符串类型

方法1:使用Annotated+PlainSerializer强制序列化为字符串

通过Pydantic序列化工具,既保留输入时的Enum验证能力,又让生成Schema时将字段识别为字符串类型:

from pydantic_spark.base import SparkBase
from enum import Enum
from pydantic import Annotated, PlainSerializer

class TestEnum(Enum):
    SOMETHING = "something"
    OTHER_THING = "something else"

# 定义带序列化规则的类型,将Enum转为对应字符串值
StringEnum = Annotated[
    TestEnum,
    PlainSerializer(lambda x: x.value, return_type=str)
]

class TestModel(SparkBase):
    thing: StringEnum

my_model = TestModel(thing=TestEnum.SOMETHING)
print(my_model.spark_schema())
# 输出为StringType,无运行时错误

方法2:自定义字符串字段处理Enum转换

继承str类型,自定义验证与序列化逻辑,让SparkSchema直接识别为字符串:

from pydantic_spark.base import SparkBase
from enum import Enum
from pydantic import GetCoreSchemaHandler
from pydantic_core import core_schema

class TestEnum(Enum):
    SOMETHING = "something"
    OTHER_THING = "something else"

class EnumAsString(str):
    @classmethod
    def __get_pydantic_core_schema__(cls, source_type, handler: GetCoreSchemaHandler) -> core_schema.CoreSchema:
        return core_schema.no_info_wrap_validator_function(
            cls.validate,
            core_schema.str_schema(),
        )
    
    @classmethod
    def validate(cls, v: str | TestEnum) -> str:
        if isinstance(v, TestEnum):
            return v.value
        if v in [item.value for item in TestEnum]:
            return v
        raise ValueError(f"无效值,必须是{[item.value for item in TestEnum]}之一")

class TestModel(SparkBase):
    thing: EnumAsString

my_model = TestModel(thing=TestEnum.SOMETHING)
print(my_model.spark_schema())

方法3:重写spark_schema方法手动指定字段类型

直接覆盖模型的spark_schema方法,手动替换目标字段的类型为字符串:

from pydantic_spark.base import SparkBase
from enum import Enum
from pyspark.sql.types import StructType, StringType, StructField

class TestEnum(Enum):
    SOMETHING = "something"
    OTHER_THING = "something else"

class TestModel(SparkBase):
    thing: TestEnum

    def spark_schema(self) -> StructType:
        # 获取默认schema后替换指定字段类型
        base_schema = super().spark_schema()
        return StructType([
            StructField(field.name, StringType(), field.nullable)
            if field.name == "thing"
            else field
            for field in base_schema.fields
        ])

my_model = TestModel(thing=TestEnum.SOMETHING)
print(my_model.spark_schema())

内容的提问来源于stack exchange,提问作者WhatAmIDoing

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 08:17:10