作者:后端工程师的学习之路
难度:进阶 · 适合有 FastAPI 基础的同学
阅读时间:约 15 分钟


前言

在入门阶段,我们用 FastAPI 写过 CRUD、Pydantic 模型、简单的路由。但在生产项目中,面对的挑战远不止这些:依赖管理混乱、数据库会话泄露、认证体系脆弱、权限控制缺失、文件上传与缓存架构散乱——这些正是拦在"能跑"与"能上线"之间的坎。

本文从依赖注入原理出发,贯穿 SQLAlchemy 2.0 + 异步 ORMJWT 双令牌RBAC 权限控制中间件链路Redis 缓存文件上传后台任务等主题,完整呈现一个生产级 FastAPI 项目的架构思路。所有代码可直接当作脚手架使用。


一、依赖注入系统深度解析

FastAPI 的依赖注入(DI)是整个框架的灵魂。它远比表面看起来强大。

1.1 Depends 的工作原理

from fastapi import Depends, FastAPI, Query

app = FastAPI()

# 依赖可以是函数
def pagination(
    page: int = Query(1, ge=1),
    size: int = Query(20, ge=1, le=100),
):
    return {"page": page, "size": size}

@app.get("/items")
def list_items(paginator: dict = Depends(pagination)):
    return paginator

Depends 做的核心三件事:

  1. 解析子依赖:递归解析所有 Depends 链,构造 DAG(有向无环图)
  2. 缓存结果:在同一请求周期内,同一个依赖只执行一次(按参数路径缓存)
  3. 生命周期管理:请求结束后自动触发 yield 之后的清理逻辑

1.2 可调用类依赖

函数依赖够用,但类依赖更强大——可以携带状态、注入其他服务:

from fastapi import Depends, HTTPException, status
from typing import Optional

class PermissionChecker:
    def __init__(self, required_role: str):
        self.required_role = required_role

    def __call__(self, user: "UserOut" = Depends(get_current_user)):
        if user.role != self.required_role:
            raise HTTPException(
                status_code=status.HTTP_403_FORBIDDEN,
                detail=f"需要 {self.required_role} 角色",
            )
        return user

# 路由中使用
require_admin = PermissionChecker("admin")

@app.get("/admin/dashboard")
def admin_dashboard(
    user=Depends(require_admin),  # 直接传入可调用对象
):
    return {"message": "欢迎管理员", "user": user}

1.3 全局依赖

某些依赖(如认证)需要作用于所有路由,不必在每个路由上都写一遍 Depends

app = FastAPI(dependencies=[Depends(verify_api_key)])

局部路由组也可以挂载:

router = APIRouter(prefix="/api/v1", dependencies=[Depends(verify_token)])

1.4 yield 依赖与资源清理

这是 FastAPI DI 最优雅的设计之一——用生成器管理资源生命周期

from sqlalchemy.ext.asyncio import AsyncSession

async def get_db() -> AsyncSession:
    async with async_session_factory() as session:
        yield session
        # 离开 with 块时 session 自动关闭

FastAPI 保证:即使路由抛出异常,yield 之后的清理代码也一定会执行


二、数据库集成:SQLAlchemy 2.0 + Alembic + 异步 ORM

2.1 项目结构

app/
  db/
    __init__.py
    base.py          # DeclarativeBase
    session.py       # 引擎、会话工厂
    migrations/      # Alembic 目录
  models/
    user.py
    item.py
  repositories/
    user_repo.py
  schemas/
    user.py
    item.py

2.2 引擎与会话工厂

# app/db/session.py
from sqlalchemy.ext.asyncio import (
    AsyncSession,
    async_sessionmaker,
    create_async_engine,
)
from app.core.config import settings

engine = create_async_engine(
    settings.DATABASE_URL,  # e.g. postgresql+asyncpg://user:pass@localhost/db
    echo=settings.DB_ECHO,
    pool_size=20,
    max_overflow=10,
)

async_session_factory = async_sessionmaker(
    engine,
    class_=AsyncSession,
    expire_on_commit=False,
)

2.3 声明式基类和模型

# app/db/base.py
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
from datetime import datetime
from sqlalchemy import func

class Base(DeclarativeBase):
    pass

class TimestampMixin:
    created_at: Mapped[datetime] = mapped_column(
        server_default=func.now()
    )
    updated_at: Mapped[datetime] = mapped_column(
        server_default=func.now(), onupdate=func.now()
    )
