FastAPI进阶
中间件
作用
使用中间件为每个请求前后添加统一的处理逻辑:
- 多个接口都需要验证用户身份
- 多个接口 都需要记录日志、性能数据

含义
中间件(Middleware)是一个在每次请求进入 FastAPI 应用时都会被执行的函数。 它在请求到达实际的路径操作(路由处理函数)之前运行,并且在响应返回给客户端之前再运行一次。

写法
函数的顶部使用装饰器 @app.middleware("http")
异步函数async def middleware(request,call_next)
有两个参数:
- request:请求
- call_naxt:传递请求给路径处理函数
一个中间件情况:
from fastapi import FastAPI
# 创建实例
app = FastAPI()
@app.get("/")
async def root():
return {"message": "你好哈哈哈哈"}
@app.middleware("http")
async def middleware1(request,call_next):
print("中间件1开始")
response = await call_next(request)
print("中间件1结束")
return response

两个中间件(自下而上执行顺序):
from fastapi import FastAPI
# 创建实例
app = FastAPI()
@app.get("/")
async def root():
return {"message": "你好哈哈哈哈"}
@app.middleware("http")
async def middleware1(request,call_next):
print("中间件1开始")
response = await call_next(request)
print("中间件1结束")
return response
@app.middleware("http")
async def middleware2(request,call_next):
print("中间件2开始")
response = await call_next(request)
print("中间件2结束")
return response

依赖注入
作用
使用依赖注入系统来共享通用逻辑,减少代码重复
与中间件区别
中间件控制所有接口,而依赖注入控制谁自己说了算
含义
依赖注入就是我们封装依赖项注入到路由的处理函数。
- 依赖项:可重用的组件(函数/类),负责提供某种功能或数据。
- 注入:FastAPI 自动帮你调用依赖项,并将结果"注入"到路径操作函数中。
- 优点:
- 代码复用:一次编写,多处使用
- 解耦:业务逻辑与基础设施代码分离
- 易于测试:轻松地用模拟依赖替换真实依赖进行测试
应用场景

调用依赖注入
调用依赖注入需要三步:
创建依赖项
# 1.创建依赖项
async def common(skip:int=Query(0,ge=0),limit:int=Query(10,le=100)):
return {"skip":skip,"limit":limit}
导入Dpends
from fastapi import Depends #导入Depends
声明依赖项
# 3.依赖注入
@app.get("/item/item_list")
async def items(commons=Depends(common)):
return commons

