FastAPI中如何跨方法传递带状态的类实例(适配ML模型更新)
问题描述
我正在查阅FastAPI的Depends文档及官方示例:
from typing import Annotated from fastapi import Depends, FastAPI app = FastAPI() async def common_parameters(q: str | None = None, skip: int = 0, limit: int = 100): return {"q": q, "skip": skip, "limit": limit} @app.get("/items/") async def read_items(commons: Annotated[dict, Depends(common_parameters)]): return commons
但我的使用场景是部署一个需按周期(小时级、日级等)更新的ML模型。文档中的方案依赖可调用函数,它会被缓存而非每次调用都生成。我的场景不需要每次调用都初始化销毁资源,而是需要一个带状态的自定义类,让ML模型作为类属性可被定时或异步更新,且/invocations接口能使用更新后的模型提供服务。
目前我使用全局变量实现,在单脚本应用中运行良好,但随着应用扩展要使用Router时,我担心全局状态会引发问题。
请问是否有合适的方式在FastAPI的多个方法间传递带状态的类实例?
以下是示例类及接口代码:
import os import boto3 import joblib import pandas as pd from fastapi import FastAPI, status from fastapi.responses import JSONResponse from pydantic import BaseModel class InferenceRequest(BaseModel): feature1: float feature2: float class StateManager: def __init__(self): self.bucket = os.environ.get("BUCKET_NAME", "artifacts_bucket") self.s3_model_path = "./model.joblib" self.local_model_path = './model.joblib' self.s3 = None self.model = None def get_clients(self): self.s3 = boto3.client('s3') def download_model(self): if not self.s3: self.get_clients() self.s3.download_file(self.bucket, self.s3_model_path, self.local_model_path) self.model = joblib.load(self.local_model_path) # 当前的全局变量实现 state = StateManager() state.download_model() app = FastAPI() @app.post("/invocations") def invocations(request: InferenceRequest): input_data = pd.DataFrame(dict(request), index=[0]) try: predictions = state.model.predict(input_data) return JSONResponse({"predictions": predictions.tolist()}, status_code=status.HTTP_200_OK) except Exception as e: return JSONResponse({"error": str(e)}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR)
解决方案
方法1:使用类依赖(Depends)
FastAPI的类依赖默认是单例模式(只会初始化一次),刚好适配带状态的场景,且能在路由和Router中共享实例。
实现步骤:
- 完善
StateManager类的初始化与模型更新逻辑:
class StateManager: def __init__(self): self.bucket = os.environ.get("BUCKET_NAME", "artifacts_bucket") self.s3_model_path = "./model.joblib" self.local_model_path = './model.joblib' self.s3 = boto3.client('s3') self.model = None # 初始化时加载模型 self.download_model() def download_model(self): self.s3.download_file(self.bucket, self.s3_model_path, self.local_model_path) self.model = joblib.load(self.local_model_path) def update_model(self): # 用于定时更新模型的方法 self.download_model()
- 在路由中通过
Depends注入实例:
from typing import Annotated from fastapi import Depends, FastAPI, APIRouter app = FastAPI() router = APIRouter(prefix="/v1") # 定义依赖项 def get_state_manager(): return StateManager() # 简化类型注解 StateDep = Annotated[StateManager, Depends(get_state_manager)] # 主应用路由 @app.post("/invocations") def invocations(request: InferenceRequest, state: StateDep): input_data = pd.DataFrame(dict(request), index=[0]) try: predictions = state.model.predict(input_data) return JSONResponse({"predictions": predictions.tolist()}, status_code=status.HTTP_200_OK) except Exception as e: return JSONResponse({"error": str(e)}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR) # Router中的路由共享同一实例 @router.get("/model-status") def get_model_status(state: StateDep): return {"model_loaded": state.model is not None} app.include_router(router)
注意:FastAPI的类依赖默认缓存实例(单例),所有路由都会拿到同一个
StateManager对象,无需担心全局变量的问题。
方法2:使用FastAPI的app.state属性
FastAPI提供app.state用于存储应用级状态,适合存放全局共享的实例。
实现步骤:
- 应用初始化时绑定
StateManager实例:
app = FastAPI() app.state.state_manager = StateManager()
- 在路由中通过
request.app.state访问实例:
from fastapi import Request @app.post("/invocations") def invocations(request: Request, req: InferenceRequest): state_manager = request.app.state.state_manager input_data = pd.DataFrame(dict(req), index=[0]) try: predictions = state_manager.model.predict(input_data) return JSONResponse({"predictions": predictions.tolist()}, status_code=status.HTTP_200_OK) except Exception as e: return JSONResponse({"error": str(e)}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR) # Router中同样可通过Request访问 @router.get("/model-status") def get_model_status(request: Request): state_manager = request.app.state.state_manager return {"model_loaded": state_manager.model is not None}
方法3:添加定时更新逻辑
无论是类依赖还是app.state,都可以结合APScheduler实现模型周期更新:
- 安装依赖:
pip install apscheduler
- 集成定时任务:
from apscheduler.schedulers.asyncio import AsyncIOScheduler def setup_scheduler(app: FastAPI): scheduler = AsyncIOScheduler(timezone="UTC") # 每天凌晨2点更新模型 scheduler.add_job( func=app.state.state_manager.update_model, trigger="cron", hour=2 ) # 或者每小时更新一次 # scheduler.add_job( # func=app.state.state_manager.update_model, # trigger="interval", # hours=1 # ) scheduler.start() # 应用启动时初始化定时任务 @app.on_event("startup") async def startup_event(): setup_scheduler(app)
方案对比
- 类依赖(Depends):符合FastAPI设计理念,依赖注入方式清晰,无需传递
Request即可在不同模块共享实例。 - app.state:实现简单直接,适合快速搭建应用级全局状态,需通过
Request访问,Router中同样适用。
两种方案都能避免全局变量的潜在问题(如多进程环境下的状态不一致,若使用多进程部署,需确保模型更新逻辑进程安全,或使用共享存储/分布式缓存)。
内容的提问来源于stack exchange,提问作者jbuddy_13
相关产品推荐
相关产品推荐