# app/models/user.py
from app.db.base import Base, TimestampMixin
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy import String, Boolean, Enum as SAEnum
import enum

class UserRole(str, enum.Enum):
    ADMIN = "admin"
    EDITOR = "editor"
    VIEWER = "viewer"

class User(Base, TimestampMixin):
    __tablename__ = "users"

    id: Mapped[int] = mapped_column(primary_key=True)
    email: Mapped[str] = mapped_column(
        String(255), unique=True, index=True
    )
    hashed_password: Mapped[str] = mapped_column(String(255))
    role: Mapped[UserRole] = mapped_column(
        SAEnum(UserRole), default=UserRole.VIEWER
    )
    is_active: Mapped[bool] = mapped_column(Boolean, default=True)

2.4 Alembic 迁移配置

alembic init -t async app/db/migrations

编辑 env.py 使用异步引擎:

from app.db.base import Base
from app.core.config import settings
from app.models import User, Item  # 确保所有模型被导入

target_metadata = Base.metadata

async def run_migrations_online():
    connectable = create_async_engine(settings.DATABASE_URL)

    async with connectable.connect() as connection:
        await connection.run_sync(do_run_migrations)

常用命令:

alembic revision --autogenerate -m "add user table"
alembic upgrade head

2.5 异步查询示例

from sqlalchemy import select
from app.db.session import async_session_factory
from app.models.user import User

async def get_user_by_email(email: str) -> User | None:
    async with async_session_factory() as session:
        result = await session.execute(
            select(User).where(User.email == email)
        )
        return result.scalar_one_or_none()

三、Repository Pattern:数据访问层抽象

Repository Pattern 将数据逻辑与业务逻辑解耦,让单元测试时可以直接 mock 数据层。

3.1 泛型基类仓库

# app/repositories/base.py
from typing import Generic, TypeVar, Type
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, delete as sa_delete

ModelType = TypeVar("ModelType")

class BaseRepository(Generic[ModelType]):
    def __init__(self, model: Type[ModelType], session: AsyncSession):
        self.model = model
        self.session = session

    async def get(self, id: int) -> ModelType | None:
        result = await self.session.get(self.model, id)
        return result

    async def list(
        self, skip: int = 0, limit: int = 100
    ) -> list[ModelType]:
        result = await self.session.execute(
            select(self.model).offset(skip).limit(limit)
        )
        return list(result.scalars().all())

    async def create(self, **kwargs) -> ModelType:
        instance = self.model(**kwargs)
        self.session.add(instance)
        await self.session.commit()
        await self.session.refresh(instance)
        return instance

    async def delete(self, id: int) -> bool:
        result = await self.session.execute(
            sa_delete(self.model).where(self.model.id == id)
        )
        await self.session.commit()
        return result.rowcount > 0

3.2 具体仓库

# app/repositories/user_repo.py
from app.repositories.base import BaseRepository
from app.models.user import User
from sqlalchemy import select

class UserRepository(BaseRepository[User]):
    def __init__(self, session):
        super().__init__(User, session)

    async def get_by_email(self, email: str) -> User | None:
        result = await self.session.execute(
            select(User).where(User.email == email)
        )
        return result.scalar_one_or_none()

    async def count_active(self) -> int:
        result = await self.session.execute(
            select(User).where(User.is_active == True)
        )
        return len(result.scalars().all())

3.3 在路由中使用

@app.get("/users/{user_id}")
async def get_user(
    user_id: int,
    session: AsyncSession = Depends(get_db),
):
    repo = UserRepository(session)
    user = await repo.get(user_id)
    if not user:
        raise HTTPException(status_code=404)
    return user

四、JWT 认证:Access Token + Refresh Token

4.1 令牌工具函数

# app/core/security.py
from datetime import datetime, timedelta, timezone
from jose import jwt, JWTError
from passlib.context import CryptContext
from app.core.config import settings

pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")

def verify_password(plain: str, hashed: str) -> bool:
    return pwd_context.verify(plain, hashed)

def hash_password(plain: str) -> str:
    return pwd_context.hash(plain)

def create_access_token(data: dict) -> str:
    to_encode = data.copy()
    expire = datetime.now(timezone.utc) + timedelta(
        minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES
    )
    to_encode.update({"exp": expire, "type": "access"})
    return jwt.encode(
        to_encode,
        settings.SECRET_KEY,
        algorithm=settings.ALGORITHM,
    )

