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层,符合面向对象设计的内聚原则
posted @ 2026-07-03 00:42  Edmond辉仔  阅读(10)  评论(0)    收藏  举报