如何在Pydantic BaseModel中为第三方类实现泛型子类型注解支持
问题描述
已实现一个Pydantic注解包装器,可将JSON列表解析为NumPy数组,但无法支持指定dtype的校验需求——例如希望使用NumpyWrapper[np.float64]这样的语法同时校验数组类型。尝试用泛型改造时出现错误:‘类型“ndarray[Any, dtype[Unknown]]”已被特化’。
解决方案
要实现支持指定dtype的NumPy数组注解包装器,需要正确结合泛型与Pydantic的核心Schema机制,同时在验证逻辑中加入dtype校验。以下是修改后的完整代码:
from typing import Annotated, Any, Generic, TypeVar import numpy as np import numpy.typing as npt from pydantic import GetCoreSchemaHandler, GetJsonSchemaHandler from pydantic.json_schema import JsonSchemaValue from pydantic_core import core_schema # 定义受Numpy dtype约束的类型变量 DTypeLike = TypeVar("DTypeLike", bound=npt.DTypeLike) class _NumpyPydanticAnnotation(Generic[DTypeLike]): @classmethod def __get_pydantic_core_schema__( cls, source_type: Any, handler: GetCoreSchemaHandler, ) -> core_schema.CoreSchema: # 从泛型参数中提取目标dtype target_dtype = source_type.__args__[0] if hasattr(source_type, "__args__") else None def validate_from_list(value: list) -> np.ndarray: arr = np.array(value, dtype=target_dtype) # 校验实际生成的数组dtype是否匹配目标dtype(处理自动类型提升的情况) if not np.issubdtype(arr.dtype, target_dtype): raise ValueError(f"Array dtype {arr.dtype} does not match required dtype {target_dtype}") return arr def validate_numpy_array(value: np.ndarray) -> np.ndarray: if not np.issubdtype(value.dtype, target_dtype): raise ValueError(f"Array dtype {value.dtype} does not match required dtype {target_dtype}") return value from_list_schema = core_schema.chain_schema( [ core_schema.list_schema(), core_schema.no_info_plain_validator_function(validate_from_list), ] ) numpy_instance_schema = core_schema.chain_schema( [ core_schema.is_instance_schema(np.ndarray), core_schema.no_info_plain_validator_function(validate_numpy_array), ] ) return core_schema.json_or_python_schema( json_schema=from_list_schema, python_schema=core_schema.union_schema([numpy_instance_schema, from_list_schema]), serialization=core_schema.plain_serializer_function_ser_schema(lambda instance: instance.tolist()), ) @classmethod def __get_pydantic_json_schema__( cls, _core_schema: core_schema.CoreSchema, handler: GetJsonSchemaHandler ) -> JsonSchemaValue: # 沿用列表的JSON Schema定义 return handler(core_schema.list_schema()) # 定义泛型的Annotated包装类型 NumpyWrapper = Annotated[npt.NDArray[DTypeLike], _NumpyPydanticAnnotation[DTypeLike]]
关键说明
- 泛型绑定:定义
DTypeLike类型变量并绑定到npt.DTypeLike,确保传入的类型是合法的NumPy dtype。 - 提取目标dtype:通过
source_type.__args__从泛型参数中获取指定的目标dtype,作为校验依据。 - 双重校验逻辑:
- 从JSON列表解析时,直接用目标dtype创建数组,并校验实际生成的dtype是否匹配(避免自动类型提升导致的不符合预期)。
- 对已有的NumPy数组实例,直接校验其dtype是否符合要求。
- Schema兼容性:保持JSON Schema为列表类型,同时支持Python端的NumPy数组实例直接传入。
使用示例
from pydantic import BaseModel class SomeDataModel(BaseModel): float_array: NumpyWrapper[np.float64] int_array: NumpyWrapper[np.int32] # 合法场景 valid_data = {"float_array": [1.0, 2.5], "int_array": [1, 2]} model = SomeDataModel(**valid_data) print(model.float_array.dtype) # 输出: float64 print(model.int_array.dtype) # 输出: int32 # 非法场景:传入不兼容类型会触发错误 invalid_data = {"float_array": ["not a number"], "int_array": [1.5]} try: model = SomeDataModel(**invalid_data) except ValueError as e: print(e) # 抛出dtype不匹配或数组创建失败的错误信息
内容的提问来源于stack exchange,提问作者Roland Deschain
相关产品推荐
相关产品推荐