def create_refresh_token(data: dict) -> str:
    to_encode = data.copy()
    expire = datetime.now(timezone.utc) + timedelta(days=7)
    to_encode.update({"exp": expire, "type": "refresh"})
    return jwt.encode(
        to_encode,
        settings.SECRET_KEY,
        algorithm=settings.ALGORITHM,
    )

def decode_token(token: str) -> dict:
    try:
        payload = jwt.decode(
            token,
            settings.SECRET_KEY,
            algorithms=[settings.ALGORITHM],
        )
        return payload
    except JWTError:
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="无效的令牌",
        )

4.2 认证依赖

# app/api/deps.py
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from app.core.security import decode_token
from app.repositories.user_repo import UserRepository
from app.db.session import async_session_factory

oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/v1/auth/login")

async def get_current_user(
    token: str = Depends(oauth2_scheme),
) -> "UserOut":
    payload = decode_token(token)
    if payload.get("type") != "access":
        raise HTTPException(status_code=401, detail="需要 access token")

    user_id = payload.get("sub")
    if not user_id:
        raise HTTPException(status_code=401)

    async with async_session_factory() as session:
        repo = UserRepository(session)
        user = await repo.get(int(user_id))

    if not user or not user.is_active:
        raise HTTPException(status_code=401)
    return user

4.3 登录与刷新端点

# app/api/v1/auth.py
router = APIRouter(prefix="/auth", tags=["认证"])

@router.post("/login")
async def login(form: OAuth2PasswordRequestForm = Depends()):
    async with async_session_factory() as session:
        repo = UserRepository(session)
        user = await repo.get_by_email(form.username)

    if not user or not verify_password(form.password, user.hashed_password):
        raise HTTPException(status_code=401, detail="邮箱或密码错误")

    return {
        "access_token": create_access_token({"sub": str(user.id)}),
        "refresh_token": create_refresh_token({"sub": str(user.id)}),
        "token_type": "bearer",
    }

@router.post("/refresh")
async def refresh_token(refresh: RefreshTokenIn):
    payload = decode_token(refresh.refresh_token)
    if payload.get("type") != "refresh":
        raise HTTPException(status_code=401, detail="无效的 refresh token")

    return {
        "access_token": create_access_token({"sub": payload["sub"]}),
        "refresh_token": create_refresh_token({"sub": payload["sub"]}),
        "token_type": "bearer",
    }

五、RBAC 权限控制

5.1 角色检查依赖

# app/api/deps.py
from functools import wraps

def require_roles(*roles: str):
    """装饰器版本"""
    def decorator(func):
        @wraps(func)
        async def wrapper(*args, **kwargs):
            # 从依赖注入中获取 user
            return await func(*args, **kwargs)
        return wrapper
    return decorator

# 推荐:使用 Depends 版本
class RoleChecker:
    def __init__(self, allowed_roles: list[str]):
        self.allowed_roles = allowed_roles

    def __call__(self, user: User = Depends(get_current_user)):
        if user.role not in self.allowed_roles:
            raise HTTPException(
                status_code=403,
                detail=f"需要 {self.allowed_roles} 角色之一",
            )
        return user

allow_admin = RoleChecker(["admin"])
allow_editor_or_admin = RoleChecker(["admin", "editor"])

5.2 路由中使用

@app.get("/admin/users")
async def list_all_users(
    user: User = Depends(allow_admin),  # 仅 admin 可访问
    repo: UserRepository = Depends(get_user_repo),
):
    return await repo.list()

@app.put("/posts/{post_id}")
async def update_post(
    post_id: int,
    payload: PostUpdate,
    user: User = Depends(allow_editor_or_admin),
):
    # 业务逻辑...
    pass

5.3 细粒度权限

对于更复杂的权限(如"只能编辑自己的文章"),在业务层判断:

@app.patch("/posts/{post_id}")
async def patch_post(
    post_id: int,
    payload: PostUpdate,
    user: User = Depends(get_current_user),
    repo: PostRepository = Depends(get_post_repo),
):
    post = await repo.get(post_id)
    if not post:
        raise HTTPException(status_code=404)

    # 权限校验:admin 可编辑所有,普通用户只能编辑自己的
    if user.role != "admin" and post.author_id != user.id:
        raise HTTPException(status_code=403, detail="无权编辑此文章")

    return await repo.update(post_id, **payload.model_dump())

