作者:后端工程师的学习之路
难度:进阶 · 适合有 FastAPI 基础的同学
阅读时间:约 15 分钟
前言
在入门阶段,我们用 FastAPI 写过 CRUD、Pydantic 模型、简单的路由。但在生产项目中,面对的挑战远不止这些:依赖管理混乱、数据库会话泄露、认证体系脆弱、权限控制缺失、文件上传与缓存架构散乱——这些正是拦在"能跑"与"能上线"之间的坎。
本文从依赖注入原理出发,贯穿 SQLAlchemy 2.0 + 异步 ORM、JWT 双令牌、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 做的核心三件事:
- 解析子依赖:递归解析所有
Depends链,构造 DAG(有向无环图) - 缓存结果:在同一请求周期内,同一个依赖只执行一次(按参数路径缓存)
- 生命周期管理:请求结束后自动触发
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 微服务与异步流式响应
评论