依赖注入
前言
依赖注入(Dependency Injection, DI)是编程模式中依赖倒置原则的一种应用实现:对象不自行创建依赖,而是通过某种机制自动接收依赖实例。依赖注入中涉及 4 个关键概念:
- 服务:可以是类或函数,负责提供具体功能实现。
- 客户:使用服务的对象(如视图函数)。
- 接口:客户与服务之间的连接,客户只需声明需要的服务,无需了解实现细节。
- 注入器:负责把服务实例注入到客户中。
FastAPI 的依赖树机制在某种程度上扮演了 IoC 容器角色,可自动解析并处理每层依赖的注册和执行。常见应用场景:
- 业务逻辑共享:统一定义逻辑,避免每个函数重复编写。
- 数据库连接:共享同一上下文中的连接/会话,避免重复创建。
- 认证和授权:统一管理认证鉴权逻辑。
- 缓存管理、外部 API 调用、参数校验和转换等。
在 FastAPI 中,依赖通过声明函数参数实现,使用 Depends 注入:
from fastapi import Depends
函数式依赖项
from fastapi import Depends, Query
from fastapi.exceptions import HTTPException
# 定义函数式依赖项
def username_check(username: str = Query(...)):
if username != "zhangsan":
raise HTTPException(status_code=403, detail="Permission Denied")
return username
# 通过 Depends 将依赖项注入到视图函数中
@app.get("/user/login")
def login(username: str = Depends(username_check)):
return username
# 通过依赖注入实现复用同一段业务处理逻辑
@app.get("/user/info")
def userinfo(username: str = Depends(username_check)):
return username
类方式依赖项
fake_items_db = [
{"item_name": "Foo"},
{"item_name": "Bar"},
{"item_name": "Baz"},
]
class CommonQueryParams:
def __init__(self, q: str | None = None, page: int = 1, limit: int = 100):
self.q = q
self.page = page
self.limit = limit
# 三种写法等价
# async def classes_as_dependency(commons: CommonQueryParams = Depends(CommonQueryParams)):
# async def classes_as_dependency(commons: CommonQueryParams = Depends()):
@app.get("/classes_as_dependency")
async def classes_as_dependency(commons=Depends(CommonQueryParams)):
response = {}
if commons.q:
response.update({"q": commons.q})
items = fake_items_db[commons.page : commons.page + commons.limit]
response.update({"items": items})
return response
多个依赖项注入
from fastapi import Depends, Query
from fastapi.exceptions import HTTPException
def username_check(username: str = Query(...)):
if username != "zhangsan":
raise HTTPException(status_code=403, detail="Permission denied")
return username
def age_check(age: int = Query(...)):
if age < 18:
raise HTTPException(status_code=403, detail="用户未满18岁")
return age
@router.get("/user/login/", summary="user login")
def user_login(
username: str = Depends(username_check),
age: int = Depends(age_check),
):
return {"username": username, "age": age}
依赖项传参(带参数的依赖)
依赖函数本身需要配置参数时,不能直接给 Depends 传参,需要用工厂函数返回依赖函数,或者用 functools.partial 预先绑定参数:
from functools import partial
from fastapi import Depends, HTTPException
# 方式一:工厂函数返回依赖
def require_permission(role: str):
def checker(username: str = Depends(username_check)) -> str:
if role == "admin":
return username
raise HTTPException(status_code=403, detail="需要管理员权限")
return checker
@router.get("/admin/info", summary="管理员接口")
def admin_info(username: str = Depends(require_permission("admin"))):
return {"username": username, "role": "admin"}
# 方式二:functools.partial 预先绑定参数
def role_checker(username: str, role: str = "user"):
if username == "zhangsan" and role == "admin":
return username
raise HTTPException(status_code=403, detail="无权限")
admin_checker = partial(role_checker, role="admin")
@router.get("/admin/partial")
def admin_partial(username: str = Depends(admin_checker)):
return {"username": username}
提示
实际项目中更常见的用法是依赖数据库配置:get_db_session() 由工厂根据配置创建,见 连接数据库。
多层依赖项嵌套注入
复杂业务场景中,依赖可以继续依赖其它依赖:
def query(q: str | None = None):
return q
def sub_query(q: str = Depends(query), last_query: str | None = None):
if not q:
return last_query
return q
@app.get("/sub_dependency")
async def sub_dependency(final_query: str = Depends(sub_query, use_cache=True)):
"""use_cache 默认为 True:多个依赖共享同一个子依赖时,每次请求只调用子依赖一次"""
return {"sub_dependency": final_query}
全局依赖项
FastAPI(...) 和 APIRouter(...) 都提供 dependencies 参数,用于实现全局/路由组级别的依赖:
from fastapi import APIRouter, Depends, Query
from fastapi.exceptions import HTTPException
def username_check(username: str = Query(...)):
if username != "zhangsan":
raise HTTPException(status_code=403, detail="Permission denied")
return username
def age_check(age: int = Query(...)):
if age < 18:
raise HTTPException(status_code=403, detail="用户未满18岁")
return age
router = APIRouter(
prefix="/dependency",
tags=["FastAPI笔记 - 依赖注入"],
dependencies=[
Depends(username_check),
Depends(age_check),
],
)
@router.get("/user/login/", summary="user login")
def user_login():
return {"code": 200}
@router.get("/user/info", summary="user info")
def user_info(username: str, age: int):
return {"username": username, "age": age}
注意
全局依赖的返回值默认不会被注入到视图函数;需要在视图中使用返回值时,仍然要在视图参数里显式声明 Depends(...)。