六、中间件实现

6.1 CORS 中间件

from fastapi.middleware.cors import CORSMiddleware

app.add_middleware(
    CORSMiddleware,
    allow_origins=settings.ALLOWED_ORIGINS,  # ["https://example.com"]
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)

6.2 请求日志与追踪 ID

import uuid
import time
import logging
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request

logger = logging.getLogger("access")

class RequestLogMiddleware(BaseHTTPMiddleware):
    async def dispatch(self, request: Request, call_next):
        request_id = str(uuid.uuid4())
        start = time.perf_counter()

        # 将 request_id 注入 request.state,路由中可访问
        request.state.request_id = request_id

        response = await call_next(request)

        elapsed = time.perf_counter() - start
        logger.info(
            "request_id=%s method=%s path=%s status=%d elapsed=%.3fms",
            request_id,
            request.method,
            request.url.path,
            response.status_code,
            elapsed * 1000,
        )
        response.headers["X-Request-ID"] = request_id
        return response

app.add_middleware(RequestLogMiddleware)

6.3 在路由中获取 request_id

@app.get("/trace")
async def trace(request: Request):
    return {"request_id": request.state.request_id}

七、Redis 缓存策略

7.1 缓存客户端

# app/core/cache.py
import json
from redis.asyncio import Redis
from app.core.config import settings

redis_client = Redis.from_url(
    settings.REDIS_URL,  # redis://localhost:6379/0
    decode_responses=True,
)

async def get_cache(key: str):
    data = await redis_client.get(key)
    return json.loads(data) if data else None

async def set_cache(key: str, value, expire: int = 300):
    await redis_client.setex(key, expire, json.dumps(value))

async def invalidate_cache(pattern: str):
    """按模式删除缓存,例如 invalidate_cache("user:*")"""
    keys = await redis_client.keys(pattern)
    if keys:
        await redis_client.delete(*keys)

7.2 缓存装饰器

from functools import wraps
import hashlib

def cached(expire: int = 300):
    def decorator(func):
        @wraps(func)
        async def wrapper(*args, **kwargs):
            # 根据函数名+参数构建缓存键
            key_data = f"{func.__name__}:{args}:{kwargs}"
            cache_key = hashlib.md5(key_data.encode()).hexdigest()

            cached_data = await get_cache(cache_key)
            if cached_data is not None:
                return cached_data

            result = await func(*args, **kwargs)
            await set_cache(cache_key, result, expire)
            return result
        return wrapper
    return decorator

# 使用
@app.get("/stats")
@cached(expire=60)
async def get_stats(repo=Depends(get_stats_repo)):
    # 耗时较长的统计查询
    return await repo.get_dashboard_stats()

7.3 缓存失效策略

写入数据时主动失效相关缓存:

@app.post("/users")
async def create_user(payload: UserCreate, repo=Depends(get_user_repo)):
    user = await repo.create(**payload.model_dump())
    # 使用户列表缓存失效
    await invalidate_cache("UserRepository:list:*")
    return user

八、文件上传与静态文件服务

8.1 文件上传

from fastapi import UploadFile, File
import aiofiles
import os

UPLOAD_DIR = "uploads"

@app.post("/upload")
async def upload_file(file: UploadFile = File(...)):
    # 安全检查:文件类型
    allowed_types = {"image/jpeg", "image/png", "application/pdf"}
    if file.content_type not in allowed_types:
        raise HTTPException(status_code=400, detail="不支持的文件类型")

    # 安全文件名
    safe_name = f"{uuid.uuid4()}_{file.filename}"
    file_path = os.path.join(UPLOAD_DIR, safe_name)

    # 流式写入
    async with aiofiles.open(file_path, "wb") as f:
        while chunk := await file.read(1024 * 1024):  # 1MB 块
            await f.write(chunk)

    return {
        "filename": safe_name,
        "size": os.path.getsize(file_path),
        "url": f"/static/{safe_name}",
    }

8.2 多文件上传

@app.post("/upload/multiple")
async def upload_multiple(files: list[UploadFile] = File(...)):
    results = []
    for file in files:
        # 处理每个文件...
        results.append({"filename": file.filename})
    return results

8.3 静态文件服务

from fastapi.staticfiles import StaticFiles

app.mount("/static", StaticFiles(directory="uploads"), name="static")

九、后台任务与 Celery

9.1 FastAPI BackgroundTasks 轻量方案

