如何在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
相关产品推荐
相关产品推荐