完整测试:
from fastapi import FastAPI,Query,Depends #导入Depends
# 创建实例
app = FastAPI()
@app.get("/")
async def root():
return {"message": "你好哈哈哈哈"}
# 分页参数逻辑共用:项目列表和用户列表共用
# 1.创建依赖项
async def common(skip:int=Query(0,ge=0),limit:int=Query(10,le=100)):
return {"skip":skip,"limit":limit}
# 3.依赖注入
@app.get("/item/item_list")
async def items(commons=Depends(common)):
return commons
@app.get("/user/user_list")
async def user(commons=Depends(common)):
return commons
ORM
ORM(Object-RelationalMapping,对象关系映射)是一种编程技术,用于在面向对象编程语言和关系型数据库之间 建立映射。它允许开发者通过操作对象的方式与数据库进行交互,而无需直接编写复杂的SQL语句。
- 优势:
- 减少重复的 SQL 代码
- 代码更简洁易读
- 自动处理数据库连接和事务
- 自动防止 SQL 注入攻击
分类
| 排名 | ORM 工具 | 特点 | 适应场景 |
| 1 | SQLAlchemy ORM | 功能最强、最灵活、企业级 | 各类 API、微服务、数据应用 |
| 2 | Django ORM | 封装好、上手快 | Django 项目、管理后台 |
| 3 | Tortoise ORM | 全异步 | 异步 Web 服务、高并发 API |
使用流程(SQLAlchemy ORM)
安装
- sqlalchemy[asyncio]
- aiomysql(异步数据库驱动)
使用清华源:
pip install "sqlalchemy[asyncio]" aiomysql -i https://pypi.tuna.tsinghua.edu.cn/simple
建表
- 流程:
- 1.创建数据库引擎
- 2.定义模型类
- 3.启动应用时建表
创建数据库引擎
使用 create_async_engine 创建异步引擎
from sqlalchemy.ext.asyncio import create_async_engine # root:123456是数据库密码;fastapi_project:项目名称; ASYNC_DATABASE_URL = "mysql+aiomysql://root:123456@localhost:3306/fastapi_project?charset=utf8" # 创建异步引擎 async_engine = create_async_engine( ASYNC_DATABASE_URL, echo=True, # 可选:输出SQL日志 pool_size=10, # 设置连接池活跃的连接数 max_overflow=20, # 设置连接池允许创建的额外连接数 )
定义模型类
- 基类,继承 DeclarativeBase(包含通用属性和字段的映射)
- 定义数据库表对应的模型类
若出现下面错误:
RuntimeError: 'cryptography' package is required for sha256_password or caching_sha2_password auth methods就安装
pip install cryptography -i https://pypi.tuna.tsinghua.edu.cn/simple
from datetime import datetime
from fastapi import FastAPI
from sqlalchemy import DateTime, func, String, Float
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
# 创建实例
app = FastAPI()
# root:123456是数据库密码;fastapi_project:项目名称;
ASYNC_DATABASE_URL = "mysql+aiomysql://root:123456@localhost:3306/fastapi_project?charset=utf8mb4"
# 创建异步引擎
async_engine = create_async_engine(
ASYNC_DATABASE_URL,
echo=True, # 可选:输出SQL日志
pool_size=10, # 设置连接池活跃的连接数
max_overflow=20, # 设置连接池允许创建的额外连接数
)
# 定义模型类:基类+表对应的模型类
# 基类:创建时间,更新时间;项目表:id,项目名,负责人,价格,甲方
class Base(DeclarativeBase):
create_time: Mapped[datetime] = mapped_column(DateTime, insert_default=func.now(), default=func.now, comment="创建时间")
update_time: Mapped[datetime] = mapped_column(DateTime, insert_default=func.now(), default=func.now, comment="修改时间")
class Project(Base):
__tablename__ = "project"
id: Mapped[int] = mapped_column(primary_key=True, comment="项目id")
name: Mapped[str] = mapped_column(String(255), comment="项目名")
author: Mapped[str] = mapped_column(String(255), comment="负责人")
price: Mapped[float] = mapped_column(Float, comment="价格")
boss: Mapped[str] = mapped_column(String(255), comment="甲方")
# 建表 定义函数建表 fastapi启动的时候调用建表的函数
async def create_table():
# 获取异步引擎,创建事务,建表
async with async_engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all) # Base 模型类的源数据创建
@app.on_event("startup")
async def startup():
await create_table()
@app.get("/")
async def root():
return {"message": "你好哈哈哈哈"}
在路由中使用ORM
核心:创建依赖项,使用 Depends 注入到处理函数
from datetime import datetime
from fastapi import FastAPI, Depends
from sqlalchemy import DateTime, func, String, Float, select
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncSession
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
# 创建实例
app = FastAPI()
# root:123456是数据库密码;fastapi_project:项目名称;
ASYNC_DATABASE_URL = "mysql+aiomysql://root:123456@localhost:3306/fastapi_project?charset=utf8mb4"
# 创建异步引擎
async_engine = create_async_engine(
ASYNC_DATABASE_URL,
echo=True, # 可选:输出SQL日志
pool_size=10, # 设置连接池活跃的连接数
max_overflow=20, # 设置连接池允许创建的额外连接数
)
# 定义模型类:基类+表对应的模型类
# 基类:创建时间,更新时间;项目表:id,项目名,负责人,价格,甲方
class Base(DeclarativeBase):
create_time: Mapped[datetime] = mapped_column(DateTime, insert_default=func.now(), default=func.now, comment="创建时间")
update_time: Mapped[datetime] = mapped_column(DateTime, insert_default=func.now(), default=func.now, comment="修改时间")
class Project(Base):
__tablename__ = "project"
id: Mapped[int] = mapped_column(primary_key=True, comment="项目id")
name: Mapped[str] = mapped_column(String(255), comment="项目名")
author: Mapped[str] = mapped_column(String(255), comment="负责人")
price: Mapped[float] = mapped_column(Float, comment="价格")
boss: Mapped[str] = mapped_column(String(255), comment="甲方")
# 建表 定义函数建表 fastapi启动的时候调用建表的函数
async def create_table():
# 获取异步引擎,创建事务,建表
async with async_engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all) # Base 模型类的源数据创建
@app.on_event("startup")
async def startup():
await create_table()
@app.get("/")
async def root():
return {"message": "你好哈哈哈哈"}
# 需求:查询功能的接口,查询图书,依赖注入:创建依赖项获取数据库会话+Depends注入路由处理函数
# 创建异步会话工厂
AsyncSessionLocal = async_sessionmaker(
bind=async_engine, # 绑定数据库引擎
class_=AsyncSession, # 指定会话类
expire_on_commit=False # 会话对象不过期,不重新查询数据库
)
# 依赖项,用于获取数据库会话
async def get_database():
async with AsyncSessionLocal() as session:
try:
yield session # 返回数据库会话给路由处理函数
await session.commit() # 无异常提交事务
except Exception:
await session.rollback() # 有异常则回滚(保证数据的一致性)
raise
finally:
await session.close() # 关闭会话
@app.get("/projects/project_list")
async def read_items(db: AsyncSession= Depends(get_database)):
# 查询所有项目
result = await db.execute(select(Project)) # 从数据库中查询所有 Project 表的数据。
project = result.scalars().all() # .scalars() - 将行转换为标量值,.all() - 获取所有结果 返回: List[Project]
return project

