FastAPI
FastAPI
1、FastAPI 是什么
它是 Python 现代高性能 Web 框架,对标 Java 的 SpringBoot + SpringMVC。优点是速度快、自动生成接口文档、支持异步、自带类型校验、代码极度简洁,现在后端开发、AI 接口服务、微服务项目几乎都在用。
2、安装与启动
Java 用 Maven 或 Gradle 引入依赖;Python 只需要用 pip 安装,启动命令也非常简单,支持热重载,改代码自动刷新,开发效率很高。
Java(Maven)
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-web</artifactId>
</dependency>
Python(FastAPI)
pip install fastapi uvicorn
uvicorn main:app --reload
3、请求流程
和 SpringBoot 的请求处理流程几乎一致:请求到达 → 路由匹配 → 参数解析与校验 → 执行业务逻辑 → 返回响应。FastAPI 基于协程处理请求,高并发场景下性能更强。
Java Spring MVC:
Filter → DispatcherServlet → HandlerMapping → Controller → ViewResolver → 响应
FastAPI:
Middleware → Router匹配 → 参数解析(Pydantic 请求模型) → 依赖注入(Depends) → 路由函数 → 校验和转化JSON(Pydantic 响应模型) → 响应
from fastapi import FastAPI, Request
app = FastAPI()
# 中间件:请求进来最先执行,最后返回
@app.middleware("http")
async def log_middleware(request: Request, call_next):
print(f"1. 中间件前置: {request.method} {request.url}")
response = await call_next(request) # 进入路由处理
print(f"4. 中间件后置: {response.status_code}")
return response
# 依赖注入:中间件之后、路由函数之前执行
async def get_db():
print("2. 依赖注入: 获取数据库连接")
yield "db_connection"
print("5. 依赖清理: 释放数据库连接")
# 路由函数
@app.get("/hello")
async def hello(name: str, db: str = Depends(get_db)):
print(f"3. 路由处理: name={name}, db={db}")
return {"message": f"Hello {name}"}
执行顺序:
1. 中间件前置: GET /hello?name=world
2. 依赖注入: 获取数据库连接
3. 路由处理: name=world, db=db_connection
4. 中间件后置: 200
5. 依赖清理: 释放数据库连接
4、模块化路由-路由分层
将庞大的应用拆分为多个独立的业务模块,对应 SpringBoot 的 Controller 分包,FastAPI 支持路由拆分,不同业务模块分开管理,大型项目结构清晰、易于维护
Java(Controller)
@RestController
@RequestMapping("/user")
public class UserController {
@GetMapping("/list")
public List<User> list() {
return List.of(new User(1), new User(2));
}
}
Python(APIRouter)
from fastapi import FastAPI, APIRouter
app = FastAPI()
# 1.实例化路由对象 实际:各子路由对象,由各自不同的业务模块生产,主路由对象统一管理即可
v1_router = APIRouter(prefix="/api/v1") # 所有接口的 统一前缀
user_router = APIRouter(tags=['用户应用'])
item_router = APIRouter(tags=['商品应用'], prefix="/item") # prefix可写在 APIRouter中,也可写在 .include_router()中
# 2.通过该路由对象,注册请求函数
@user_router.get('/user/info')
async def get_info1():
return 'user: 内容获取成功!'
@item_router.get('/item/info')
async def get_info2():
return 'item: 内容获取成功!'
# 3.将路由对象,注册到 FastAPI应用 或 其他路由对象中
v1_router.include_router(user_router)
v1_router.include_router(item_router)
# app.include_router(user_router, prefix="/user")
app.include_router(v1_router)
5、查询参数、请求体
对应 SpringBoot 的 @RequestParam和@RequestBody,FastAPI 直接定义参数和用 Pydantic 接收 JSON,自带类型校验、默认值
Java
@GetMapping("/search")
public String search(@RequestParam String keyword, @RequestParam(defaultValue = "1") int page) {
return keyword;
}
public class User {
private String name;
private Integer age;
}
@PostMapping("/user")
public User create(@RequestBody User user) {
return user;
}
Python
### 查询参数 Query
@app.get("/search") # /search?keyword=haha&page=1
def search(
keyword: str | None = Query(default=None, max_length=50, description="搜索关键字") ,
page_size: int = Query(10, ge=1, le=100, alias="pageSize", description="每页条数"): # alias 对应前端具体的字段名
return {"keyword": keyword, "pageSize": page_size}
### 路径参数 Path
@app.get("/news/{news_id}") # /news/123
async def get_news(
news_id: int = Path(
..., # 必填项,必须传 ...
ge=1, # 校验:大于等于 1
title="新闻ID",
description="要查询的新闻唯一标识符"
)
):
return {"news_id": news_id}
### 请求体参数 通常使用pydantic 构建请求模型【自带验证等】
from pydantic import BaseModel
class User(BaseModel):
name: str
age: int
@app.post("/user")
def create_user(user: User):
return {"msg": "创建成功", "data": user}
6、Pydantic 数据校验
对标 Java 的参数校验框架,支持非空、长度、范围、格式等多种校验规则,接口更安全、更规范。
# ===== 常见约束速查 =====
"""
字符串:
min_length - 最小长度
max_length - 最大长度
pattern - 正则表达式
数字:
gt - 大于
ge - 大于等于
lt - 小于
le - 小于等于
multiple_of - 倍数
通用:
default - 默认值
default_factory - 默认值工厂函数(适合[]、{})
description - 字段描述(生成到文档)
examples - 示例值(生成到文档)
deprecated - 标记弃用
alias - 字段别名(如JSON的camelCase)
"""
# ===== 校验装饰器 =====
field_validator, model_validator
# 校验顺序
1.预处理阶段 @model_validator(mode="before")
2.逐字段类型解析与强制转换 # 即:字段上的校验
3.单字段业务校验 @field_validator # 即:单字段的校验函数
4.实例级跨字段校验 @model_validator(mode="after")
@model_validator(mode="before") # 针对整个模型的校验 【类方法】
def (cls, data)->dict:
# 注意:
1.预处理阶段:数据还是原始的字典,尚未进行任何类型转换
2.被装饰的函数,必须再加上 @classmethod,接收 cls 和原始 data 字典,且必须返回一个字典
@field_validator(字段) # 针对'字段'的 单字段业务校验 【类方法】
# 注意:
1.默认 mode='after',即在对应字段的类型转换之后,执行
2.被装饰的函数,必须再加上 @classmethod 装饰器,并且推荐加上强类型注解
@model_validator(mode="after") # 针对整个模型的校验 实例级跨字段校验 【对象方法】
# 注意:
1.在所有字段都解析和校验完毕后,执行
2.非常使用 做 多字段的联合校验
# ===== 序列装饰器 =====
@field_serializer(字段) # 针对'字段'的 自定义的导出格式 序列化函数【对象方法】
- 案例:
from pydantic import BaseModel, Field, field_validator, model_validator
from typing import Optional, List, Dict
from datetime import datetime
from enum import Enum
import re
# ===== 基础字段约束 =====
class ProductCreate(BaseModel):
name: str = Field(
..., # 默认值填... ... 表示必填
min_length=1,
max_length=100,
description="商品名称",
examples=["iPhone 15"]
)
price: float = Field(
...,
gt=0, # 大于0
le=999999.99, # 小于等于
description="价格"
)
stock: int = Field(
default=0,
ge=0, # 大于等于0
description="库存"
)
tags: List[str] = Field(
default_factory=list,
max_length=10, # 列表最多10个元素
)
metadata: Dict[str, str] = Field(default_factory=dict)
# 定义枚举类
class Status(str, Enum):
DRAFT = "draft"
PUBLISHED = "published"
ARCHIVED = "archived"
# 枚举字段
status: Status = Field(default=Status.DRAFT)
# 日期时间
publish_at: Optional[datetime] = None
# ===== 字段级校验 =====
class UserRegister(BaseModel):
username: str = Field(..., min_length=3, max_length=20)
password: str = Field(..., min_length=8)
confirm_password: str
email: str
phone: str
@field_validator("username")
@classmethod
def validate_username(cls, v: str) -> str:
"""用户名只能包含字母数字下划线"""
if not re.match(r"^[a-zA-Z0-9_]+$", v):
raise ValueError("用户名只能包含字母、数字和下划线")
return v.strip()
@field_validator("phone")
@classmethod
def validate_phone(cls, v: str) -> str:
"""手机号格式校验"""
if not re.match(r"^1[3-9]\d{9}$", v):
raise ValueError("手机号格式不正确")
return v
@field_validator("email")
@classmethod
def validate_email(cls, v: str) -> str:
v = v.strip().lower()
if "@" not in v:
raise ValueError("邮箱格式不正确")
return v
# ===== 跨字段校验 =====
class PasswordChange(BaseModel):
old_password: str
new_password: str
confirm_new_password: str
# 跨字段校验的最佳实践:在实例完全构建后进行
@model_validator(mode="after")
def validate_passwords(self):
# 1. 新密码不能与旧密码相同
if self.new_password == self.old_password:
raise ValueError("新密码不能与旧密码相同")
# 2. 确认密码必须和新密码一致
if self.confirm_new_password != self.new_password:
raise ValueError("两次输入的新密码不一致")
# 注意:必须返回 self
return self
# ===== 模型级校验(所有字段验证后) =====
class OrderCreate(BaseModel):
items: List[str]
coupon_code: Optional[str] = None
@model_validator(mode="after")
def validate_order(self):
"""订单级别校验"""
if len(self.items) == 0:
raise ValueError("订单必须包含至少一个商品")
if self.coupon_code and len(self.items) < 2:
raise ValueError("优惠券需要至少2件商品才能使用")
return self
# ===== 序列化(模型对象被导出时) =====
class Article(BaseModel):
title: str
created_at: datetime
# 自定义时间字段的输出格式
@field_serializer("created_at")
def serialize_created_at(self, dt: datetime, _info) -> str:
return dt.strftime("%Y-%m-%d %H:%M:%S")
# 测试
article = Article(title="Pydantic V2", created_at=datetime.now())
print(article.model_dump())
请求响应模型校验
数据校验与转换流程
# 整体流程
Middleware → Router匹配 → 参数解析(Pydantic 请求模型) → 依赖注入(Depends) → 路由函数 → 校验和转化JSON(Pydantic 响应模型) → 响应
# schema/user.py 请求模型 和 响应模型
# 方案对比:
传统的方案:
手动转化【ORM对象 -> Json】 [jsonable_encoder] 需要先转成 Python 字典,再转成 JSON
实际方案:
Pydantic 响应模型 [from_attributes] + 路由中[response_model] 是直接读取ORM对象属性
并交由Rust引擎输出,省去了中间环节,内存占用和 CPU 消耗大幅降低。
# 核心原理:转换:Pydantic 的 from_attributes 机制
这是整个流程中最核心的一步!FastAPI 底层调用了 Pydantic 的验证器:
逻辑:
Pydantic 看到 响应模型 配置了 model_config = ConfigDict(from_attributes=True)
动作:
Pydantic 遍历 数据库模型类 列表中的每一个 ORM 对象,然后像查字典一样去读取对象的属性(例如执行 getattr(orm_obj, "name"))
过滤:
它只提取 响应模型 中明确声明的字段(id, name, sort_order),
自动丢弃 ORM 对象里的其他所有内部属性,并将其转换为一个全新的、纯粹的 响应模型 Pydantic 对象
转化:JSON
FastAPI底层使用Rust编写的高性能序列化引擎(如 orjson)
将这个纯净的 Pydantic 对象列表 瞬间 转化为JSON字符串,并通过 HTTP 响应发送给前端
from pydantic import BaseModel, EmailStr, ConfigDict
# 存放于 schema/user.py 注:请求和响应模型 字段重复,还可以提取公共字段 父类
# 请求模型:用于接收前端传来的数据,自动进行类型校验
class UserCreate(BaseModel):
name: str
email: EmailStr # 自动校验是否为合法的邮箱格式 但需要安装邮箱扩展 pip install 'pydantic[email]'
age: int | None = None # 可选字段
# 响应模型:用于格式化返回给前端的数据,隐藏敏感信息
class UserResponse(BaseModel):
id: int
name: str
email: str
age: int | None = None
class Config:
# 关键:允许从 SQLAlchemy ORM 对象中读取数据
from_attributes = True
# 它告诉Pydantic:传入的不再是字典,而是一个带有属性的ORM对象(如user.name),请通过属性名去提取数据
# 新版本写法 - 设置:允许从ORM对象读取属性
model_config = ConfigDict(
alias_generator=to_camel, # 输出 序列化时,自动转为驼峰
populate_by_name=True, # 允许在Python中 同时使用属性字段名和别名 即:下划线和驼峰
from_attributes=True # 允许从 ORM对象属性读取 转换
)
# response_model:指定响应模型
@app.post("/users/", response_model=UserResponse, status_code=201)
async def create_user(
user_data: UserCreate, # FastAPI 会自动将 JSON Body 解析为 UserCreate对象
db: AsyncSession = Depends(get_db)
):
new_user = User(**user_data.model_dump()) # 使用 ** 解包传入数据,非常优雅
db.add(new_user) # 加入待提交队列,还没写入数据库
await db.commit() # 有get_db 兜底 提交事务,不用再手动提交事务
return new_user
常见Pydantic方法
Pydantic模型,本质上就是Python类
from pydantic import BaseModel, EmailStr, ConfigDict
### 实例【对象】方法
# 1.数据序列化 Pydantic模型对象 ---> 字典/JSON
.model_dump() # 转换为Python字典
# 特有参数:
mode='json' # 是将字典中的内部值,转成JSON格式
.model_dump_json() # 直接序列化为 JSON 字符串
# 参数:
include # 只包含 指定的字段
exclude # 排除 指定的字段
# 下三个 默认都为 False
exclude_unset # 排除未显式设置的字段 但不能排除 name=None
exclude_none # 排除值为None的字段 不管是默认值,还是前端显示赋值为None的
# 【通常:作为更新数据的请求模型时,设置该两项为True,实现动态局部更新】
by_alias # 输出的字典属性名,是否 用 模型类 字段 定义的别名
# 【默认为False,默认不自动使用别名进行序列化】 通常设置为 False
# 2.数据验证 字典/JSON ---> Pydantic模型对象
.model_validate(data) # 从Python字典【开启配置后,运行从ORM对象创建】 创建并校验 模型实例
.model_validate_json(json_str) # 从JSON字符串 创建并校验 模型实例
# 3.复制现有的Pydantic模型实例
.model_copy()
# 参数
update (可选): 传入一个字典,用于在创建副本时,修改或添加 特定字段的值
# 注意:该字典 不会再进行字段类型校验
deep (可选, 默认False): 深拷贝,避免修改副本时影响原对象
### 类属性 配置 通常定义在 Pydantic模型类 内部
model_config = ConfigDict(
extra='forbid', # 额外字段处理:传入多余字段 即抛出ValidationError
# 默认:为 'ignore':直接忽略多余字段
alias_generator=to_camel, # 别名生成器: 输出 序列化时,自动转为驼峰
populate_by_name=True, # 允许在Python中 同时使用属性字段名和别名 即:下划线和驼峰
from_attributes=True, # ORM 模式: 允许从ORM对象【属性读取】,创建并校验 模型实例
validate_assignment=True, # 赋值时验证: 开启后,修改属性时也要校验 【通常不开启,性能代价】
# 默认情况下,Pydantic 只在模型初始化(创建实例)时进行校验
)
# 注意:model_config 支持继承,子类和父类的配置,会进行深度合并
### 类方法
模型类.model_json_schema() # 根据Pydantic模型,自动生成符合JSON Schema规范的字典描述
- 案例
# eg:1.组装数据
# 先从 ORM 对象验证,并转化Pydantic模型
base_data = NewsDetailOut.model_validate(news_detail) # news_detail 是 ORM对象
# 复制更新Pydantic模型【复制并注入related_news(ORM对象)】,然后强制重新验证整个模型
detail_data = NewsDetailOut.model_validate(
base_data.model_copy(update={"related_news": related_news_list})
)
# eg:2.生成JSON Schema
class User(BaseModel):
username: str = Field(..., title="用户名", min_length=1, max_length=32)
age: int = Field(0, ge=0, le=150, description="用户年龄")
schema = User.model_json_schema()
print(schema)
# 结果:
{
"title": "User",
"type": "object",
"properties": {
"username": {
"title": "用户名",
"minLength": 1,
"maxLength": 32,
"type": "string"
},
"age": {
"title": "Age",
"description": "用户年龄",
"default": 0,
"minimum": 0,
"maximum": 150,
"type": "integer"
}
},
"required": ["username"]
}
7、自动生成接口文档
SpringBoot 需要集成 Swagger;FastAPI 零配置自带文档,相当于开箱即用的 Swagger。
from fastapi import FastAPI
app = FastAPI(
title="我的API",
version="1.0.0",
description="这是API描述"
)
@app.get("/hello")
async def hello(name: str = "World"):
return {"message": f"Hello {name}"}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8100)
http://127.0.0.1:8000/docs
http://127.0.0.1:8000/redoc
8、依赖注入
对应 SpringBoot @Autowired,FastAPI 使用 Depends,用于数据库连接、身份认证与权限控制、公共逻辑
Java
@Service
public class UserService {}
@RestController
public class UserController {
@Autowired
private UserService userService;
}
Python
from fastapi import APIRouter, Depends, HTTPException
def get_db():
db = "数据库连接"
yield db
db = "关闭连接"
@app.get("/data")
def get_data(db = Depends(get_db)): # 注:依赖注入 Depends只能出现在 Router(路由层: 即装饰器修饰的) 或者 其他被 Depends 调用的依赖函数 中
return {"db": db}
======== 其他被 Depends 调用的依赖函数 案例
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
# 依赖函数 1:解析 Token
async def get_current_user(token: str = Depends(oauth2_scheme)):
# 1. 解码 Token
# 2. 查数据库验证用户是否存在
user = await verify_token(token) # 需要额外封装,采用PyJWT等库 解码并验证签名、过期时间
if not user:
raise HTTPException(status_code=401, detail="Invalid token")
return user # <--- 这个 user 对象会被注入到路由中
# 依赖函数 2:检查权限 (依赖于上面的 get_current_user)
def get_current_active_admin(current_user: User = Depends(get_current_user)): # 其他被 Depends 调用的依赖函数
if current_user.role != "admin":
raise HTTPException(status_code=403, detail="Not enough permissions")
return current_user
认证与鉴权
- 原理流程:
# 生产环境的标准实践方案
在成熟的 FastAPI 生产项目中,通常采用 JWT + OAuth2 + 依赖注入 的 三位一体 防护体系:
1.定义安全方案:
使用 OAuth2PasswordBearer(tokenUrl="token") ,声明Token获取地址
2.编写认证依赖:
创建一个 get_current_user 异步函数,接收 Token,使用 PyJWT等库解码并验证签名、过期时间
如果验证失败,抛出 HTTPException(401);如果成功,返回用户信息字典
3.编写权限依赖(可选): # 核心:是支持嵌套与复杂权限(RBAC)
创建如 require_admin 的依赖,内部调用 get_current_user,并检查用户角色,不满足则抛出 HTTPException(403)
4.路由注入: # 核心:按需鉴权 若是采用中间件实现,还得额外编写白名单 放行(eg:登录、注册接口)
在需要保护的路由参数中,使用 user: dict = Depends(get_current_user)
5.数据隔离:
在数据库查询时,强制使用注入的 user.id 作为过滤条件,防止越权访问(水平越权)
- 实现:
# 安装依赖:
pip install pyjwt python-jose[cryptography] passlib[bcrypt] python-multipart
config/config.py (集中管理配置)
from pydantic_settings import BaseSettings, SettingsConfigDict # 专门处理 环境变量和配置文件
class Settings(BaseSettings):
PROJECT_NAME: str = "FastAPI Production Auth" # 定义项目名称
SECRET_KEY: str = "change-me-in-production" # JWT加密密钥 生产环境通过 .env 注入
ALGORITHM: str = "HS256" # JWT加密算法 HS256 是一种对称加密算法,速度快,适合单体应用
ACCESS_TOKEN_EXPIRE_MINUTES: int = 30 # 过期时间 分钟
# 内部配置类,用于指定 pydantic_settings 的行为
model_config = SettingsConfigDict(
env_file=".env", # 指定读取 项目根目录下的 .env 文件,并将同名变量映射到类属性上
env_file_encoding="utf-8",# 防止中文注释乱码
extra="ignore" # 关键:忽略环境变量中 未定义的多余字段,防止pydantic解析报错
)
settings = Settings()
# 其他字段值
extra="forbid" # 禁止多余字段 遇到未定义的字段,直接抛出ValidationError
extra="allow" # 允许并保留多余字段) 遇到未定义的字段,会将这些字段作为动态属性 保留在模型中
utils/security.py (底层安全工具)
from datetime import datetime, timedelta, timezone
from typing import Optional
import jwt
from passlib.context import CryptContext
from config import settings
# 初始化密码哈希上下文 指定使用 bcrypt 算法
# bcrypt 是一种自带盐值(salt)的慢哈希算法,能有效防止彩虹表攻击和暴力破解
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
# 校验密码 参数:明文密码、hash加密后的密码[一般是数据库存储的]
def verify_password(plain_password: str, hashed_password: str) -> bool:
return pwd_context.verify(plain_password, hashed_password)
# hash加密
def get_password_hash(password: str) -> str:
return pwd_context.hash(password) # 自动生成盐值并进行哈希运算
# 生成JWT Token
# 按照传入的数据(eg:用户名或id) + 过期时间(参数指定 或 配置文件设置)
# ➡️ 根据配置文件中的加密密钥和算法,生成Token(加密字符串)
def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -> str:
# 对传入的数据进行浅拷贝,防止修改原始字典
to_encode = data.copy()
expire = datetime.now(timezone.utc) + (expires_delta or timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES))
to_encode.update({"exp": expire})
# 使用 PyJWT 对数据进行签名加密
# 它会用 SECRET_KEY 和 ALGORITHM 生成一个包含 Header、Payload、Signature 的 Base64 字符串
return jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
### 注意:data 即 JWT的playload部分
{
# === 1. JWT 官方标准字段 (Registered Claims) ===
"sub": "1234567890", # Subject: 主题,通常存放用户的唯一标识(如 user_id)
"jti": "a1b2c3d4-e5f6", # JWT ID: 唯一标识符,通常用于生成 Redis 黑名单以实现主动登出
"iat": 1698765432, # Issued At: 签发时间(Unix 时间戳)
"exp": 1698769032, # Expiration Time: 过期时间(Unix 时间戳)
# === 2. 自定义业务字段 (Custom Claims) ===
"username": "alice", # 用户名(可选,方便后续日志记录)
"roles": ["admin", "user"], # 用户角色(强烈建议放入,方便权限校验)
"is_active": true # 账号状态(可选,用于快速判断账号是否被封禁)
}
schemas/user.py (响应模型类)
from typing import List
from pydantic import BaseModel
class User(BaseModel):
username: str
# 用户的角色列表,默认包含 "user" 角色。支持 RBAC(基于角色的访问控制)
roles: List[str] = ["user"]
# 定义 Token 响应模型,用于登录接口返回给前端的数据结构
class Token(BaseModel):
# JWT Token 字符串
access_token: str
# Token 类型,OAuth2 规范要求固定为 "bearer"
token_type: str = "bearer"
common/auth.py (核心鉴权逻辑)
### 前提:
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/token")
# 作用:
它本身只负责在其他受保护的接口中,自动从请求头里提取 Bearer Token,且提取token值
# tokenUrl 的真实作用
不会创建任何实际的路由,它只是一个配置参数,指向登录接口即可
主要作用是:【仅仅是给 Swagger UI 文档看的】
1.给前端/客户端指路:告诉前端,如果需要获取Token, 向 ip:端口/token 地址,发送用户名和密码
2.给Swagger UI文档指路:当你在 /docs 页面,点击右上角的 "Authorize" 按钮时,弹出的表单,知道该把数据提交给哪个接口
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer # FastAPI内置的 OAuth2 密码模式安全工具
from jose import jwt, JWTError
from sqlalchemy.ext.asyncio import AsyncSession
from config.config import settings
from config.db_config import get_db
from crud import users
from models.user import User
# 初始化 OAuth2 安全方案
# tokenUrl="/token" 告诉 Swagger UI 去哪个接口获取 Token
# 当这个依赖被调用时,它会自动从 HTTP 请求头[ Authorization: Bearer <token> ] 中提取 Token
# 路由函数中,从获取请求头参数 使用 Head()
# 如果请求头没带 Token,它会自动拦截并返回 401 错误,路由函数根本不会执行
# 同时,它会在 Swagger 文档中,生成一个 Authorize 按钮
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/token")
# 定义核心鉴权 依赖函数:校验Token,并返回当前用户对象
# token: str = Depends(oauth2_scheme)
# 以依赖注入的形式,自动执行oauth2_scheme,获取请求头中 token值,并赋值给 token 参数
async def get_current_user(db: AsyncSession = Depends(get_db), token: str = Depends(oauth2_scheme)) -> User:
# 定义一个标准的 401 异常,当 Token 无效或用户不存在时抛出
credentials_exception = HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED, # 401 表示“未认证”
detail="Invalid token / 无效的token",
headers={"WWW-Authenticate": "Bearer"}, # 告诉客户端需要使用 Bearer Token 认证
)
try:
# 使用 PyJWT 解码 Token
# 它会自动验证签名是否被篡改、是否已过期。如果验证失败,会抛出 PyJWTError
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
# 从解码后的 Payload 字典中提取 "sub" (Subject) 字段, 是JWT规范中用来存放用户唯一标识的字段
username: str = payload.get("username")
# 如果 Token 中没有 sub 字段,说明这不是一个合法的认证 Token,抛出 401
if not username:
raise credentials_exception
# 捕获所有 PyJWT 相关的异常(如签名错误、Token 过期等)
except JWTError:
raise credentials_exception
# 【模拟数据库查询】根据 username 从字典中获取用户数据
user = await users.get_user_by_username(db, username)
# 如果数据库中找不到该用户(例如用户被删除,但 Token 还没过期),抛出 401
if not user:
raise credentials_exception
# 验证通过,返回 数据库查到的user对象
# 这个返回值会被 FastAPI 注入到路由函数的参数中
return user
# 【按需设计】 定义RBAC权限校验 依赖函数:嵌套调用 get_current_user,并校验角色
async def require_admin(current_user: User = Depends(get_current_user)) -> User:
# 先执行 get_current_user,如果成功,将返回的User对象赋值给 current_user
# 检查当前用户的角色列表中是否包含 "admin"
if "admin" not in current_user.roles:
# 如果不包含,抛出 403 Forbidden 403 表示“已认证但无权限”
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="You do not have permission to perform this action"
)
# 权限校验通过,返回当前用户对象
return current_user
router/users.py (业务路由)
from fastapi import APIRouter, Depends
from common.auth import get_current_user, require_admin
from models.user import User
router = APIRouter(prefix="/users", tags=["Users"])
# 定义一个需要登录才能访问的接口
@router.get("/me")
async read_users_me(
# 依赖注入的形式:【核心】
# FastAPI 在执行这个路由函数之前,会自动运行 get_current_user 依赖
# 如果鉴权失败,路由函数根本不会被执行,直接返回 401
# 如果成功,解析出的 User 对象会直接赋值给 current_user 变量
current_user: User = Depends(get_current_user)
):
return {"message": f"Hello, {current_user.username}", "roles": current_user.roles}
# 定义一个仅管理员可访问的接口
@router.get("/admin/dashboard")
async admin_dashboard(
# 同理,FastAPI 会先执行 require_admin 依赖
# require_admin 内部又会调用 get_current_user,形成依赖链
admin_user: User = Depends(require_admin)
):
return {"message": f"Welcome Admin: {admin_user.username}"}
9、AOP 与中间件
对应 Spring AOP / 拦截器,FastAPI 中间件做全局日志、跨域、请求头、异常处理。
Java(拦截器)
public class MyInterceptor implements HandlerInterceptor {
// 前置/后置处理
}
Python
@app.middleware("http")
async def my_middleware(request, call_next):
# 前置处理
response = await call_next(request)
# 后置处理
return response
CORS 中间件
解决跨域问题
from fastapi.middleware.cors import CORSMiddleware
# 允许的来源(可以是域名列表)
origins = [
"http://localhost",
"http://localhost:3000",
"https://your-frontend-domain.com"
]
# 添加 CORS 中间件
app.add_middleware(
CORSMiddleware,
allow oriains=["*"], # 允许访问的源 开发环境写* 实际生产 自定义配置 origins
allow_credentials=True, # 允许携带 Cookie
allow_methods=["*"], # 允许所有请求方法
allow_headers=["*"], # 允许所有请求头
)
中间件 vs 依赖注入
# 总结: FastAPI官方推荐 且经过大量生产环境验证的最佳架构模式
# 依赖注入处理的是 路由级别的、与业务紧密相关的逻辑
依赖注入(Depends): # 路由级别(按需注入)、完美集成OpenAPI文档(如自动生成 Authorize 按钮)
身份认证与权限控制(RBAC) # JWT 校验、OAuth2 授权
提取与校验请求参数 # 将分页参数(skip, limit)、复杂的查询条件、或者特定的请求头封装成依赖,避免在每个接口中重复编写参数校验代码
数据库连接与事务管理 # 利用依赖注入的 yield 机制,优雅地管理数据库 Session 的创建与自动释放,或者管理 Redis 客户端连接
获取当前上下文数据 # 将解析出的当前登录用户(User 对象)或租户信息直接注入到路由函数的参数中,供业务代码直接使用
代码复用与单元测试 # 将通用的业务逻辑(如:检查某个商品是否属于当前用户)抽象为依赖,不仅可以在多个接口中复用,在测试时还可以轻松 Mock 替换
# 中间件的设计初衷是 处理全局的、与业务和权限无关的 横切关注点
中间件(Middleware): # 全局生效(所有请求)
全局请求/响应日志记录
跨域配置
全局异常捕获
请求耗时统计 # 在请求进入前记录时间,响应返回前计算耗时,方便定位慢接口
链路追踪 # 为每个请求生成唯一的 Trace ID,方便在分布式系统中串联日志
10、异步支持
对应 Java CompletableFuture / WebFlux,FastAPI 原生 async/await,IO 密集型接口性能极强。
Java
@GetMapping("/async")
public CompletableFuture<String> async() {
return CompletableFuture.supplyAsync(() -> "异步结果");
}
Python
@app.get("/async")
async def async_api():
return {"msg": "异步接口"}
11、ORM 数据库操作
对应 Java MyBatis / JPA,FastAPI 搭配 SQLAlchemy / Tortoise,实现对象映射、增删改查。
Java(JPA)
@Entity
public class User {
@Id
private Long id;
private String name;
}
Python(SQLAlchemy)
from sqlalchemy import Column, Integer, String
from sqlalchemy.ext.declarative import declarative_base
Base = declarative_base()
class User(Base):
__tablename__ = "user"
id = Column(Integer, primary_key=True)
name = Column(String)
12、项目工程结构
按软件功能目录结构
最经典的 MVC 变体。代码按照“它是做什么的”来分类
project/
├── main.py # 主入口文件 初始化FastAP应用,并集成路由和中间件
├── config/
│ ├── __init__.py
│ └── database.py # 数据库连接配置 (数据库引擎、Session工厂)
├── models/ # 【数据层】存放所有数据模型 ,用于构建数据库表
│ ├── __init__.py
│ ├── user.py # User 表定义
│ └── item.py # Item 表定义
├── schemas/ # 【接口层】存放所有 Pydantic 验证模型,用于输入输出验证
│ ├── __init__.py
│ ├── user.py # UserCreate, UserResponse
│ └── item.py # ItemCreate, ItemResponse
├── routers/ # 【路由层】存放 API路径定义 ≈ controller + service
│ ├── __init__.py
│ ├── user.py # @router.get("/users")...
│ └── item.py # @router.get("/items")...
├── crud/ # 数据库操作层 (CRUD Operations) ≈ mapper
│ ├── __init__.py
│ ├── base.py # 基础 CRUD 类 (可选,封装通用增删改查)
│ ├── crud_user.py # 专门处理 User 表的 SQL 操作
│ └── crud_item.py # 专门处理 Item 表的 SQL 操作
├── services/ # 【业务层】(可选) 处理复杂逻辑
│ ├── __init__.py
│ ├── user_service.py # 比如:注册时发送验证码逻辑
│ └── item_service.py
├── common/ # 【公共组件】跨模块复用的代码
│ ├── deps.py # 依赖注入 (如 get_current_user)
│ └── response.py # 统一返回格式封装
└── middleware/ # 【中间件】全局拦截器
└── auth_middleware.py
### 注意:CRUD vs Service
CRUD 层:只做原子操作(增、删、改、查单表)
Service 层 (业务逻辑层):处理复杂流程
但对于大多数中小型 FastAPI 项目,Router + CRUD + Model 的三层结构是性价比最高的选择
按业务模块 目录结构
project/
├── main.py # 入口,负责聚合各模块路由
├── core/ # 【核心配置】全局通用的东西
│ ├── __init__.py
│ ├── config.py # 环境变量读取 (.env)
│ ├── database.py # 数据库引擎与 Session
│ └── security.py # JWT, Hash密码等通用安全工具
├── modules/ # 【业务模块总目录】
│ ├── user/ # --- 用户模块 ---
│ │ ├── __init__.py
│ │ ├── models.py # 仅包含 User 相关的 DB 模型
│ │ ├── schemas.py # 仅包含 User 相关的 Pydantic 模型
│ │ ├── service.py # (可选) User 的业务逻辑
│ │ └── router.py # User 的 API 接口
│ ├── item/ # --- 商品模块 ---
│ │ ├── __init__.py
│ │ ├── models.py
│ │ ├── schemas.py
│ │ └── router.py
│ └── order/ # --- 订单模块 ---
│ ├── ...
├── common/ # 【公共组件】跨模块复用的代码
│ ├── deps.py # 依赖注入 (如 get_current_user)
│ └── response.py # 统一返回格式封装
└── middleware/
└── process_time.py
核心建议:如何选择?
| 维度 | 结构一 (按职能) | 结构二 (按模块) |
|---|---|---|
| 项目规模 | 小 (< 10个接口) | 中/大 (> 20个接口) |
| 团队规模 | 单人 / 双人 | 多人协作 |
| 维护周期 | 短期 / 一次性 | 长期迭代 |
| 重构难度 | 后期重构极其痛苦 | 随时可独立拆分微服务 |
13、统一返回格式
企业级生产落地方案:采用了 泛型响应基类 + 快捷函数 + 异常拦截 的黄金组合
即:正确结果 正常统一返回,错误结果 通过 自定义异常+ 全局异常捕获处理
return 负责优雅地交付数据,raise 负责果断地中断流程,而全局异常处理器负责体面地收拾残局
- common/response.py
from typing import Generic, TypeVar, Any
from pydantic import BaseModel
T = TypeVar('T')
class BaseResponse(BaseModel, Generic[T]):
"""
统一响应基类
"""
code: int = 200
msg: str = "success"
# 使用 T | None ,确保 FastAPI 能正确解析泛型类型,并生成完美的Swagger文档
data: T | None = None
class PaginationData(BaseModel, Generic[T]):
"""通用分页响应数据模型"""
# 使用 Python 标准的 snake_case 命名
list: List[T] = Field(..., description="数据列表")
total: int = Field(..., description="总数据条数")
has_more: bool = Field(..., description="是否还有下一页")
# 核心配置:开启别名生成器,自动将 snake_case 蛇形命名 转为 camelCase 小驼峰命名
model_config = ConfigDict(
alias_generator=to_camel, # 输出时自动转为驼峰 (hasMore)
populate_by_name=True, # 允许在 Python 中同时使用 has_more 或 hasMore
from_attributes=True # 允许从 ORM对象属性读取 转换
)
def success_response(code: int = 200, msg: str = "success", data: Any = None):
return BaseResponse(code=code, msg=msg, data=data)
def fail_response(code: int = 400, msg: str = "fail", data: Any = None):
return BaseResponse(code=code, msg=msg, data=data)
- common/exceptions.py
class BusinessException(Exception):
"""自定义业务异常"""
def __init__(self, code: int = 400, msg: str = "业务处理失败"):
self.code = code
self.msg = msg
- main.py
from fastapi import FastAPI, APIRouter, Request
from fastapi.middleware.cors import CORSMiddleware
from common.exceptions import BusinessException
router = APIRouter(prefix="/api")
app = FastAPI()
# 注册路由
app.include_router(router)
# 注册CORS中间件
origins = [
"http://localhost",
"http://localhost:3000",
"https://your-frontend-domain.com"
]
app.add_middleware(
CORSMiddleware,
allow_origins=origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@router.get("/list", response_model=BaseResponse[PaginationData[NewsResponse]], summary="获取新闻列表")
async def get_news_list(
db: AsyncSession = Depends(get_db),
category_id: int = Query(..., alias="categoryId", description="新闻分类ID"),
page: int = Query(1, ge=1, description="页码"),
page_size: int = Query(10, alias="pageSize", description="每页数据条数")
):
offset = (page - 1) * page_size
# 先获取总量,校验是否有该分类
total = await news.get_news_count(db, category_id)
if not total:
# return fail(msg="该分类不存在") # 使用统一的失败响应
raise BusinessException(code=404, msg="该分类不存在") # 使用异常捕获的方式 推荐
news_list = await news.get_news_list(db, category_id, offset, page_size)
# 是否还有剩余
has_more = total > offset + len(news_list) # 跳过的数据条数 + 当前页数据条数
return success_response(msg="获取新闻列表成功", data=PaginationData(list=news_list, total=total, has_more=has_more))
# 拦截自定义业务异常 !!!
# 注意:也可以将装饰器函数,用普通函数调用方式 注册异常处理
# eg: app.add_exception_handler(BusinessException, business_exception_handler)
@app.exception_handler(BusinessException)
async def business_exception_handler(request: Request, exc: BusinessException):
return JSONResponse(
status_code=200, # 注意:HTTP状态码仍为 200,异常信息由业务异常的code区分
content={"code": exc.code, "msg": exc.msg, "data": None}
)
14、缓存Redis
lifespan机制管理生命周期
# 通过FastAPI的lifespan机制,来管理它的生命周期
from uvicorn import lifespan
# 即:在FastAPI应用 启动/结束时,做特定操作 eg: 创建数据库表、创建或销毁Redis连接池
@asynccontextmanager
async def lifespan(app: FastAPI):
# 应用启动时 操作
print("正在连接 Redis...")
redis_pool = redis.ConnectionPool(
host="localhost",
port=6379,
db=0,
decode_responses=True, # 是否将返回的数据 从字节码流 解码成 字符串
max_connections=50 # 连接池大小
)
yield # 抛出去,应用开始真正执行【请求➡️响应】
# 应用结束时 操作
print("正在关闭 Redis 连接...")
await redis_pool.disconnect()
# app 挂载lifespan
app = FastAPI(lifespan=lifespan)
- 安装
# 安装
pip install redis # 默认包含 redis[asyncio] 异步操作
- 配置Redis客户端
# redis_client.py
### 普通配置-Redis客户端
import redis.asyncio as redis
redis_client = redis.Redis(
host,
port,
db, # 数据库编号(0-15)
decode_responses=True # 是否将返回的数据 从字节码流 解码成 字符串
)
### 通过连接池获取-Redis客户端
import redis.asyncio as redis
from contextlib import asynccontextmanager
from typing import AsyncGenerator, Annotated
from fastapi import FastAPI
# 定义全局连接池变量
redis_pool: redis.ConnectionPool | None = None
@asynccontextmanager
async def lifespan(app: FastAPI):
"""应用启动时创建连接池,关闭时销毁"""
global redis_pool
print("正在连接 Redis...")
redis_pool = redis.ConnectionPool(
host="localhost",
port=6379,
db=0,
password="123456",
decode_responses=True,
max_connections=50 # 连接池大小
)
yield
print("正在关闭 Redis 连接...")
await redis_pool.disconnect()
async def get_redis() -> AsyncGenerator[redis.Redis, None]:
"""依赖注入函数:为每个请求提供 Redis 客户端"""
if redis_pool is None:
raise RuntimeError("Redis 连接池未初始化")
# 从连接池获取客户端
# 利用上下文管理器,自动处理连接的借出和归还,无需手动 client.aclose()
async with redis.Redis(connection_pool=redis_pool) as client:
yield client
# 【可选】 类型别名: 将依赖和类型绑定在一起 告别冗长的 Depends
RedisClient = Annotated[redis.Redis, Depends(get_redis)]
==============采用app.state 【单进程内的全局共享变量】 存储 redis_pool
# 1. 在 lifespan 中挂载到 app.state
@asynccontextmanager
async def lifespan(app: FastAPI):
app.state.redis_pool = redis.ConnectionPool(...)
yield
await app.state.redis_pool.disconnect()
# 2. 在依赖中通过request获取
async def get_redis(request: Request) -> AsyncGenerator[redis.Redis, None]:
redis_pool = request.app.state.redis_pool # 此时才真正需要 request
async with redis.Redis(connection_pool=redis_pool) as client:
yield client
- 依赖注入使用
from fastapi import APIRouter, Depends
import redis.asyncio as aioredis
from redis_client import get_redis, RedisClient
router = APIRouter(prefix="/items")
@router.get("/{item_id}")
async def get_item(item_id: int, r: RedisClient):
cache_key = f"item:{item_id}"
# 1. 尝试从缓存获取
cached_data = await r.get(cache_key)
if cached_data:
return {"source": "cache", "data": cached_data}
# 2. 缓存未命中,模拟耗时数据库查询
# db_data = await db.query(...)
db_data = {"id": item_id, "name": "示例商品"}
# 3. 写入缓存,设置 60 秒过期时间
await r.setex(cache_key, 60, str(db_data))
return {"source": "database", "data": db_data}
常用操作封装
import asyncio
import json
import logging
import random
import uuid
from contextlib import asynccontextmanager
from typing import Any, Optional, Callable, Awaitable
from config.cache_config import RedisClient
logger = logging.getLogger(__name__)
# 缓存雪崩,随机偏移区间(秒)
CACHE_EXPIRE_OFFSET = (-30, 30)
# 空值缓存,默认过期时间(防穿透)
NULL_CACHE_EXPIRE = 60
# 分布式锁释放Lua脚本:原子校验value+删除,避免误删其他线程锁
# Lua 脚本:保证判断和删除的原子性,防止误删其他协程/进程的锁
RELEASE_LOCK_SCRIPT = """
if redis.call("get", KEYS[1]) == ARGV[1] then
return redis.call("del", KEYS[1])
else
return 0
end
"""
# 计数器原子操作Lua脚本:自增并设置过期时间(防止内存泄漏)
INCR_WITH_EXPIRE_SCRIPT = """
local current = redis.call('incrby', KEYS[1], ARGV[1])
if current == tonumber(ARGV[1]) then
redis.call('expire', KEYS[1], ARGV[2])
end
return current
"""
class LockAcquireFailedError(Exception):
"""获取分布式锁失败异常"""
pass
# ================= 1. 缓存读写删 =================
async def get_cache(redis_client: RedisClient, key: str) -> str | None:
"""获取普通字符串缓存"""
try:
return await redis_client.get(key)
except Exception as e:
logger.error(f"[Redis] 获取缓存 失败 key={key}, error={e}", exc_info=True)
return None
async def get_json_cache(redis_client: RedisClient, key: str) -> Optional[Any]:
"""读取JSON格式缓存,自动反序列化"""
try:
data = await redis_client.get(key)
return json.loads(data) if data else None
except Exception as e:
logger.error(f"[Redis] 获取JSON缓存 失败 key={key}, error={e}", exc_info=True)
return None
async def set_cache(redis_client: RedisClient, key: str, value: Any, expire: int = 3600,
random_offset: bool = True) -> bool:
"""写入缓存:自动序列化dict/list为JSON,普通值直接存储"""
try:
# 自动序列化字典/列表
if isinstance(value, (dict, list)):
value = json.dumps(value, ensure_ascii=False) # 禁止进行ascii转码,即中文正常保存
# 增加随机过期偏移,打散大量key同时失效【防止缓存雪崩】
real_expire = expire
if random_offset:
offset = random.randint(*CACHE_EXPIRE_OFFSET)
real_expire = max(1, expire + offset)
return await redis_client.setex(key, real_expire, value)
except Exception as e:
logger.error(f"[Redis] 设置缓存 失败 key={key}, error={e}", exc_info=True)
return False
async def delete_cache(redis_client: RedisClient, key: str) -> bool:
"""删除缓存"""
try:
return await redis_client.delete(key) > 0
except Exception as e:
logger.error(f"[Redis] 删除缓存 失败 key={key}, error={e}", exc_info=True)
return False
# ================= 2. 分布式锁 (安全版) =================
@asynccontextmanager
async def distributed_lock(redis_client: RedisClient, key: str, expire: int = 10, retry_times: int = 0,
retry_delay: float = 0.1):
"""
异步分布式锁上下文管理器
:param key: 锁key
:param expire: 锁自动过期时间(秒)
:param retry_times: 获取锁重试次数,0=不重试直接抛异常
:param retry_delay: 每次重试间隔(秒)
用法:
async with distributed_lock(redis, "lock:user:1", retry_times=3):
# 临界区业务
pass
"""
lock_value = str(uuid.uuid4()) # 生成唯一标识,防止误删
acquired = False
try:
# 重试获取锁
for _ in range(retry_times + 1):
acquired = await redis_client.set(key, lock_value, nx=True, ex=expire)
if acquired:
break
# 未获取到锁,等待后重试
await asyncio.sleep(retry_delay)
if not acquired:
raise LockAcquireFailedError(f"获取分布式锁失败,已重试{retry_times}次 key={key}")
yield # 获取锁后,让出上下文 继续执行正常操作
finally:
if acquired:
try:
# 使用 Lua 脚本原子性释放锁
await redis_client.eval(RELEASE_LOCK_SCRIPT, 1, key, lock_value)
except Exception as e:
logger.error(f"[Redis] 释放分布式锁 失败 key={key}, error={e}", exc_info=True)
# ================= 3. 核心业务抽象:Cache Aside 旁路缓存模式 =================
async def get_or_set_cache(
redis_client: RedisClient,
key: str,
expire: int,
db_func: Callable[..., Awaitable[dict | list | None]],
*args, # 截断点,之后都必须使用关键字参数
lock_expire: int = 5,
lock_retry: int = 2,
**kwargs # 捕获所有业务关键字参数(透传给 db_func)
) -> Optional[Any]:
"""
标准通用的Cache Aside旁路缓存模式
防护策略:
1. 缓存穿透:空值短时间缓存
2. 缓存击穿:分布式锁+双重缓存校验
3. 缓存雪崩:set_cache默认开启随机过期偏移
:param lock_expire: 分布式锁过期时间 必须以关键字形式传入
:param lock_retry: 获取锁重试次数 必须以关键字形式传入
:param db_func: 异步查库 且 ORM转字典 必须返回可被JSON序列化的基础数据类型(dict, list, None)的 函数
# Awaitable[Any] 是一个可等待的异步对象,当它执行完毕(被await)后,返回结果是任意类型
"""
# 1.第一次查缓存
cached_data = await get_json_cache(redis_client, key)
if cached_data is not None: # 防止缓存了 0 或 "" 等假值
return cached_data
# 2.缓存未命中,加分布式锁 防缓存击穿
lock_key = f"cache:lock:{key}"
# 加锁处理
async with distributed_lock(redis_client, lock_key, expire=lock_expire, retry_times=lock_retry):
# 3.【双重检查】拿到锁后,再次查缓存(防止其他协程已经查库,并回写了)
cached_data = await get_json_cache(redis_client, key)
if cached_data is not None:
return cached_data
# 4.查库
try:
db_data = await db_func(**kwargs)
except Exception as e:
logger.error(f"[Redis] 数据源查询 失败 key={key}, args={args}, kwargs={kwargs}, error={e}", exc_info=True)
raise # 建议向上抛出异常,让业务层决定如何处理
# 5.回写缓存(包含空值,防穿透)
if db_data is not None:
await set_cache(redis_client, key, db_data, expire)
else:
# 防止缓存穿透:有一个不存在的key被高频请求,每次都会直接打穿到数据库
# 空值缓存,设置较短过期时间(如60秒)
await set_cache(redis_client, key, None, expire=NULL_CACHE_EXPIRE)
return db_data
# ===================== 4. 辅助工具:计数器 =====================
async def cache_incr(redis_client: RedisClient, key: str, step: int = 1, expire: int = 3600) -> int:
"""
缓存自增计数器,key不存在自动初始化为0再自增
使用Lua脚本保证 自增+设置过期时间 的原子性,防止并发下内存泄漏
"""
try:
result = await redis_client.eval(INCR_WITH_EXPIRE_SCRIPT, 1, key, step, expire)
return int(result)
except Exception as e:
logger.error(f"[Redis] 计数器自增失败 key={key}, error={e}", exc_info=True)
return 0
ORM对象 转 字典 的通用封装
# 采用Pydantic.model_dump(mode="json") 底层是采用 Rust 性能很高
# 保证了 进行数据库查询,且返回的 Redis 可直接操作的 字典 或者列表类型
# 该函数可以直接 作为 缓存的get_or_set_cache(db_func=xxx) 传入
### 方案一:封装为工厂函数 / 装饰器(推荐,最直观)
写一个高阶函数,它接收 Pydantic 模型类,返回一个“自动转换”的查询函数
from typing import TypeVar, Type, Callable, Awaitable, Optional
from pydantic import BaseModel
T = TypeVar("T", bound=BaseModel)
def auto_serialize(model_class: Type[T]) -> Callable:
"""
高阶函数:自动将 ORM 对象转换为 Pydantic 字典
"""
async def decorator(func: Callable[..., Awaitable]) -> Callable[..., Awaitable[Optional[dict]]]:
async def wrapper(*args, **kwargs):
# 1. 执行原始的查库函数
result = await func(*args, **kwargs)
# 2. 如果查不到数据,直接返回 None
if result is None:
return None
# 3. 如果返回的是列表(批量查询),则批量转换
if isinstance(result, list):
return [model_class.model_validate(item).model_dump(mode="json") for item in result]
# 4. 单个对象转换
return model_class.model_validate(result).model_dump(mode="json")
return wrapper
return decorator
# 使用方式:
# 定义一个带自动序列化功能的查库函数
@auto_serialize(UserResponse)
async def fetch_user_from_db(user_id: int):
return await session.get(User, user_id)
# 直接传给缓存层
user_data = await get_or_set_cache(
redis_client=redis,
key=f"user:{user_id}",
expire=3600,
db_func=fetch_user_from_db,
user_id=user_id
)
### 方案二:封装为 Pydantic 模型基类(最优雅)
如果项目中所有的 Schema 都继承自同一个基类,可以直接把转换能力“注入”到基类中
from typing import TypeVar, Callable, Awaitable, Optional, Any
from pydantic import BaseModel, ConfigDict
# 泛型变量,确保类型提示准确
T = TypeVar("T", bound=BaseModel)
class BaseSchema(BaseModel):
model_config = ConfigDict(alias_generator=to_camel, populate_by_name=True, from_attributes=True)
@classmethod
async def from_orm_async(
cls: type[T],
func: Callable[..., Awaitable[Any]],
*args,
**kwargs
)-> Optional[dict | list[dict]]:
"""
异步执行ORM查询,并自动将结果序列化为JSON兼容的字典 或包含字典的列表
"""
result = await func(*args, **kwargs)
if result is None:
return None
if isinstance(result, list):
return [cls.model_validate(item).model_dump(mode="json") for item in result]
return cls.model_validate(result).model_dump(mode="json")
# 使用方式:
# Schema 继承基类
class UserResponse(BaseSchema):
id: int
username: str
# 在 Service/router 层极其优雅地调用
from schemas.user import UserResponse
from crud.user import get_user_by_id
from utils.cache import get_or_set_cache # 写的旁路缓存工具
async def get_user_with_cache(redis_client, db: AsyncSession, user_id: int):
cache_key = f"user:info:{user_id}"
# 使用偏函数:固定函数的部分参数,生成一个新的函数对象
# 把ORM转换逻辑和查库逻辑绑定在一起
bound_db_func = partial(
UserResponse.from_orm_async,
func=get_user_by_id
)
return await get_or_set_cache(
redis_client=redis_client,
key=cache_key,
expire=3600,
db_func=bound_db_func, # 传入绑定好的函数对象
# 把CRUD方法需要的业务参数,以关键字参数 **kwargs 透传给get_user_by_id
db=db,
user_id=user_id
)
# 总结与建议
方案一(装饰器/高阶函数):
适合不想修改现有 Schema 基类的情况,逻辑独立,即插即用
方案二(基类方法):
如果使用的是 FastAPI,强烈推荐使用方案二
它将“查库”和“序列化”完美地绑定在了Schema层,符合面向对象设计的内聚原则

浙公网安备 33010602011771号