如何为FastAPI端点动态生成对应Pipeline的Pydantic输入模型?
解决方案:动态生成FastAPI Pipeline端点并解决mypy报错
你的泛型基类设计是合理的,完全符合Python类型提示的规范。要实现动态生成对应输入模型的API端点,同时解决mypy的静态检查报错,有两种可行方案:
方案1:为基类添加显式类属性存储类型信息
通过在BasePipeline中定义类属性,让子类明确指定输入输出类型,这样mypy能直接识别这些类型:
修改基类
from abc import ABC, abstractmethod from typing import Generic, TypeVar, Type from pydantic import BaseModel pipeline_input_T = TypeVar("pipeline_input_T", bound=BaseModel) pipeline_output_T = TypeVar("pipeline_output_T", bound=BaseModel) class BasePipeline(Generic[pipeline_input_T, pipeline_output_T], ABC): # 定义类属性,由子类继承时赋值 input_type: Type[pipeline_input_T] output_type: Type[pipeline_output_T] @abstractmethod async def run(self, input_data: pipeline_input_T) -> pipeline_output_T: pass
子类实现
class SomeItem(BaseModel): a: str b: int class SomePipeline(BasePipeline[SomeItem, bytes]): # 显式指定类型属性 input_type = SomeItem output_type = bytes async def run(self, input_data: SomeItem) -> bytes: # 业务逻辑实现 return b"processed_result"
动态生成端点
用一个辅助函数封装端点逻辑,通过cast告诉mypy类型的实际信息:
from fastapi import FastAPI, status from typing import Dict, Type, cast app = FastAPI() # 假设你的应用实例包含pipelines字典 application = ... def create_pipeline_endpoint(pipeline: BasePipeline): # 用cast让mypy识别变量的类型身份 InputType = cast(Type[BaseModel], pipeline.input_type) OutputType = cast(Type, pipeline.output_type) async def endpoint(input_data: InputType) -> OutputType: return await pipeline.run(input_data) return endpoint # 遍历生成所有Pipeline端点 for name, pipeline in application.pipelines.items(): endpoint = create_pipeline_endpoint(pipeline) app.post( f"/pipelines/{name}/run", tags=["Pipeline"], status_code=status.HTTP_200_OK, response_model=pipeline.output_type # 显式指定响应模型,提升API文档准确性 )(endpoint)
方案2:自动提取泛型类型参数(无需子类手动指定)
如果不想让每个子类都手动写input_type和output_type,可以通过Python的泛型元数据自动提取类型参数:
修改基类添加自动提取方法
from abc import ABC, abstractmethod from typing import Generic, TypeVar, Type from pydantic import BaseModel pipeline_input_T = TypeVar("pipeline_input_T", bound=BaseModel) pipeline_output_T = TypeVar("pipeline_output_T", bound=BaseModel) class BasePipeline(Generic[pipeline_input_T, pipeline_output_T], ABC): @classmethod def get_input_type(cls) -> Type[pipeline_input_T]: # 从类的泛型基元数据中提取输入类型 for base in cls.__orig_bases__: if hasattr(base, "__origin__") and base.__origin__ is BasePipeline: return base.__args__[0] raise ValueError("无法提取Pipeline的输入类型") @classmethod def get_output_type(cls) -> Type[pipeline_output_T]: # 从类的泛型基元数据中提取输出类型 for base in cls.__orig_bases__: if hasattr(base, "__origin__") and base.__origin__ is BasePipeline: return base.__args__[1] raise ValueError("无法提取Pipeline的输出类型") @abstractmethod async def run(self, input_data: pipeline_input_T) -> pipeline_output_T: pass
动态生成端点(修改辅助函数)
def create_pipeline_endpoint(pipeline: BasePipeline): # 调用类方法获取类型,并用cast消除mypy警告 InputType = cast(Type[BaseModel], pipeline.__class__.get_input_type()) OutputType = cast(Type, pipeline.__class__.get_output_type()) async def endpoint(input_data: InputType) -> OutputType: return await pipeline.run(input_data) return endpoint
为什么之前的get_input_type会报错?
mypy在静态检查时,会严格区分类型和值:你之前定义的get_input_type返回的是一个变量(值),mypy无法将其视为合法的类型注解。而上面两种方案中,要么用类属性(属于类定义层面的类型信息),要么通过泛型元数据提取类型,再用cast明确告诉mypy变量的类型身份,就能解决报错问题。
内容的提问来源于stack exchange,提问作者Roland Deschain
相关产品推荐
相关产品推荐

