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

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中共享实例。

实现步骤:

  1. 完善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()
  1. 在路由中通过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用于存储应用级状态,适合存放全局共享的实例。

实现步骤:

  1. 应用初始化时绑定StateManager实例:
app = FastAPI()
app.state.state_manager = StateManager()
  1. 在路由中通过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实现模型周期更新:

  1. 安装依赖:
pip install apscheduler
  1. 集成定时任务:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 14:40:09