中间件

作用

使用中间件为每个请求前后添加统一的处理逻辑:

  • 多个接口都需要验证用户身份
  • 多个接口 都需要记录日志、性能数据

含义

中间件(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,    # 设置连接池允许创建的额外连接数
)
定义模型类
  1. 基类,继承 DeclarativeBase(包含通用属性和字段的映射)
  2. 定义数据库表对应的模型类

若出现下面错误:
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

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