FastAPI 基础
目标
- 理解 FastAPI 的核心概念和优势
- 掌握路由、路径参数、查询参数、请求体的使用
- 掌握路由分组(APIRouter)、状态码、响应模型
- 熟练使用 Pydantic 进行数据验证
- 掌握依赖注入机制
- 能够连接 MySQL 数据库并完成 CRUD 操作
- 独立完成 3 个综合练习
一、FastAPI 简介
1.1 什么是 FastAPI?
FastAPI 是一个现代、高性能的 Python Web 框架,用于构建 API。
核心特性:
- 快速:基于 Starlette 和 Pydantic,性能堪比 NodeJS 和 Go
- 高效编码:利用 Python 类型提示,减少重复代码
- 减少Bug:编辑器自动补全和类型检查
- 易于上手:设计直观,文档自动生成
- 标准驱动:基于 OpenAPI(Swagger)和 JSON Schema
1.2 与其他框架对比
| 特性 | FastAPI | Flask | Django REST |
|---|---|---|---|
| 异步支持 | 原生支持 | 有限支持 | 有限支持 |
| 自动文档 | 内置(Swagger/Redoc) | 需要插件 | 需要插件 |
| 类型验证 | Pydantic | 需要插件 | 序列化器 |
| 性能 | 高 | 中 | 中 |
| 学习曲线 | 低 | 低 | 高 |
1.3 为什么选择 FastAPI?
自动文档 ← 写完代码,文档就有了
类型安全 ← Pydantic 模型验证
异步支持 ← async/await 原生支持
标准化 ← OpenAPI 标准,前后端协作顺畅
二、环境搭建与第一个应用
2.1 安装依赖
# 创建虚拟环境
python -m venv venv
venv\Scripts\activate # Windows
# source venv/bin/activate # Mac/Linux
# 安装 FastAPI 和 Uvicorn
pip install fastapi uvicorn[standard]
# 后续课程需要的依赖
pip install sqlalchemy pymysql pydantic python-jose[cryptography] passlib[bcrypt] python-multipart
2.2 第一个 FastAPI 应用
创建文件 main.py:
from fastapi import FastAPI
app = FastAPI(title="我的第一个FastAPI应用")
@app.get("/")
def read_root():
return {"message": "Hello, FastAPI!"}
@app.get("/items/{item_id}")
def read_item(item_id: int):
return {"item_id": item_id}
2.3 启动服务
uvicorn main:app --reload
main:文件名(不含 .py)app:FastAPI 实例变量名--reload:开发模式,文件修改后自动重启
2.4 访问自动文档
启动后访问以下地址:
| 地址 | 说明 |
|---|---|
| http://localhost:8000/docs | Swagger UI(交互式文档) |
| http://localhost:8000/redoc | ReDoc(只读文档) |
| http://localhost:8000/openapi.json | OpenAPI JSON Schema |
练习:创建一个
/hello/{name}接口,返回{"message": "Hello, {name}!"}
三、路由与 HTTP 方法
3.1 HTTP 方法装饰器
from fastapi import FastAPI
app = FastAPI()
@app.get("/users")
def get_users():
"""获取用户列表"""
return [{"id": 1, "name": "张三"}, {"id": 2, "name": "李四"}]
@app.post("/users")
def create_user():
"""创建用户"""
return {"id": 3, "name": "王五", "status": "created"}
@app.put("/users/{user_id}")
def update_user(user_id: int):
"""更新用户"""
return {"id": user_id, "status": "updated"}
@app.delete("/users/{user_id}")
def delete_user(user_id: int):
"""删除用户"""
return {"id": user_id, "status": "deleted"}
3.2 路径参数
路径参数是 URL 路径的一部分,用 {} 包裹:
@app.get("/users/{user_id}")
def get_user(user_id: int):
return {"id": user_id, "name": "张三"}
@app.get("/users/{user_id}/posts/{post_id}")
def get_user_post(user_id: int, post_id: int):
return {"user_id": user_id, "post_id": post_id}
注意:参数类型会自动转换。如果访问
/users/abc(期望int),会返回 422 验证错误。
3.3 查询参数
查询参数是 URL 中 ? 后面的键值对:
@app.get("/users")
def list_users(skip: int = 0, limit: int = 10, active: bool = True):
"""
skip: 跳过几条(分页偏移)
limit: 返回几条(每页数量)
active: 是否只查询活跃用户
"""
return {"skip": skip, "limit": limit, "active": active}
访问示例:/users?skip=20&limit=5&active=true
3.4 路径参数 vs 查询参数
| 场景 | 使用 |
|---|---|
| 标识一个具体资源 | 路径参数 /users/1 |
| 过滤、排序、分页 | 查询参数 /users?skip=0&limit=10 |
| 资源的层级关系 | 路径参数 /users/1/posts/2 |
3.5 响应状态码
可以在装饰器中指定 HTTP 状态码:
from fastapi import FastAPI, status
app = FastAPI()
@app.post("/users", status_code=201)
def create_user():
return {"id": 1, "name": "新用户"}
@app.delete("/users/{user_id}", status_code=204)
def delete_user(user_id: int):
pass
常用状态码:
| 状态码 | 含义 |
|---|---|
200 OK | 成功(默认) |
201 Created | 资源创建成功 |
204 No Content | 成功但无返回内容 |
400 Bad Request | 请求参数错误 |
401 Unauthorized | 未认证 |
403 Forbidden | 无权限 |
404 Not Found | 资源不存在 |
422 Unprocessable Entity | 数据验证失败 |
500 Internal Server Error | 服务器内部错误 |
提示:推荐使用
status_code=201而不是裸数字,IDE 有自动补全提示。
3.6 response_model 响应模型
在路由装饰器中指定 response_model 可以控制返回结构,过滤敏感字段:
from pydantic import BaseModel
class UserCreate(BaseModel):
username: str
password: str
email: str
class UserResponse(BaseModel):
id: int
username: str
email: str
@app.post("/users", response_model=UserResponse)
def create_user(user: UserCreate):
# 即使返回 dict 包含 password,也会被 response_model 过滤掉
return {
"id": 1,
"username": user.username,
"email": user.email,
"password": "hashed_secret",
}
发送数据:
{
"username": "abc",
"password": "1234",
"email": "44162416@qq.com"
}
返回数据:
{
"id": 1,
"username": "abc",
"email": "44162416@qq.com"
}
password字段被自动过滤,不会暴露给前端。
3.7 路由分组(APIRouter)
当项目变大时,将所有路由写在同一个文件中会导致代码臃肿。使用 APIRouter 可以将路由按模块拆分:
创建独立路由模块 users.py:
from fastapi import APIRouter
router = APIRouter(prefix="/users", tags=["用户管理"])
@router.get("/")
def list_users():
return [{"id": 1, "name": "张三"}]
@router.get("/{user_id}")
def get_user(user_id: int):
return {"id": user_id, "name": "张三"}
@router.post("/", status_code=201)
def create_user():
return {
"id": 3,
"name": "王五",
"status": "created"
}
在主文件中挂载路由:
from fastapi import FastAPI
from users import router as user_router
app = FastAPI()
# 挂载用户路由
app.include_router(user_router)
多个路由模块:
from fastapi import FastAPI
from users import router as user_router
from items import router as item_router
from orders import router as order_router
app = FastAPI()
app.include_router(user_router) # 前缀 /users
app.include_router(item_router) # 前缀 /items
app.include_router(order_router) # 前缀 /orders
3.8 APIRouter 的参数说明
| 参数 | 说明 |
|---|---|
prefix="/users" | 路由前缀,所有路径自动加上 /users |
tags=["用户管理"] | 在自动文档中分组显示 |
dependencies=[Depends(verify_token)] | 该路由下所有接口共用依赖 |
responses={404: {"description": "Not found"}} | 自定义响应 |
3.9 多种参数来源
除了路径参数、查询参数和请求体,FastAPI 还支持从其他来源获取数据:
请求头(Header):
from fastapi import Header
@app.get("/items")
def list_items(user_agent: str | None = Header(None)):
return {"user_agent": user_agent}
curl -X 'GET' \
'http://localhost:8000/items/' \
-H 'accept: application/json' \
-H 'user-agent: abc' # 头信息中发送数据到服务器
Cookie:
from fastapi import Cookie
@app.get("/items")
def list_items(session_id: str | None = Cookie(None)):
return {"session_id": session_id}
curl -X 'GET' \
'http://localhost:8000/items/' \
-H 'accept: application/json' \
-H 'Cookie: session_id=56' # cookie 携带参数到服务器
表单数据(Form):
from fastapi import Form
@app.post("/login")
def login(username: str = Form(...), password: str = Form(...)):
return {"username": username}
curl -X 'POST' \
'http://localhost:8000/items/login' \
-H 'accept: application/json' \
-H 'Content-Type: application/x-www-form-urlencoded' \
-d 'username=abc&password=123' # 表单数据格式提交数据
注意:使用 Form 需要先安装
python-multipart。
3.10 路由优先级
FastAPI 按定义顺序匹配路由。先定义的优先级更高:
# 正确顺序:具体的路由在前
@app.get("/users/me")
def read_current_user():
return {"id": 1, "name": "当前用户"}
@app.get("/users/{user_id}")
def read_user(user_id: int):
return {"id": user_id}
陷阱:
如果先定义
"/users/{user_id}",那么访问/users/me时,me会被当作user_id,导致 422 错误。
四、请求体与 Pydantic 模型
4.1 Pydantic 模型定义
请求体使用 Pydantic 的 BaseModel 来定义数据结构:
from pydantic import BaseModel, Field
from typing import Optional
class UserCreate(BaseModel):
username: str
email: str
age: int
is_active: bool = True
bio: Optional[str] = None
@app.post("/users")
def create_user(user: UserCreate):
return {"id": 1, **user.model_dump()}
user: User自动解析并验证请求体 JSON user.model_dump()将 Pydantic 模型转为字典 **...把字典内容“展开”合并到新字典中
4.2 字段验证
Pydantic 支持丰富的字段验证:
from pydantic import BaseModel, Field, EmailStr
from typing import Optional
class UserCreate(BaseModel):
username: str = Field(..., min_length=2, max_length=20, description="用户名")
email: str = Field(..., pattern=r"^[a-zA-Z0-9_.+-]+@[a-zA-Z0-9-]+\.[a-zA-Z0-9-.]+$")
age: int = Field(..., ge=0, le=150, description="年龄")
score: float = Field(default=0.0, ge=0.0, le=100.0)
tags: list[str] = Field(default_factory=list, max_length=5)
bio: Optional[str] = Field(None, max_length=500)
@app.post("/users")
def create_user(user: UserCreate):
return {"id": 1, **user.model_dump()}
常用验证器:
| 验证器 | 说明 |
|---|---|
Field(..., min_length=N) | 最小长度 |
Field(..., max_length=N) | 最大长度 |
Field(..., ge=N) | 大于等于 |
Field(..., le=N) | 小于等于 |
Field(..., gt=N) | 大于 |
Field(..., lt=N) | 小于 |
Field(..., pattern=r"正则") | 正则匹配 |
Field(default_factory=list) | 默认值工厂 |
4.3 嵌套模型
class Address(BaseModel):
province: str
city: str
detail: str
class UserCreate(BaseModel):
username: str
email: str
address: Address
@app.post("/users")
def create_user(user: UserCreate):
return user
请求示例:
{
"username": "张三",
"email": "zhang@example.com",
"address": {
"province": "北京",
"city": "北京",
"detail": "朝阳区xxx路xxx号"
}
}
五、依赖注入
5.1 什么是依赖注入?
依赖注入(DI)是 FastAPI 的核心机制,用于共享逻辑、减少重复代码。
from fastapi import Depends, FastAPI, HTTPException
app = FastAPI()
def common_params(skip: int = 0, limit: int = 10):
"""通用分页参数"""
return {"skip": skip, "limit": limit}
@app.get("/items")
def list_items(params: dict = Depends(common_params)):
return {"items": [], "pagination": params}
@app.get("/users")
def list_users(params: dict = Depends(common_params)):
return {"users": [], "pagination": params}
5.2 依赖链
依赖可以嵌套,形成依赖链:
def verify_token(token: str):
if token != "secret-token":
raise HTTPException(status_code=401, detail="Invalid token")
return token
def get_current_user(token: str = Depends(verify_token)):
return {"id": 1, "username": "admin", "token": token}
@app.get("/me")
def get_me(user=Depends(get_current_user)):
return user
5.3 在路由上统一使用依赖
from fastapi import APIRouter, Depends
router = APIRouter(dependencies=[Depends(verify_token)])
@router.get("/items")
def list_items():
return {"items": []}
@router.post("/items")
def create_item():
return {"item": "created"}
此路由下所有接口都会自动验证 token。
六、连接 MySQL 数据库
6.1 安装依赖
pip install sqlalchemy pymysql
6.2 数据库配置
创建 database.py:
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker, declarative_base
DATABASE_URL = "mysql+pymysql://root:password@localhost:3306/fastapi_demo?charset=utf8mb4"
engine = create_engine(DATABASE_URL)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
Base = declarative_base()
def get_db():
"""获取数据库会话的依赖函数"""
db = SessionLocal() # 1. 创建数据库会话
try:
yield db # 2. 将 db 交给调用者(如路由函数)
finally:
db.close() # 3. 无论成功或出错,都关闭会话
yield db到底做了什么?
yield在这里把get_db()变成了一个 生成器函数(generator function)。- 当 FastAPI 调用这个依赖时:
- 执行到
yield db,暂停函数执行,并将db对象“返回”给使用它的路由函数。- 路由函数使用这个
db对象执行数据库操作(比如查询、写入)。- 当路由处理完成(无论成功还是抛出异常),FastAPI 会自动 继续执行
get_db()中yield之后的代码。- 于是
finally块被执行,db.close()被调用,确保数据库连接被释放。💡 这种模式叫做 上下文管理(Context Management),类似于
with语句,但通过生成器实现。
6.3 定义模型
创建 models.py:
from sqlalchemy import Column, Integer, String, DateTime, Boolean, func
from database import Base
class Student(Base):
__tablename__ = "students"
id = Column(Integer, primary_key=True, index=True)
name = Column(String(50), nullable=False)
email = Column(String(100), unique=True, nullable=False)
age = Column(Integer, default=0)
is_active = Column(Boolean, default=True)
created_at = Column(DateTime, server_default=func.now())
6.4 创建数据表
# 在 Python 交互环境或脚本中运行
from database import engine, Base
Base.metadata.create_all(bind=engine)
6.5 CRUD 操作示例
创建 main.py:
from fastapi import FastAPI, Depends, HTTPException
from sqlalchemy.orm import Session
from database import get_db, engine, Base
from models import Student
from pydantic import BaseModel, EmailStr
from typing import Optional
# 首次运行时创建表
Base.metadata.create_all(bind=engine)
app = FastAPI()
# --- Pydantic 模型 ---
class StudentCreate(BaseModel):
name: str
email: str
age: int = 0
class StudentUpdate(BaseModel):
name: Optional[str] = None
email: Optional[str] = None
age: Optional[int] = None
is_active: Optional[bool] = None
class StudentResponse(BaseModel):
id: int
name: str
email: str
age: int
is_active: bool
class Config:
from_attributes = True
# --- CRUD 接口 ---
@app.post("/students", response_model=StudentResponse, status_code=201)
def create_student(student: StudentCreate, db: Session = Depends(get_db)):
db_student = Student(**student.model_dump())
db.add(db_student)
db.commit()
db.refresh(db_student)
return db_student
@app.get("/students", response_model=list[StudentResponse])
def list_students(skip: int = 0, limit: int = 10, db: Session = Depends(get_db)):
students = db.query(Student).offset(skip).limit(limit).all()
return students
@app.get("/students/{student_id}", response_model=StudentResponse)
def get_student(student_id: int, db: Session = Depends(get_db)):
student = db.query(Student).filter(Student.id == student_id).first()
if not student:
raise HTTPException(status_code=404, detail="学生不存在")
return student
@app.put("/students/{student_id}", response_model=StudentResponse)
def update_student(student_id: int, student_data: StudentUpdate, db: Session = Depends(get_db)):
student = db.query(Student).filter(Student.id == student_id).first()
if not student:
raise HTTPException(status_code=404, detail="学生不存在")
for key, value in student_data.model_dump(exclude_unset=True).items():
setattr(student, key, value)
db.commit()
db.refresh(student)
return student
@app.delete("/students/{student_id}", status_code=204)
def delete_student(student_id: int, db: Session = Depends(get_db)):
student = db.query(Student).filter(Student.id == student_id).first()
if not student:
raise HTTPException(status_code=404, detail="学生不存在")
db.delete(student)
db.commit()
return None
七、练习 1:学生管理系统 API
目标
完成一个完整的学生管理 API,支持以下功能:
需求
-
添加学生
POST /students- 验证姓名不能为空
- 邮箱格式验证
- 邮箱不能重复
-
学生列表
GET /students- 支持分页(skip、limit)
- 支持按姓名搜索(
?name=张) - 支持按活跃状态筛选(
?is_active=true)
-
学生详情
GET /students/{id} -
更新学生
PUT /students/{id}- 支持部分更新
-
删除学生
DELETE /students/{id}
参考代码:
from fastapi import FastAPI, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from database import get_db, engine, Base
from models import Student
from pydantic import BaseModel, Field
from typing import Optional
Base.metadata.create_all(bind=engine)
app = FastAPI()
class StudentCreate(BaseModel):
name: str = Field(..., min_length=1, max_length=50)
email: str = Field(..., min_length=5, max_length=100)
age: int = Field(default=0, ge=0, le=150)
class StudentUpdate(BaseModel):
name: Optional[str] = Field(None, min_length=1, max_length=50)
email: Optional[str] = None
age: Optional[int] = None
is_active: Optional[bool] = None
class StudentResponse(BaseModel):
id: int
name: str
email: str
age: int
is_active: bool
class Config:
from_attributes = True
@app.post("/students", response_model=StudentResponse, status_code=201)
def create_student(student: StudentCreate, db: Session = Depends(get_db)):
# 检查邮箱是否重复
existing = db.query(Student).filter(Student.email == student.email).first()
if existing:
raise HTTPException(status_code=400, detail="邮箱已注册")
db_student = Student(**student.model_dump())
db.add(db_student)
db.commit()
db.refresh(db_student)
return db_student
@app.get("/students", response_model=list[StudentResponse])
def list_students(
skip: int = 0,
limit: int = 10,
name: Optional[str] = Query(None, min_length=1),
is_active: Optional[bool] = None,
db: Session = Depends(get_db),
):
query = db.query(Student)
if name:
query = query.filter(Student.name.contains(name))
if is_active is not None:
query = query.filter(Student.is_active == is_active)
return query.offset(skip).limit(limit).all()
@app.get("/students/{student_id}", response_model=StudentResponse)
def get_student(student_id: int, db: Session = Depends(get_db)):
student = db.query(Student).filter(Student.id == student_id).first()
if not student:
raise HTTPException(status_code=404, detail="学生不存在")
return student
@app.put("/students/{student_id}", response_model=StudentResponse)
def update_student(student_id: int, data: StudentUpdate, db: Session = Depends(get_db)):
student = db.query(Student).filter(Student.id == student_id).first()
if not student:
raise HTTPException(status_code=404, detail="学生不存在")
for key, value in data.model_dump(exclude_unset=True).items():
setattr(student, key, value)
db.commit()
db.refresh(student)
return student
@app.delete("/students/{student_id}", status_code=204)
def delete_student(student_id: int, db: Session = Depends(get_db)):
student = db.query(Student).filter(Student.id == student_id).first()
if not student:
raise HTTPException(status_code=404, detail="学生不存在")
db.delete(student)
db.commit()
return None
八、练习 2:商品分类系统
目标
实现分类和商品的一对多关系 API。
数据库模型
from sqlalchemy import Column, Integer, String, Float, ForeignKey, DateTime, func
from database import Base
class Category(Base):
__tablename__ = "categories"
id = Column(Integer, primary_key=True, index=True)
name = Column(String(50), unique=True, nullable=False)
description = Column(String(200), default="")
created_at = Column(DateTime, server_default=func.now())
# 关系
products = relationship("Product", back_populates="category")
class Product(Base):
__tablename__ = "products"
id = Column(Integer, primary_key=True, index=True)
name = Column(String(100), nullable=False)
price = Column(Float, nullable=False)
description = Column(String(500), default="")
stock = Column(Integer, default=0)
category_id = Column(Integer, ForeignKey("categories.id"))
created_at = Column(DateTime, server_default=func.now())
# 关系
category = relationship("Category", back_populates="products")
注意:需要在文件顶部添加
from sqlalchemy.orm import relationship
需求
-
分类管理
POST /categories- 创建分类GET /categories- 分类列表GET /categories/{id}- 分类详情
-
商品管理
POST /products- 创建商品GET /products- 商品列表(支持按分类筛选、分页、排序)GET /products/{id}- 商品详情(包含分类信息)PUT /products/{id}- 更新商品DELETE /products/{id}- 删除商品
参考代码:
from fastapi import FastAPI, Depends, HTTPException, Query
from sqlalchemy.orm import Session, relationship
from database import get_db, engine, Base
from pydantic import BaseModel, Field
from typing import Optional
from datetime import datetime
# --- 模型 ---
from sqlalchemy import Column, Integer, String, Float, ForeignKey, DateTime, func
from sqlalchemy.orm import declarative_base, relationship
Base = declarative_base()
class CategoryModel(Base):
__tablename__ = "categories"
id = Column(Integer, primary_key=True, index=True)
name = Column(String(50), unique=True, nullable=False)
description = Column(String(200), default="")
created_at = Column(DateTime, server_default=func.now())
products = relationship("ProductModel", back_populates="category")
class ProductModel(Base):
__tablename__ = "products"
id = Column(Integer, primary_key=True, index=True)
name = Column(String(100), nullable=False)
price = Column(Float, nullable=False)
description = Column(String(500), default="")
stock = Column(Integer, default=0)
category_id = Column(Integer, ForeignKey("categories.id"))
created_at = Column(DateTime, server_default=func.now())
category = relationship("CategoryModel", back_populates="products")
Base.metadata.create_all(bind=engine)
# --- Pydantic 模型 ---
class CategoryCreate(BaseModel):
name: str = Field(..., min_length=1, max_length=50)
description: str = ""
class CategoryResponse(BaseModel):
id: int
name: str
description: str
class Config:
from_attributes = True
class ProductCreate(BaseModel):
name: str = Field(..., min_length=1, max_length=100)
price: float = Field(..., gt=0)
description: str = ""
stock: int = Field(default=0, ge=0)
category_id: int
class ProductResponse(BaseModel):
id: int
name: str
price: float
description: str
stock: int
category_id: int
category: Optional[CategoryResponse] = None
class Config:
from_attributes = True
# --- 路由 ---
app = FastAPI()
@app.post("/categories", response_model=CategoryResponse, status_code=201)
def create_category(cat: CategoryCreate, db: Session = Depends(get_db)):
existing = db.query(CategoryModel).filter(CategoryModel.name == cat.name).first()
if existing:
raise HTTPException(status_code=400, detail="分类已存在")
db_cat = CategoryModel(**cat.model_dump())
db.add(db_cat)
db.commit()
db.refresh(db_cat)
return db_cat
@app.get("/categories", response_model=list[CategoryResponse])
def list_categories(db: Session = Depends(get_db)):
return db.query(CategoryModel).all()
@app.post("/products", response_model=ProductResponse, status_code=201)
def create_product(product: ProductCreate, db: Session = Depends(get_db)):
cat = db.query(CategoryModel).filter(CategoryModel.id == product.category_id).first()
if not cat:
raise HTTPException(status_code=404, detail="分类不存在")
db_product = ProductModel(**product.model_dump())
db.add(db_product)
db.commit()
db.refresh(db_product)
return db_product
@app.get("/products", response_model=list[ProductResponse])
def list_products(
skip: int = 0,
limit: int = 10,
category_id: Optional[int] = None,
sort_by: str = Query("id", pattern="^(id|price|stock|name)$"),
db: Session = Depends(get_db),
):
query = db.query(ProductModel)
if category_id:
query = query.filter(ProductModel.category_id == category_id)
if sort_by == "price":
query = query.order_by(ProductModel.price)
elif sort_by == "stock":
query = query.order_by(ProductModel.stock)
elif sort_by == "name":
query = query.order_by(ProductModel.name)
return query.offset(skip).limit(limit).all()
@app.get("/products/{product_id}", response_model=ProductResponse)
def get_product(product_id: int, db: Session = Depends(get_db)):
product = db.query(ProductModel).filter(ProductModel.id == product_id).first()
if not product:
raise HTTPException(status_code=404, detail="商品不存在")
return product
九、练习 3:用户认证基础
目标
实现基于 JWT 的用户认证。
安装依赖
pip install python-jose[cryptography] passlib[bcrypt]
参考答案
from fastapi import FastAPI, Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm
from jose import JWTError, jwt
from passlib.context import CryptContext
from pydantic import BaseModel
from datetime import datetime, timedelta
from typing import Optional
app = FastAPI()
# --- 配置 ---
SECRET_KEY = "your-secret-key-change-in-production"
ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 30
# --- 密码加密 ---
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="login")
def hash_password(password: str) -> str:
return pwd_context.hash(password)
def verify_password(plain: str, hashed: str) -> bool:
return pwd_context.verify(plain, hashed)
def create_token(data: dict, expires_delta: timedelta | None = None) -> str:
to_encode = data.copy()
expire = datetime.utcnow() + (expires_delta or timedelta(minutes=15))
to_encode.update({"exp": expire})
return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
# --- 模拟数据库 ---
fake_users_db = {
"admin": {
"username": "admin",
"hashed_password": hash_password("admin123"),
"role": "admin",
}
}
# --- Pydantic 模型 ---
class Token(BaseModel):
access_token: str
token_type: str
class UserResponse(BaseModel):
username: str
role: str
# --- 接口 ---
@app.post("/login", response_model=Token)
def login(form: OAuth2PasswordRequestForm = Depends()):
user = fake_users_db.get(form.username)
if not user or not verify_password(form.password, user["hashed_password"]):
raise HTTPException(status_code=401, detail="用户名或密码错误")
token = create_token({"sub": user["username"], "role": user["role"]}, timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES))
return {"access_token": token, "token_type": "bearer"}
def get_current_user(token: str = Depends(oauth2_scheme)) -> dict:
try:
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
username: str = payload.get("sub")
if username is None:
raise HTTPException(status_code=401, detail="Invalid token")
return {"username": username, "role": payload.get("role")}
except JWTError:
raise HTTPException(status_code=401, detail="Token 无效")
@app.get("/me", response_model=UserResponse)
def get_me(user: dict = Depends(get_current_user)):
return user
@app.get("/protected")
def protected_route(user: dict = Depends(get_current_user)):
return {"message": f"Hello {user['username']}, this is a protected resource"}
测试方式
# 1. 登录获取 Token
curl -X POST http://localhost:8000/login -d "username=admin&password=admin123"
# 2. 使用 Token 访问受保护接口
curl http://localhost:8000/me -H "Authorization: Bearer <token>"
总结
| 知识点 | 核心内容 |
|---|---|
| 路由 | GET/POST/PUT/DELETE 装饰器 |
| 路径参数 | /users/{id} 类型自动转换 |
| 查询参数 | ?skip=0&limit=10 默认值 |
| 状态码 | status_code=201 等 |
| 响应模型 | response_model 过滤返回字段 |
| 路由分组 | APIRouter 模块化拆分 |
| 参数来源 | Header、Cookie、Form |
| 路由优先级 | 具体路径在前,通用路径在后 |
| 请求体 | Pydantic BaseModel 验证 |
| 依赖注入 | Depends() 共享逻辑 |
| 数据库 | SQLAlchemy ORM + PyMySQL |
| 认证 | JWT Token + OAuth2PasswordBearer |

1378

被折叠的 条评论
为什么被折叠?