问题1:
如果使用上述代码出现问题:
on_event is deprecated, use lifespan event handlers instead.
可以尝试下面的代码:
from datetime import datetime
from contextlib import asynccontextmanager
from fastapi import FastAPI, Depends
from sqlalchemy import DateTime, func, String, Float, select
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncSession
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
from pydantic import BaseModel
# 创建异步引擎
ASYNC_DATABASE_URL = "mysql+aiomysql://root:123456@localhost:3306/fastapi_project?charset=utf8mb4"
async_engine = create_async_engine(
ASYNC_DATABASE_URL,
echo=True,
pool_size=10,
max_overflow=20,
)
# 定义模型类
class Base(DeclarativeBase):
create_time: Mapped[datetime] = mapped_column(DateTime, insert_default=func.now(), default=func.now, comment="创建时间")
update_time: Mapped[datetime] = mapped_column(DateTime, insert_default=func.now(), default=func.now, comment="修改时间")
class Project(Base):
__tablename__ = "project"
id: Mapped[int] = mapped_column(primary_key=True, comment="项目id")
name: Mapped[str] = mapped_column(String(255), comment="项目名")
author: Mapped[str] = mapped_column(String(255), comment="负责人")
price: Mapped[float] = mapped_column(Float, comment="价格")
boss: Mapped[str] = mapped_column(String(255), comment="甲方")
# 建表函数
async def create_table():
async with async_engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
# 创建异步会话工厂
AsyncSessionLocal = async_sessionmaker(
bind=async_engine,
class_=AsyncSession,
expire_on_commit=False
)
# 使用 lifespan 上下文管理器替代 @app.on_event
@asynccontextmanager
async def lifespan(app: FastAPI):
# 启动时执行建表
await create_table()
print("数据库表创建完成")
yield
# 关闭时清理资源
await async_engine.dispose()
print("数据库连接已关闭")
# 创建 FastAPI 实例时传入 lifespan
app = FastAPI(lifespan=lifespan)
@app.get("/")
async def root():
return {"message": "你好哈哈哈哈"}
# 依赖项,用于获取数据库会话
async def get_database():
async with AsyncSessionLocal() as session:
try:
yield session
await session.commit()
except Exception:
await session.rollback()
raise
finally:
await session.close()
# Pydantic 模型
class ProjectBase(BaseModel):
id: int
name: str
author: str
price: float
boss: str
@app.post("/project/add_project")
async def get_list(project: ProjectBase, db: AsyncSession = Depends(get_database)):
project_obj = Project(**project.__dict__)
db.add(project_obj)
await db.commit()
return project
问题2docs文档加载失败:
如果docs文档加载不出来可以使用下面代码:
app = FastAPI(lifespan=lifespan,
docs_url=None, # 禁用 Swagger UI
redoc_url=None # 禁用默认的在线 Redoc
)
# 添加一个路由来提供离线文档
from fastapi.openapi.docs import get_swagger_ui_html, get_redoc_html
from fastapi.staticfiles import StaticFiles
# 挂载静态文件目录(如果需要的话)
# app.mount("/static", StaticFiles(directory="static"), name="static")
@app.get("/docs", include_in_schema=False)
async def custom_swagger_ui_html():
return get_swagger_ui_html(
openapi_url=app.openapi_url,
title=app.title + " - Swagger UI",
oauth2_redirect_url=app.swagger_ui_oauth2_redirect_url,
swagger_js_url="https://unpkg.com/swagger-ui-dist@5/swagger-ui-bundle.js",
swagger_css_url="https://unpkg.com/swagger-ui-dist@5/swagger-ui.css",
)
操作数据
查询(select)
核心语句:await db.execute( select(模型类) ),返回一个 ORM 对象
- 获取所有数据
- scalars().all()
- 获取单条数据
- scalars().first()
- get(模型类, 主键值)
@app.get("/projects/project_list")
async def read_items(db: AsyncSession= Depends(get_database)):
# # 查询所有项目
# result = await db.execute(select(Project)) # 从数据库中查询所有 Project 表的数据。
# # project = result.scalars().all() #获取所有结果
# project = result.scalars().first() # 获取单条数据
# 使用get(模型类, 主键值)获取
project = await db.get(Project, 2)
return project
查询条件
- 比较判断
- ==; >; =; <= 等
# 需求:路径参数 项目id @app.get("/projects/{project_id}") async def read_items(project_id: int, db: AsyncSession= Depends(get_database)): result = await db.execute(select(Project).where(Project.id == project_id)) project = result.scalar_one_or_none() return project # 需求:路径参数 价格>=999999 @app.get("/project_price") async def read_items(db: AsyncSession = Depends(get_database)): result = await db.execute(select(Project).where(Project.price >= 999999.0)) project = result.scalars().all() return project- 模糊查询:like()
- %:零个、一个或多个字符
- _:一个单个字符
- 与非查询:
- &:与
- |:或
- ~:非
- 包含查询:in_()
- 聚合查询:func.方法(模型类.属性)
- count:统计行数量
- avg:求平均值
- max:求最大值
- min:求最小值
- sum:求和
# 需求:查询以张开头的负责人 @app.get("/project/author") async def author(db: AsyncSession = Depends(get_database)): # result = await db.execute(select(Project).where((Project.author.like("张%"))&(Project.name.like("项目_")))) # # 需求:项目id在自己定义的列表里 # id_list = [1, 3, 5, 6] # result = await db.execute(select(Project).where(Project.id.in_(id_list))) # project = result.scalars().all() # return project # 统计数量 result = await db.execute(select(func.count(Project.id))) project = result.scalar() return project
- 分页查询
- select().offset().limit()
- offset:跳过的记录数
- offset值 = (当前页码-1)*每页数量limit
- limit:返回的记录数
@app.get("/project") async def get_list( page: int = 1, page_size: int = 3, db: AsyncSession = Depends(get_database) ): skip = (page - 1) * page_size curr_size = select(Project).offset(skip).limit(page_size) result = await db.execute(curr_size) project = result.scalars().all() return project
新增(add)
核心步骤:定义 ORM 对象 → 添加对象到事务:add(对象) → commit 提交到数据库
# 需求:用户输入项目信息 class ProjectBase(BaseModel): id: int name: str author: str price: float boss: str @app.post("/project/add_project") async def get_list(project: ProjectBase, db: AsyncSession = Depends(get_database)): # orm对象 add,commit project_obj = Project(**project.__dict__) db.add(project_obj) await db.commit() return project
更新
# 需求:修改项目信息,先查再改
# 设计思路:路径参数id:作用是查找;请求体参数:作用是新数据(项目名、负责人、价格、老板)
@app.put("/project/update_project/{project_id}")
async def update_project(project_id: int, data: ProjectBase, db: AsyncSession = Depends(get_database)):
db_project = await db.get(Project, project_id)
if db_project is None:
raise HTTPException(
status_code=404,
detail="查无此书"
)
# 重新赋值
db_project.name = data.name
db_project.author = data.author
db_project.price = data.price
db_project.boss = data.boss
await db.commit()
return db_project
删除
# 删除
@app.delete("/project/delete_project/{project_id}")
async def delete_project(project_id: int, db: AsyncSession = Depends(get_database)):
# 先查再删除在提交
db_project = await db.get(Project, project_id)
if db_project is None:
raise HTTPException(
status_code=404,
detail="没找到"
)
await db.delete(db_project)
await db.commit()
return db_project
更多推荐
所有评论(0)