适合简单、短任务(发送邮件通知、写日志):

from fastapi import BackgroundTasks

def send_welcome_email(email: str):
    # 同步操作
    import time
    time.sleep(2)
    print(f"发送欢迎邮件到 {email}")

@app.post("/register")
async def register(user: UserCreate, bg: BackgroundTasks):
    # 创建用户...
    bg.add_task(send_welcome_email, user.email)
    return {"message": "注册成功,邮件稍后发送"}

9.2 Celery 集成(生产级)

# app/core/celery_app.py
from celery import Celery
from app.core.config import settings

celery_app = Celery(
    "worker",
    broker=settings.CELERY_BROKER_URL,     # redis://localhost:6379/1
    backend=settings.CELERY_RESULT_BACKEND,  # redis://localhost:6379/2
)

celery_app.conf.task_routes = {
    "tasks.email.*": {"queue": "email"},
    "tasks.report.*": {"queue": "report"},
}

# app/tasks/email.py
from app.core.celery_app import celery_app

@celery_app.task(name="send_verification_email")
def send_verification_email(user_id: int, email: str):
    # 复杂任务:渲染模板、与 SMTP 交互、重试逻辑
    pass

# 路由中调用
@app.post("/register")
async def register(user: UserCreate):
    new_user = await repo.create(...)
    send_verification_email.delay(new_user.id, new_user.email)
    return {"message": "注册成功"}

9.3 任务状态跟踪

@app.get("/tasks/{task_id}")
async def get_task_status(task_id: str):
    task = send_verification_email.AsyncResult(task_id)
    return {
        "task_id": task_id,
        "status": task.status,
        "result": task.result if task.ready() else None,
    }

十、全局异常处理

兜住所有未处理的异常,统一返回格式:

from fastapi import Request
from fastapi.responses import JSONResponse

class AppException(Exception):
    def __init__(self, code: int, message: str):
        self.code = code
        self.message = message

@app.exception_handler(AppException)
async def app_exception_handler(request: Request, exc: AppException):
    return JSONResponse(
        status_code=exc.code,
        content={"detail": exc.message, "code": exc.code},
    )

@app.exception_handler(Exception)
async def global_exception_handler(request: Request, exc: Exception):
    logger.error("Unhandled error: %s", exc, exc_info=True)
    return JSONResponse(
        status_code=500,
        content={"detail": "服务器内部错误"},
    )

十一、项目结构总览

将以上所有组合起来,一个生产级项目结构如下:

app/
  main.py                      # 应用入口
  core/
    config.py                  # Pydantic Settings
    security.py                # JWT + 密码哈希
    cache.py                   # Redis 客户端
    celery_app.py              # Celery 配置
    exceptions.py              # 自定义异常
  db/
    base.py                    # 声明式基类
    session.py                 # 异步引擎与会话工厂
    migrations/
  models/
    user.py, item.py
  schemas/
    user.py, item.py
  repositories/
    base.py                    # 泛型基类仓库
    user_repo.py, item_repo.py
  services/
    auth_service.py
    email_service.py
  api/
    deps.py                    # 通用依赖(认证、权限、DB)
    v1/
      auth.py, users.py, items.py
  tasks/
    email.py, report.py
  middleware/
    request_log.py
  uploads/

总结

本文从工程落地的角度,梳理了 FastAPI 进阶开发的 9 个核心板块:

模块 关键知识点
依赖注入 Depends 链、类依赖、yield 清理、全局依赖
数据库 SQLAlchemy 2.0 异步 ORM、Alembic 迁移
Repository 泛型基类、接口隔离、可测试性
认证 JWT 双令牌、OAuth2PasswordBearer
权限 RBAC、RoleChecker 可调用类
中间件 CORS、日志链路、请求 ID
缓存 Redis 装饰器、缓存失效策略
文件上传 流式写入、安全校验、静态挂载
后台任务 BackgroundTasks 轻量 / Celery 生产级

每个项目都有自己的上下文与约束,上面的模式并非银弹。但希望这份"脚手架"能为你的 FastAPI 生产级实践提供一个扎实的起点。

如果你有任何实战中的疑问或更好的实践,欢迎在评论区讨论。


本文代码示例基于 Python 3.11+、FastAPI 0.110+、SQLAlchemy 2.0+
下一篇预告:FastAPI 性能优化:数据库连接池、gRPC 微服务与异步流式响应