Fastapi框架深度解析

Fastapi 2026-04-20 12
预计阅读时间:66 分钟

FastAPI 深度解析:现代 Python Web 框架的典范

一、FastAPI 的核心理念

FastAPI 是一个现代、快速(高性能)的 Web 框架,用于构建 API。它的设计哲学建立在三个核心支柱之上:

  • 快速:性能媲美 NodeJS 和 Go(基于 Starlette 和 Pydantic)
  • 快速开发:开发速度提升约 200% 到 300%
  • 减少错误:减少约 40% 的人为错误
  • 直观:出色的编辑器支持,无处不在的自动补全
  • 简单:设计易于使用和学习
  • 简短:代码重复最小化
  • 健壮:生产就绪,具有自动交互式文档
  • 标准化:基于 OpenAPI(原 Swagger)和 JSON Schema

1.1 第一个 FastAPI 应用

"""
安装 FastAPI 和 ASGI 服务器:
pip install fastapi uvicorn[standard]
"""

from fastapi import FastAPI
from typing import Optional

# 创建 FastAPI 应用实例
app = FastAPI(
    title="我的第一个 FastAPI",
    description="FastAPI 学习笔记示例",
    version="1.0.0"
)

@app.get("/")
async def root():
    """根路径,返回欢迎信息"""
    return {"message": "Hello, FastAPI!"}

@app.get("/hello/{name}")
async def say_hello(name: str, age: Optional[int] = None):
    """
    打招呼接口

    - **name**: 名字(路径参数)
    - **age**: 年龄(可选查询参数)
    """
    response = {"message": f"Hello, {name}!"}
    if age:
        response["age"] = age
    return response

@app.get("/items/")
async def list_items(skip: int = 0, limit: int = 10):
    """
    获取物品列表(分页)

    - **skip**: 跳过的数量
    - **limit**: 返回的最大数量
    """
    # 模拟数据库查询
    items = [{"id": i, "name": f"Item {i}"} for i in range(100)]
    return items[skip:skip + limit]

# 运行应用:
# uvicorn main:app --reload
# 访问 http://127.0.0.1:8000/docs 查看自动生成的 API 文档
# 访问 http://127.0.0.1:8000/redoc 查看 ReDoc 文档

1.2 FastAPI 项目结构

"""
推荐的 FastAPI 项目结构:

myapp/
├── app/
│   ├── __init__.py
│   ├── main.py                 # 应用入口
│   ├── config.py               # 配置管理
│   ├── database.py             # 数据库连接
│   ├── models/                 # 数据模型
│   │   ├── __init__.py
│   │   ├── user.py
│   │   └── post.py
│   ├── schemas/                # Pydantic 模型(请求/响应)
│   │   ├── __init__.py
│   │   ├── user.py
│   │   └── post.py
│   ├── api/                    # API 路由
│   │   ├── __init__.py
│   │   ├── v1/
│   │   │   ├── __init__.py
│   │   │   ├── endpoints/
│   │   │   │   ├── users.py
│   │   │   │   └── posts.py
│   │   │   └── router.py
│   │   └── deps.py             # 依赖项
│   ├── crud/                   # CRUD 操作
│   │   ├── __init__.py
│   │   ├── base.py
│   │   ├── user.py
│   │   └── post.py
│   ├── core/                   # 核心功能
│   │   ├── __init__.py
│   │   ├── security.py         # 认证和安全
│   │   ├── exceptions.py       # 自定义异常
│   │   └── pagination.py       # 分页工具
│   ├── db/                     # 数据库相关
│   │   ├── __init__.py
│   │   ├── base.py             # 基类
│   │   └── session.py          # 会话管理
│   └── utils/                  # 工具函数
│       ├── __init__.py
│       └── helpers.py
├── alembic/                    # 数据库迁移
│   └── versions/
├── tests/                      # 测试
│   ├── __init__.py
│   ├── conftest.py
│   ├── test_api/
│   └── test_crud/
├── .env                        # 环境变量
├── .env.example
├── requirements.txt
├── alembic.ini
└── docker-compose.yml
"""

二、请求与响应模型

2.1 Pydantic 模型定义

# app/schemas/user.py
from pydantic import BaseModel, EmailStr, Field, validator, constr
from typing import Optional, List
from datetime import datetime
from enum import Enum

class UserRole(str, Enum):
    """用户角色枚举"""
    ADMIN = "admin"
    MODERATOR = "moderator"
    USER = "user"
    GUEST = "guest"

class UserBase(BaseModel):
    """用户基础模型"""
    username: str = Field(..., min_length=3, max_length=50, description="用户名")
    email: EmailStr = Field(..., description="邮箱地址")
    full_name: Optional[str] = Field(None, max_length=100, description="全名")
    role: UserRole = Field(default=UserRole.USER, description="用户角色")

    @validator('username')
    def username_alphanumeric(cls, v):
        """验证用户名只包含字母数字和下划线"""
        if not v.replace('_', '').isalnum():
            raise ValueError('用户名只能包含字母、数字和下划线')
        return v

class UserCreate(UserBase):
    """创建用户请求模型"""
    password: str = Field(..., min_length=8, description="密码")
    password_confirm: str = Field(..., description="确认密码")

    @validator('password_confirm')
    def passwords_match(cls, v, values):
        """验证两次密码是否一致"""
        if 'password' in values and v != values['password']:
            raise ValueError('两次输入的密码不一致')
        return v

    @validator('password')
    def password_strength(cls, v):
        """验证密码强度"""
        if not any(c.isupper() for c in v):
            raise ValueError('密码必须包含至少一个大写字母')
        if not any(c.islower() for c in v):
            raise ValueError('密码必须包含至少一个小写字母')
        if not any(c.isdigit() for c in v):
            raise ValueError('密码必须包含至少一个数字')
        return v

class UserUpdate(BaseModel):
    """更新用户请求模型"""
    username: Optional[str] = Field(None, min_length=3, max_length=50)
    email: Optional[EmailStr] = None
    full_name: Optional[str] = Field(None, max_length=100)
    role: Optional[UserRole] = None
    password: Optional[str] = Field(None, min_length=8)
    is_active: Optional[bool] = None

class UserInDB(UserBase):
    """数据库用户模型(包含敏感信息)"""
    id: int
    hashed_password: str
    is_active: bool = True
    created_at: datetime
    updated_at: Optional[datetime] = None

    class Config:
        from_attributes = True  # 允许从 ORM 模型创建

class User(UserBase):
    """用户响应模型(排除敏感信息)"""
    id: int
    is_active: bool
    created_at: datetime

    class Config:
        from_attributes = True

class UserList(BaseModel):
    """用户列表响应"""
    users: List[User]
    total: int
    page: int
    size: int
    pages: int

class Token(BaseModel):
    """JWT Token 响应"""
    access_token: str
    token_type: str = "bearer"
    expires_in: int

class TokenData(BaseModel):
    """Token 中的数据"""
    user_id: Optional[int] = None
    username: Optional[str] = None
    role: Optional[UserRole] = None

# app/schemas/post.py
from pydantic import BaseModel, Field, HttpUrl
from typing import Optional, List
from datetime import datetime
from .user import User

class PostBase(BaseModel):
    """文章基础模型"""
    title: str = Field(..., min_length=1, max_length=200, description="标题")
    content: str = Field(..., min_length=1, description="内容")
    summary: Optional[str] = Field(None, max_length=500, description="摘要")
    cover_image: Optional[HttpUrl] = Field(None, description="封面图片URL")
    is_published: bool = Field(default=False, description="是否发布")

class PostCreate(PostBase):
    """创建文章请求模型"""
    category_id: Optional[int] = Field(None, description="分类ID")
    tag_names: Optional[List[str]] = Field(default=[], description="标签名称列表")

class PostUpdate(BaseModel):
    """更新文章请求模型"""
    title: Optional[str] = Field(None, min_length=1, max_length=200)
    content: Optional[str] = Field(None, min_length=1)
    summary: Optional[str] = Field(None, max_length=500)
    cover_image: Optional[HttpUrl] = None
    category_id: Optional[int] = None
    is_published: Optional[bool] = None

class PostInDB(PostBase):
    """数据库文章模型"""
    id: int
    slug: str
    author_id: int
    category_id: Optional[int]
    views_count: int = 0
    likes_count: int = 0
    comments_count: int = 0
    published_at: Optional[datetime]
    created_at: datetime
    updated_at: Optional[datetime]

    class Config:
        from_attributes = True

class Post(PostInDB):
    """文章响应模型"""
    author: Optional[User] = None
    category: Optional['Category'] = None
    tags: List['Tag'] = []

    class Config:
        from_attributes = True

class PostList(BaseModel):
    """文章列表响应"""
    posts: List[Post]
    total: int
    page: int
    size: int
    pages: int

class Category(BaseModel):
    """分类模型"""
    id: int
    name: str
    slug: str
    description: Optional[str] = None
    post_count: Optional[int] = 0

    class Config:
        from_attributes = True

class Tag(BaseModel):
    """标签模型"""
    id: int
    name: str
    slug: str
    post_count: Optional[int] = 0

    class Config:
        from_attributes = True

# 解决循环引用
Post.model_rebuild()

2.2 路径参数与查询参数

# app/api/v1/endpoints/posts.py
from fastapi import APIRouter, Path, Query, Depends, HTTPException, status
from typing import Optional, List
from enum import Enum

router = APIRouter(prefix="/posts", tags=["posts"])

class PostSortField(str, Enum):
    """文章排序字段"""
    CREATED_AT = "created_at"
    UPDATED_AT = "updated_at"
    VIEWS = "views_count"
    LIKES = "likes_count"
    TITLE = "title"

class SortOrder(str, Enum):
    """排序方向"""
    ASC = "asc"
    DESC = "desc"

@router.get("/")
async def list_posts(
    # 分页参数
    page: int = Query(1, ge=1, description="页码"),
    size: int = Query(10, ge=1, le=100, description="每页数量"),

    # 排序参数
    sort_by: PostSortField = Query(
        PostSortField.CREATED_AT, 
        description="排序字段"
    ),
    sort_order: SortOrder = Query(
        SortOrder.DESC, 
        description="排序方向"
    ),

    # 筛选参数
    category: Optional[str] = Query(None, description="分类 slug"),
    tag: Optional[str] = Query(None, description="标签 slug"),
    author_id: Optional[int] = Query(None, description="作者 ID"),
    is_published: Optional[bool] = Query(None, description="是否已发布"),

    # 搜索参数
    q: Optional[str] = Query(None, min_length=1, description="搜索关键词"),

    # 时间范围
    created_after: Optional[datetime] = Query(None, description="创建时间之后"),
    created_before: Optional[datetime] = Query(None, description="创建时间之前"),
):
    """
    获取文章列表

    支持分页、排序、筛选和搜索
    """
    # 构建查询条件
    query = db.query(Post)

    if category:
        query = query.join(Category).filter(Category.slug == category)

    if tag:
        query = query.join(Post.tags).filter(Tag.slug == tag)

    if author_id:
        query = query.filter(Post.author_id == author_id)

    if is_published is not None:
        query = query.filter(Post.is_published == is_published)

    if q:
        query = query.filter(
            db.or_(
                Post.title.ilike(f"%{q}%"),
                Post.content.ilike(f"%{q}%")
            )
        )

    if created_after:
        query = query.filter(Post.created_at >= created_after)

    if created_before:
        query = query.filter(Post.created_at <= created_before)

    # 排序
    sort_column = getattr(Post, sort_by.value)
    if sort_order == SortOrder.DESC:
        query = query.order_by(sort_column.desc())
    else:
        query = query.order_by(sort_column.asc())

    # 分页
    total = query.count()
    posts = query.offset((page - 1) * size).limit(size).all()

    return {
        "posts": posts,
        "total": total,
        "page": page,
        "size": size,
        "pages": (total + size - 1) // size
    }

@router.get("/{slug}")
async def get_post(
    slug: str = Path(
        ..., 
        min_length=1, 
        max_length=200,
        regex=r"^[a-z0-9]+(?:-[a-z0-9]+)*$",
        description="文章 slug",
        example="hello-world"
    ),
    include_body: bool = Query(True, description="是否包含正文"),
):
    """
    获取单篇文章

    - **slug**: 文章的唯一标识符
    - **include_body**: 是否在响应中包含文章正文
    """
    post = db.query(Post).filter(Post.slug == slug).first()

    if not post:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail=f"文章 '{slug}' 不存在"
        )

    # 增加浏览量
    post.views_count += 1
    db.commit()

    if not include_body:
        post.content = None  # 不返回正文

    return post

@router.get("/{year}/{month}/{day}/{slug}")
async def get_post_by_date(
    year: int = Path(..., ge=2000, le=2100, description="年份"),
    month: int = Path(..., ge=1, le=12, description="月份"),
    day: int = Path(..., ge=1, le=31, description="日期"),
    slug: str = Path(..., description="文章 slug"),
):
    """
    通过日期和 slug 获取文章

    这种 URL 格式常用于博客系统
    """
    # 验证日期是否有效
    try:
        date = datetime(year, month, day)
    except ValueError:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="无效的日期"
        )

    post = db.query(Post).filter(
        Post.slug == slug,
        db.func.date(Post.published_at) == date.date()
    ).first()

    if not post:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="文章不存在"
        )

    return post

2.3 请求体与表单数据

# app/api/v1/endpoints/posts.py (续)
from fastapi import File, UploadFile, Form
from pathlib import Path as FilePath
import shutil
import uuid

@router.post("/", response_model=Post, status_code=status.HTTP_201_CREATED)
async def create_post(
    post: PostCreate,
    current_user: User = Depends(get_current_user),
):
    """
    创建新文章

    - **post**: 文章数据
    - 需要认证
    """
    # 检查 slug 是否已存在
    slug = generate_slug(post.title)
    existing = db.query(Post).filter(Post.slug == slug).first()
    if existing:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="文章标题已存在,请修改标题"
        )

    # 创建文章
    db_post = Post(
        **post.model_dump(exclude={'category_id', 'tag_names'}),
        slug=slug,
        author_id=current_user.id,
        category_id=post.category_id
    )

    if db_post.is_published:
        db_post.published_at = datetime.utcnow()

    db.add(db_post)
    db.flush()

    # 添加标签
    if post.tag_names:
        for tag_name in post.tag_names:
            tag = db.query(Tag).filter(Tag.name == tag_name).first()
            if not tag:
                tag = Tag(name=tag_name, slug=generate_slug(tag_name))
                db.add(tag)
                db.flush()
            db_post.tags.append(tag)

    db.commit()
    db.refresh(db_post)

    return db_post

@router.put("/{post_id}", response_model=Post)
async def update_post(
    post_id: int = Path(..., description="文章 ID"),
    post_update: PostUpdate = None,
    current_user: User = Depends(get_current_user),
):
    """
    更新文章

    - 需要认证
    - 只有作者或管理员可以更新
    """
    post = db.query(Post).filter(Post.id == post_id).first()

    if not post:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="文章不存在"
        )

    # 权限检查
    if post.author_id != current_user.id and current_user.role != UserRole.ADMIN:
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="没有权限更新此文章"
        )

    # 更新字段
    update_data = post_update.model_dump(exclude_unset=True)

    # 如果更新了标题,重新生成 slug
    if 'title' in update_data:
        update_data['slug'] = generate_slug(update_data['title'])

    # 如果设置为发布且之前未发布,设置发布时间
    if update_data.get('is_published') and not post.published_at:
        update_data['published_at'] = datetime.utcnow()

    for field, value in update_data.items():
        setattr(post, field, value)

    db.commit()
    db.refresh(post)

    return post

@router.delete("/{post_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_post(
    post_id: int = Path(..., description="文章 ID"),
    current_user: User = Depends(get_current_user),
):
    """
    删除文章

    - 需要认证
    - 只有作者或管理员可以删除
    """
    post = db.query(Post).filter(Post.id == post_id).first()

    if not post:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="文章不存在"
        )

    # 权限检查
    if post.author_id != current_user.id and current_user.role != UserRole.ADMIN:
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="没有权限删除此文章"
        )

    db.delete(post)
    db.commit()

    return None

@router.post("/{post_id}/upload-cover")
async def upload_cover_image(
    post_id: int,
    file: UploadFile = File(..., description="封面图片"),
    current_user: User = Depends(get_current_user),
):
    """
    上传文章封面图片

    - 支持格式:jpg, jpeg, png, gif
    - 最大大小:5MB
    """
    # 验证文件类型
    allowed_types = ["image/jpeg", "image/png", "image/gif"]
    if file.content_type not in allowed_types:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail=f"不支持的文件类型。允许的类型:{', '.join(allowed_types)}"
        )

    # 验证文件大小
    file.file.seek(0, 2)  # 移动到文件末尾
    file_size = file.file.tell()
    file.file.seek(0)  # 重置到开头

    max_size = 5 * 1024 * 1024  # 5MB
    if file_size > max_size:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail=f"文件过大。最大允许 {max_size // (1024*1024)}MB"
        )

    # 检查文章是否存在
    post = db.query(Post).filter(Post.id == post_id).first()
    if not post:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="文章不存在"
        )

    # 权限检查
    if post.author_id != current_user.id and current_user.role != UserRole.ADMIN:
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="没有权限上传封面"
        )

    # 保存文件
    file_extension = file.filename.split('.')[-1]
    filename = f"{uuid.uuid4()}.{file_extension}"
    upload_dir = FilePath("uploads/covers")
    upload_dir.mkdir(parents=True, exist_ok=True)

    file_path = upload_dir / filename
    with file_path.open("wb") as buffer:
        shutil.copyfileobj(file.file, buffer)

    # 更新文章封面 URL
    cover_url = f"/uploads/covers/{filename}"
    post.cover_image = cover_url
    db.commit()

    return {"cover_url": cover_url}

@router.post("/{post_id}/like")
async def like_post(
    post_id: int,
    current_user: User = Depends(get_current_user),
):
    """
    点赞/取消点赞文章

    - 需要认证
    """
    post = db.query(Post).filter(Post.id == post_id).first()

    if not post:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="文章不存在"
        )

    # 检查是否已点赞
    like = db.query(Like).filter(
        Like.user_id == current_user.id,
        Like.post_id == post_id
    ).first()

    if like:
        db.delete(like)
        post.likes_count -= 1
        action = "unliked"
    else:
        like = Like(user_id=current_user.id, post_id=post_id)
        db.add(like)
        post.likes_count += 1
        action = "liked"

    db.commit()

    return {
        "action": action,
        "likes_count": post.likes_count
    }

三、依赖注入系统

3.1 依赖项基础

# app/api/deps.py
from fastapi import Depends, HTTPException, status, Request
from fastapi.security import OAuth2PasswordBearer
from jose import JWTError, jwt
from typing import Optional, Generator
from sqlalchemy.orm import Session

from app.db.session import SessionLocal
from app.core.config import settings
from app.models.user import User
from app.schemas.user import TokenData

# OAuth2 密码流(用于用户名密码登录)
oauth2_scheme = OAuth2PasswordBearer(
    tokenUrl=f"{settings.API_V1_STR}/auth/login"
)

# 数据库会话依赖
def get_db() -> Generator[Session, None, None]:
    """
    获取数据库会话

    每个请求创建一个新的会话,请求结束后自动关闭
    """
    db = SessionLocal()
    try:
        yield db
    finally:
        db.close()

# 当前用户依赖
async def get_current_user(
    db: Session = Depends(get_db),
    token: str = Depends(oauth2_scheme)
) -> User:
    """
    获取当前认证用户

    从 JWT token 中解析用户信息并验证
    """
    credentials_exception = HTTPException(
        status_code=status.HTTP_401_UNAUTHORIZED,
        detail="无法验证凭据",
        headers={"WWW-Authenticate": "Bearer"},
    )

    try:
        payload = jwt.decode(
            token, 
            settings.SECRET_KEY, 
            algorithms=[settings.ALGORITHM]
        )
        user_id: str = payload.get("sub")
        if user_id is None:
            raise credentials_exception
        token_data = TokenData(user_id=int(user_id))
    except JWTError:
        raise credentials_exception

    user = db.query(User).filter(User.id == token_data.user_id).first()
    if user is None:
        raise credentials_exception

    if not user.is_active:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="用户已被禁用"
        )

    return user

# 可选当前用户(允许未认证访问)
async def get_current_user_optional(
    db: Session = Depends(get_db),
    token: Optional[str] = Depends(oauth2_scheme)
) -> Optional[User]:
    """
    获取当前用户(可选)

    如果用户已认证则返回用户对象,否则返回 None
    """
    if not token:
        return None

    try:
        return await get_current_user(db, token)
    except HTTPException:
        return None

# 活跃用户依赖
async def get_current_active_user(
    current_user: User = Depends(get_current_user),
) -> User:
    """确保用户是活跃状态"""
    if not current_user.is_active:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="用户已被禁用"
        )
    return current_user

# 管理员权限依赖
async def get_current_admin_user(
    current_user: User = Depends(get_current_active_user),
) -> User:
    """确保用户是管理员"""
    if current_user.role != UserRole.ADMIN:
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="需要管理员权限"
        )
    return current_user

# 超级用户依赖
async def get_current_superuser(
    current_user: User = Depends(get_current_active_user),
) -> User:
    """确保用户是超级用户"""
    if not current_user.is_superuser:
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="需要超级用户权限"
        )
    return current_user

# 分页依赖
class PaginationParams:
    """分页参数"""
    def __init__(
        self,
        page: int = Query(1, ge=1, description="页码"),
        size: int = Query(20, ge=1, le=100, description="每页数量"),
    ):
        self.page = page
        self.size = size
        self.offset = (page - 1) * size

# 排序依赖
class SortParams:
    """排序参数"""
    def __init__(
        self,
        sort_by: str = Query("created_at", description="排序字段"),
        sort_order: str = Query("desc", regex="^(asc|desc)$", description="排序方向"),
    ):
        self.sort_by = sort_by
        self.sort_order = sort_order
        self.is_desc = sort_order == "desc"

# 请求限流依赖
from slowapi import Limiter, _rate_limit_exceeded_handler
from slowapi.util import get_remote_address
from slowapi.errors import RateLimitExceeded

limiter = Limiter(key_func=get_remote_address)

def get_rate_limiter():
    """获取限流器"""
    return limiter

# 缓存依赖
from app.core.cache import redis_client
import json

class CacheDependency:
    """缓存依赖"""

    def __init__(self, prefix: str = "", expire: int = 300):
        self.prefix = prefix
        self.expire = expire

    async def get(self, key: str) -> Optional[dict]:
        """从缓存获取"""
        cache_key = f"{self.prefix}:{key}"
        data = await redis_client.get(cache_key)
        if data:
            return json.loads(data)
        return None

    async def set(self, key: str, value: dict):
        """设置缓存"""
        cache_key = f"{self.prefix}:{key}"
        await redis_client.setex(
            cache_key,
            self.expire,
            json.dumps(value)
        )

    async def delete(self, key: str):
        """删除缓存"""
        cache_key = f"{self.prefix}:{key}"
        await redis_client.delete(cache_key)

    async def clear_pattern(self, pattern: str):
        """清除匹配模式的缓存"""
        keys = await redis_client.keys(f"{self.prefix}:{pattern}")
        if keys:
            await redis_client.delete(*keys)

# 使用示例:组合多个依赖
@router.get("/posts/")
async def list_posts_with_cache(
    pagination: PaginationParams = Depends(),
    sort: SortParams = Depends(),
    cache: CacheDependency = Depends(lambda: CacheDependency(prefix="posts", expire=60)),
    db: Session = Depends(get_db),
):
    """带缓存的文章列表"""

    # 生成缓存键
    cache_key = f"list:p{pagination.page}:s{pagination.size}:{sort.sort_by}:{sort.sort_order}"

    # 尝试从缓存获取
    cached = await cache.get(cache_key)
    if cached:
        return cached

    # 查询数据库
    query = db.query(Post)

    # 排序
    sort_column = getattr(Post, sort.sort_by)
    if sort.is_desc:
        query = query.order_by(sort_column.desc())
    else:
        query = query.order_by(sort_column.asc())

    total = query.count()
    posts = query.offset(pagination.offset).limit(pagination.size).all()

    result = {
        "posts": [post.to_dict() for post in posts],
        "total": total,
        "page": pagination.page,
        "size": pagination.size,
        "pages": (total + pagination.size - 1) // pagination.size
    }

    # 存入缓存
    await cache.set(cache_key, result)

    return result

3.2 高级依赖注入模式

# app/api/deps.py (续)
from functools import wraps
from typing import Callable, TypeVar, Any
import time

T = TypeVar("T")

# 类作为依赖
class CommonQueryParams:
    """通用查询参数类"""
    def __init__(
        self,
        q: Optional[str] = Query(None, description="搜索关键词"),
        skip: int = Query(0, ge=0, description="跳过的记录数"),
        limit: int = Query(100, ge=1, le=1000, description="返回的记录数"),
    ):
        self.q = q
        self.skip = skip
        self.limit = limit

# 使用类依赖
@router.get("/items/")
async def read_items(commons: CommonQueryParams = Depends()):
    """使用类依赖的示例"""
    return {"q": commons.q, "skip": commons.skip, "limit": commons.limit}

# 依赖项的可选参数
def check_permissions(required_permissions: List[str] = None):
    """
    权限检查依赖工厂

    用法:
    @app.get("/admin/", dependencies=[Depends(check_permissions(["admin:read"]))])
    """
    async def permission_checker(
        current_user: User = Depends(get_current_user)
    ):
        if required_permissions is None:
            return

        user_permissions = current_user.get_permissions()
        for perm in required_permissions:
            if perm not in user_permissions:
                raise HTTPException(
                    status_code=status.HTTP_403_FORBIDDEN,
                    detail=f"缺少权限: {perm}"
                )

    return permission_checker

# 子依赖
async def get_post_or_404(
    post_id: int,
    db: Session = Depends(get_db),
) -> Post:
    """获取文章或返回 404"""
    post = db.query(Post).filter(Post.id == post_id).first()
    if not post:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="文章不存在"
        )
    return post

async def check_post_ownership(
    post: Post = Depends(get_post_or_404),
    current_user: User = Depends(get_current_user),
) -> Post:
    """检查文章所有权"""
    if post.author_id != current_user.id and current_user.role != UserRole.ADMIN:
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="没有权限操作此文章"
        )
    return post

# 使用子依赖
@router.put("/posts/{post_id}")
async def update_owned_post(
    post: Post = Depends(check_post_ownership),
    post_update: PostUpdate = None,
    db: Session = Depends(get_db),
):
    """更新自己的文章"""
    # post 已经被验证所有权
    for field, value in post_update.model_dump(exclude_unset=True).items():
        setattr(post, field, value)

    db.commit()
    return post

# 依赖项装饰器
def require_permission(permission: str):
    """权限检查装饰器"""
    def decorator(func: Callable) -> Callable:
        @wraps(func)
        async def wrapper(
            *args,
            current_user: User = Depends(get_current_user),
            **kwargs
        ):
            if permission not in current_user.get_permissions():
                raise HTTPException(
                    status_code=status.HTTP_403_FORBIDDEN,
                    detail=f"缺少权限: {permission}"
                )
            return await func(*args, **kwargs)
        return wrapper
    return decorator

# 请求上下文依赖
class RequestContext:
    """请求上下文"""
    def __init__(self, request: Request):
        self.request = request
        self.start_time = time.time()
        self.user: Optional[User] = None

    @property
    def client_ip(self) -> str:
        """获取客户端 IP"""
        forwarded = self.request.headers.get("X-Forwarded-For")
        if forwarded:
            return forwarded.split(",")[0].strip()
        return self.request.client.host

    @property
    def user_agent(self) -> str:
        """获取 User-Agent"""
        return self.request.headers.get("User-Agent", "")

    @property
    def elapsed_time(self) -> float:
        """获取请求耗时"""
        return time.time() - self.start_time

async def get_request_context(request: Request) -> RequestContext:
    """获取请求上下文"""
    return RequestContext(request)

# 使用请求上下文
@router.get("/context")
async def test_context(
    ctx: RequestContext = Depends(get_request_context),
):
    """测试请求上下文"""
    return {
        "client_ip": ctx.client_ip,
        "user_agent": ctx.user_agent,
        "elapsed": ctx.elapsed_time
    }

# 数据库事务依赖
from contextlib import contextmanager

@contextmanager
def transaction(db: Session):
    """数据库事务上下文管理器"""
    try:
        yield
        db.commit()
    except Exception:
        db.rollback()
        raise

async def get_transaction_db(db: Session = Depends(get_db)):
    """获取事务数据库会话"""
    with transaction(db):
        yield db

# 使用事务
@router.post("/posts/with-transaction")
async def create_post_with_transaction(
    post: PostCreate,
    db: Session = Depends(get_transaction_db),
    current_user: User = Depends(get_current_user),
):
    """使用事务创建文章"""
    # 所有操作在同一事务中
    db_post = Post(**post.model_dump(), author_id=current_user.id)
    db.add(db_post)

    # 创建关联的标签
    for tag_name in post.tag_names:
        tag = Tag(name=tag_name)
        db.add(tag)
        db_post.tags.append(tag)

    # 事务会在函数结束时自动提交或回滚
    return db_post

四、数据库集成

4.1 SQLAlchemy 异步配置

# app/db/session.py
from sqlalchemy.ext.asyncio import (
    AsyncSession, 
    create_async_engine, 
    async_sessionmaker,
    AsyncEngine
)
from sqlalchemy.orm import declarative_base
from app.core.config import settings

# 创建异步引擎
engine: AsyncEngine = create_async_engine(
    settings.DATABASE_URL.replace("postgresql://", "postgresql+asyncpg://"),
    echo=settings.DB_ECHO,
    pool_size=settings.DB_POOL_SIZE,
    max_overflow=settings.DB_MAX_OVERFLOW,
    pool_pre_ping=True,
)

# 创建异步会话工厂
AsyncSessionLocal = async_sessionmaker(
    engine,
    class_=AsyncSession,
    expire_on_commit=False,
    autocommit=False,
    autoflush=False,
)

# 声明基类
Base = declarative_base()

async def get_async_db() -> AsyncGenerator[AsyncSession, None]:
    """
    获取异步数据库会话

    用法:
    @app.get("/")
    async def root(db: AsyncSession = Depends(get_async_db)):
        result = await db.execute(select(User))
        return result.scalars().all()
    """
    async with AsyncSessionLocal() as session:
        try:
            yield session
            await session.commit()
        except Exception:
            await session.rollback()
            raise
        finally:
            await session.close()

# 异步数据库工具类
class AsyncDatabase:
    """异步数据库操作工具"""

    @staticmethod
    async def create_tables():
        """创建所有表"""
        async with engine.begin() as conn:
            await conn.run_sync(Base.metadata.create_all)

    @staticmethod
    async def drop_tables():
        """删除所有表"""
        async with engine.begin() as conn:
            await conn.run_sync(Base.metadata.drop_all)

    @staticmethod
    async def check_connection() -> bool:
        """检查数据库连接"""
        try:
            async with AsyncSessionLocal() as session:
                await session.execute("SELECT 1")
            return True
        except Exception:
            return False

# app/models/user.py - 异步兼容的模型
from sqlalchemy import (
    Column, Integer, String, Boolean, DateTime, 
    Enum as SQLEnum, Text, ForeignKey, Table
)
from sqlalchemy.orm import relationship, selectinload
from sqlalchemy.sql import func
from app.db.session import Base
import enum

class UserRole(str, enum.Enum):
    ADMIN = "admin"
    MODERATOR = "moderator"
    USER = "user"

class User(Base):
    __tablename__ = "users"

    id = Column(Integer, primary_key=True, index=True)
    username = Column(String(50), unique=True, index=True, nullable=False)
    email = Column(String(255), unique=True, index=True, nullable=False)
    hashed_password = Column(String(255), nullable=False)
    full_name = Column(String(100))
    bio = Column(Text)
    avatar = Column(String(500))

    role = Column(SQLEnum(UserRole), default=UserRole.USER)
    is_active = Column(Boolean, default=True)
    is_superuser = Column(Boolean, default=False)
    email_verified = Column(Boolean, default=False)

    created_at = Column(DateTime(timezone=True), server_default=func.now())
    updated_at = Column(DateTime(timezone=True), onupdate=func.now())
    last_login = Column(DateTime(timezone=True))

    # 关系
    posts = relationship("Post", back_populates="author", lazy="dynamic")
    comments = relationship("Comment", back_populates="author", lazy="dynamic")

    def to_dict(self):
        """转换为字典"""
        return {
            "id": self.id,
            "username": self.username,
            "email": self.email,
            "full_name": self.full_name,
            "bio": self.bio,
            "avatar": self.avatar,
            "role": self.role.value,
            "is_active": self.is_active,
            "is_superuser": self.is_superuser,
            "email_verified": self.email_verified,
            "created_at": self.created_at.isoformat() if self.created_at else None,
            "last_login": self.last_login.isoformat() if self.last_login else None,
        }

# CRUD 操作基类
from typing import TypeVar, Generic, Type, Optional, List, Union, Dict, Any
from pydantic import BaseModel
from sqlalchemy import select, func, update, delete
from sqlalchemy.ext.asyncio import AsyncSession

ModelType = TypeVar("ModelType", bound=Base)
CreateSchemaType = TypeVar("CreateSchemaType", bound=BaseModel)
UpdateSchemaType = TypeVar("UpdateSchemaType", bound=BaseModel)

class CRUDBase(Generic[ModelType, CreateSchemaType, UpdateSchemaType]):
    """CRUD 操作基类"""

    def __init__(self, model: Type[ModelType]):
        self.model = model

    async def get(
        self, 
        db: AsyncSession, 
        id: int
    ) -> Optional[ModelType]:
        """根据 ID 获取记录"""
        result = await db.execute(
            select(self.model).where(self.model.id == id)
        )
        return result.scalar_one_or_none()

    async def get_multi(
        self,
        db: AsyncSession,
        *,
        skip: int = 0,
        limit: int = 100,
        filters: Optional[Dict[str, Any]] = None,
        order_by: Optional[str] = None,
    ) -> List[ModelType]:
        """获取多条记录"""
        query = select(self.model)

        if filters:
            for field, value in filters.items():
                if hasattr(self.model, field):
                    query = query.where(getattr(self.model, field) == value)

        if order_by:
            if order_by.startswith("-"):
                query = query.order_by(getattr(self.model, order_by[1:]).desc())
            else:
                query = query.order_by(getattr(self.model, order_by).asc())

        query = query.offset(skip).limit(limit)
        result = await db.execute(query)
        return result.scalars().all()

    async def create(
        self, 
        db: AsyncSession, 
        *, 
        obj_in: CreateSchemaType
    ) -> ModelType:
        """创建记录"""
        obj_in_data = obj_in.model_dump()
        db_obj = self.model(**obj_in_data)
        db.add(db_obj)
        await db.flush()
        await db.refresh(db_obj)
        return db_obj

    async def update(
        self,
        db: AsyncSession,
        *,
        db_obj: ModelType,
        obj_in: Union[UpdateSchemaType, Dict[str, Any]]
    ) -> ModelType:
        """更新记录"""
        if isinstance(obj_in, dict):
            update_data = obj_in
        else:
            update_data = obj_in.model_dump(exclude_unset=True)

        for field, value in update_data.items():
            if hasattr(db_obj, field):
                setattr(db_obj, field, value)

        db.add(db_obj)
        await db.flush()
        await db.refresh(db_obj)
        return db_obj

    async def delete(
        self, 
        db: AsyncSession, 
        *, 
        id: int
    ) -> Optional[ModelType]:
        """删除记录"""
        obj = await self.get(db, id)
        if obj:
            await db.delete(obj)
            await db.flush()
        return obj

    async def count(
        self, 
        db: AsyncSession,
        filters: Optional[Dict[str, Any]] = None
    ) -> int:
        """统计记录数"""
        query = select(func.count()).select_from(self.model)

        if filters:
            for field, value in filters.items():
                if hasattr(self.model, field):
                    query = query.where(getattr(self.model, field) == value)

        result = await db.execute(query)
        return result.scalar_one()

# 用户 CRUD
class CRUDUser(CRUDBase[User, UserCreate, UserUpdate]):
    """用户 CRUD 操作"""

    async def get_by_email(
        self, 
        db: AsyncSession, 
        *, 
        email: str
    ) -> Optional[User]:
        """根据邮箱获取用户"""
        result = await db.execute(
            select(User).where(User.email == email)
        )
        return result.scalar_one_or_none()

    async def get_by_username(
        self, 
        db: AsyncSession, 
        *, 
        username: str
    ) -> Optional[User]:
        """根据用户名获取用户"""
        result = await db.execute(
            select(User).where(User.username == username)
        )
        return result.scalar_one_or_none()

    async def create(
        self, 
        db: AsyncSession, 
        *, 
        obj_in: UserCreate
    ) -> User:
        """创建用户(包含密码哈希)"""
        from app.core.security import get_password_hash

        obj_in_data = obj_in.model_dump(exclude={"password", "password_confirm"})
        obj_in_data["hashed_password"] = get_password_hash(obj_in.password)

        db_obj = User(**obj_in_data)
        db.add(db_obj)
        await db.flush()
        await db.refresh(db_obj)
        return db_obj

    async def authenticate(
        self, 
        db: AsyncSession, 
        *, 
        email: str, 
        password: str
    ) -> Optional[User]:
        """认证用户"""
        from app.core.security import verify_password

        user = await self.get_by_email(db, email=email)
        if not user:
            return None
        if not verify_password(password, user.hashed_password):
            return None
        return user

    async def update_last_login(
        self, 
        db: AsyncSession, 
        *, 
        user: User
    ) -> User:
        """更新最后登录时间"""
        from datetime import datetime

        user.last_login = datetime.utcnow()
        db.add(user)
        await db.flush()
        return user

user_crud = CRUDUser(User)

五、认证与授权

5.1 JWT 认证实现

# app/core/security.py
from datetime import datetime, timedelta
from typing import Optional, Union, Dict, Any
from jose import JWTError, jwt
from passlib.context import CryptContext
from app.core.config import settings

# 密码上下文
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")

def verify_password(plain_password: str, hashed_password: str) -> bool:
    """验证密码"""
    return pwd_context.verify(plain_password, hashed_password)

def get_password_hash(password: str) -> str:
    """生成密码哈希"""
    return pwd_context.hash(password)

def create_access_token(
    data: Dict[str, Any],
    expires_delta: Optional[timedelta] = None
) -> str:
    """创建 JWT 访问令牌"""
    to_encode = data.copy()

    if expires_delta:
        expire = datetime.utcnow() + expires_delta
    else:
        expire = datetime.utcnow() + timedelta(
            minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES
        )

    to_encode.update({"exp": expire, "type": "access"})
    encoded_jwt = jwt.encode(
        to_encode, 
        settings.SECRET_KEY, 
        algorithm=settings.ALGORITHM
    )
    return encoded_jwt

def create_refresh_token(data: Dict[str, Any]) -> str:
    """创建 JWT 刷新令牌"""
    to_encode = data.copy()
    expire = datetime.utcnow() + timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS)

    to_encode.update({"exp": expire, "type": "refresh"})
    encoded_jwt = jwt.encode(
        to_encode, 
        settings.SECRET_KEY, 
        algorithm=settings.ALGORITHM
    )
    return encoded_jwt

def decode_token(token: str) -> Optional[Dict[str, Any]]:
    """解码 JWT 令牌"""
    try:
        payload = jwt.decode(
            token, 
            settings.SECRET_KEY, 
            algorithms=[settings.ALGORITHM]
        )
        return payload
    except JWTError:
        return None

# app/api/v1/endpoints/auth.py
from fastapi import APIRouter, Depends, HTTPException, status, BackgroundTasks
from fastapi.security import OAuth2PasswordRequestForm
from sqlalchemy.ext.asyncio import AsyncSession
from datetime import timedelta

from app.db.session import get_async_db
from app.schemas.user import User, UserCreate, Token
from app.crud.user import user_crud
from app.core.security import (
    verify_password, 
    create_access_token, 
    create_refresh_token,
    decode_token
)
from app.core.config import settings
from app.api.deps import get_current_user
from app.utils.email import send_verification_email

router = APIRouter(prefix="/auth", tags=["authentication"])

@router.post("/register", response_model=User, status_code=status.HTTP_201_CREATED)
async def register(
    *,
    db: AsyncSession = Depends(get_async_db),
    user_in: UserCreate,
    background_tasks: BackgroundTasks,
):
    """
    用户注册

    - 创建新用户
    - 发送验证邮件
    """
    # 检查邮箱是否已存在
    existing_user = await user_crud.get_by_email(db, email=user_in.email)
    if existing_user:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="该邮箱已被注册"
        )

    # 检查用户名是否已存在
    existing_user = await user_crud.get_by_username(db, username=user_in.username)
    if existing_user:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="该用户名已被使用"
        )

    # 创建用户
    user = await user_crud.create(db, obj_in=user_in)

    # 发送验证邮件(后台任务)
    background_tasks.add_task(
        send_verification_email,
        email=user.email,
        username=user.username,
        user_id=user.id
    )

    return user

@router.post("/login", response_model=Token)
async def login(
    db: AsyncSession = Depends(get_async_db),
    form_data: OAuth2PasswordRequestForm = Depends()
):
    """
    用户登录

    使用 OAuth2 密码流获取访问令牌
    """
    # 认证用户
    user = await user_crud.authenticate(
        db, 
        email=form_data.username,  # OAuth2 表单使用 username 字段
        password=form_data.password
    )

    if not user:
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="邮箱或密码错误",
            headers={"WWW-Authenticate": "Bearer"},
        )

    if not user.is_active:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="用户已被禁用"
        )

    # 更新最后登录时间
    await user_crud.update_last_login(db, user=user)

    # 创建令牌
    access_token = create_access_token(
        data={"sub": str(user.id), "username": user.username, "role": user.role.value}
    )
    refresh_token = create_refresh_token(
        data={"sub": str(user.id)}
    )

    return {
        "access_token": access_token,
        "refresh_token": refresh_token,
        "token_type": "bearer",
        "expires_in": settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60
    }

@router.post("/refresh", response_model=Token)
async def refresh_token(
    refresh_token: str,
    db: AsyncSession = Depends(get_async_db),
):
    """
    刷新访问令牌

    使用刷新令牌获取新的访问令牌
    """
    # 解码刷新令牌
    payload = decode_token(refresh_token)
    if not payload or payload.get("type") != "refresh":
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="无效的刷新令牌"
        )

    user_id = payload.get("sub")
    if not user_id:
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="无效的刷新令牌"
        )

    # 验证用户
    user = await user_crud.get(db, id=int(user_id))
    if not user or not user.is_active:
        raise HTTPException(
            status_code=status.HTTP_401_UNAUTHORIZED,
            detail="用户不存在或已被禁用"
        )

    # 创建新的访问令牌
    access_token = create_access_token(
        data={"sub": str(user.id), "username": user.username, "role": user.role.value}
    )

    return {
        "access_token": access_token,
        "refresh_token": refresh_token,  # 返回原刷新令牌
        "token_type": "bearer",
        "expires_in": settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60
    }

@router.post("/logout")
async def logout(
    current_user: User = Depends(get_current_user),
):
    """
    用户登出

    注意:JWT 是无状态的,登出需要在客户端删除令牌
    这里可以添加将令牌加入黑名单的逻辑
    """
    # 可选:将令牌加入 Redis 黑名单
    # await add_to_blacklist(token)

    return {"message": "登出成功"}

@router.get("/me", response_model=User)
async def get_current_user_info(
    current_user: User = Depends(get_current_user),
):
    """获取当前用户信息"""
    return current_user

@router.post("/verify-email/{token}")
async def verify_email(
    token: str,
    db: AsyncSession = Depends(get_async_db),
):
    """验证邮箱"""
    payload = decode_token(token)
    if not payload or payload.get("type") != "email_verification":
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="无效的验证令牌"
        )

    user_id = payload.get("sub")
    user = await user_crud.get(db, id=int(user_id))

    if not user:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="用户不存在"
        )

    if user.email_verified:
        return {"message": "邮箱已验证"}

    user.email_verified = True
    await db.commit()

    return {"message": "邮箱验证成功"}

@router.post("/forgot-password")
async def forgot_password(
    email: str,
    db: AsyncSession = Depends(get_async_db),
    background_tasks: BackgroundTasks = None,
):
    """
    忘记密码

    发送重置密码邮件
    """
    user = await user_crud.get_by_email(db, email=email)

    if user:
        # 创建重置密码令牌
        reset_token = create_access_token(
            data={"sub": str(user.id), "type": "password_reset"},
            expires_delta=timedelta(hours=1)
        )

        # 发送重置密码邮件
        if background_tasks:
            background_tasks.add_task(
                send_password_reset_email,
                email=user.email,
                username=user.username,
                token=reset_token
            )

    # 无论用户是否存在,都返回相同消息(防止邮箱枚举攻击)
    return {"message": "如果该邮箱已注册,您将收到重置密码邮件"}

@router.post("/reset-password")
async def reset_password(
    token: str,
    new_password: str,
    db: AsyncSession = Depends(get_async_db),
):
    """重置密码"""
    payload = decode_token(token)
    if not payload or payload.get("type") != "password_reset":
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="无效的重置令牌"
        )

    user_id = payload.get("sub")
    user = await user_crud.get(db, id=int(user_id))

    if not user:
        raise HTTPException(
            status_code=status.HTTP_404_NOT_FOUND,
            detail="用户不存在"
        )

    # 更新密码
    from app.core.security import get_password_hash
    user.hashed_password = get_password_hash(new_password)
    await db.commit()

    return {"message": "密码重置成功"}

# OAuth2 第三方登录
from fastapi_sso import GoogleSSO, GitHubSSO

google_sso = GoogleSSO(
    client_id=settings.GOOGLE_CLIENT_ID,
    client_secret=settings.GOOGLE_CLIENT_SECRET,
    redirect_uri=f"{settings.SERVER_HOST}/api/v1/auth/google/callback",
)

github_sso = GitHubSSO(
    client_id=settings.GITHUB_CLIENT_ID,
    client_secret=settings.GITHUB_CLIENT_SECRET,
    redirect_uri=f"{settings.SERVER_HOST}/api/v1/auth/github/callback",
)

@router.get("/google/login")
async def google_login():
    """Google 登录入口"""
    return await google_sso.get_login_redirect()

@router.get("/google/callback")
async def google_callback(
    request: Request,
    db: AsyncSession = Depends(get_async_db),
):
    """Google 登录回调"""
    user_info = await google_sso.verify_and_process(request)

    # 查找或创建用户
    user = await user_crud.get_by_email(db, email=user_info.email)

    if not user:
        # 创建新用户
        user_in = UserCreate(
            username=user_info.email.split("@")[0],
            email=user_info.email,
            password=None,  # OAuth 用户不需要密码
            full_name=user_info.display_name
        )
        user = await user_crud.create_oauth_user(db, obj_in=user_in)

    # 创建访问令牌
    access_token = create_access_token(
        data={"sub": str(user.id), "username": user.username, "role": user.role.value}
    )

    return {"access_token": access_token, "token_type": "bearer"}

六、中间件与事件处理

6.1 自定义中间件

# app/middleware.py
from fastapi import FastAPI, Request, Response
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.trustedhost import TrustedHostMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
from starlette.middleware.sessions import SessionMiddleware
from starlette.types import ASGIApp
import time
import logging
import uuid
from typing import Callable, Dict, Any

logger = logging.getLogger(__name__)

class RequestIDMiddleware(BaseHTTPMiddleware):
    """请求 ID 中间件"""

    async def dispatch(
        self, 
        request: Request, 
        call_next: RequestResponseEndpoint
    ) -> Response:
        request_id = str(uuid.uuid4())
        request.state.request_id = request_id

        response = await call_next(request)
        response.headers["X-Request-ID"] = request_id

        return response

class LoggingMiddleware(BaseHTTPMiddleware):
    """日志中间件"""

    async def dispatch(
        self, 
        request: Request, 
        call_next: RequestResponseEndpoint
    ) -> Response:
        start_time = time.time()

        # 记录请求
        logger.info(
            f"Request: {request.method} {request.url.path}",
            extra={
                "method": request.method,
                "path": request.url.path,
                "client": request.client.host if request.client else None,
                "request_id": getattr(request.state, "request_id", None)
            }
        )

        response = await call_next(request)

        # 记录响应
        duration = time.time() - start_time
        logger.info(
            f"Response: {response.status_code} ({duration:.3f}s)",
            extra={
                "status_code": response.status_code,
                "duration": duration,
                "request_id": getattr(request.state, "request_id", None)
            }
        )

        return response

class RateLimitMiddleware(BaseHTTPMiddleware):
    """限流中间件"""

    def __init__(self, app: ASGIApp, redis_client, limit: int = 100, window: int = 60):
        super().__init__(app)
        self.redis = redis_client
        self.limit = limit
        self.window = window

    async def dispatch(
        self, 
        request: Request, 
        call_next: RequestResponseEndpoint
    ) -> Response:
        # 获取客户端标识
        client_id = request.headers.get(
            "X-Forwarded-For", 
            request.client.host if request.client else "unknown"
        )

        # 构建 Redis 键
        key = f"rate_limit:{client_id}:{request.url.path}"

        # 获取当前计数
        current = await self.redis.get(key)

        if current and int(current) >= self.limit:
            return Response(
                content='{"detail": "请求过于频繁,请稍后再试"}',
                status_code=429,
                media_type="application/json",
                headers={"Retry-After": str(self.window)}
            )

        # 增加计数
        pipe = self.redis.pipeline()
        pipe.incr(key)
        pipe.expire(key, self.window)
        await pipe.execute()

        return await call_next(request)

class CacheMiddleware(BaseHTTPMiddleware):
    """缓存中间件"""

    def __init__(self, app: ASGIApp, redis_client, default_ttl: int = 60):
        super().__init__(app)
        self.redis = redis_client
        self.default_ttl = default_ttl

    async def dispatch(
        self, 
        request: Request, 
        call_next: RequestResponseEndpoint
    ) -> Response:
        # 只缓存 GET 请求
        if request.method != "GET":
            return await call_next(request)

        # 检查是否需要绕过缓存
        if request.headers.get("Cache-Control") == "no-cache":
            return await call_next(request)

        # 构建缓存键
        cache_key = f"cache:{request.url.path}:{request.url.query}"

        # 尝试从缓存获取
        cached = await self.redis.get(cache_key)
        if cached:
            return Response(
                content=cached,
                media_type="application/json",
                headers={"X-Cache": "HIT"}
            )

        # 执行请求
        response = await call_next(request)

        # 只缓存成功的响应
        if response.status_code == 200:
            # 获取缓存 TTL
            ttl = self.default_ttl
            cache_control = response.headers.get("Cache-Control")
            if cache_control and "max-age=" in cache_control:
                try:
                    ttl = int(cache_control.split("max-age=")[1].split(",")[0])
                except (ValueError, IndexError):
                    pass

            # 读取响应体
            body = b""
            async for chunk in response.body_iterator:
                body += chunk

            # 存入缓存
            await self.redis.setex(cache_key, ttl, body)

            # 重建响应
            return Response(
                content=body,
                status_code=response.status_code,
                headers=dict(response.headers),
                media_type=response.media_type
            )

        return response

class SecurityHeadersMiddleware(BaseHTTPMiddleware):
    """安全头中间件"""

    async def dispatch(
        self, 
        request: Request, 
        call_next: RequestResponseEndpoint
    ) -> Response:
        response = await call_next(request)

        # 添加安全头
        response.headers["X-Content-Type-Options"] = "nosniff"
        response.headers["X-Frame-Options"] = "DENY"
        response.headers["X-XSS-Protection"] = "1; mode=block"
        response.headers["Strict-Transport-Security"] = "max-age=31536000; includeSubDomains"
        response.headers["Content-Security-Policy"] = "default-src 'self'"
        response.headers["Referrer-Policy"] = "strict-origin-when-cross-origin"
        response.headers["Permissions-Policy"] = "geolocation=(), microphone=(), camera=()"

        return response

# 注册中间件
def setup_middlewares(app: FastAPI):
    """配置所有中间件"""

    # CORS
    app.add_middleware(
        CORSMiddleware,
        allow_origins=settings.CORS_ORIGINS,
        allow_credentials=True,
        allow_methods=["*"],
        allow_headers=["*"],
    )

    # 信任的主机
    app.add_middleware(
        TrustedHostMiddleware,
        allowed_hosts=settings.ALLOWED_HOSTS,
    )

    # GZip 压缩
    app.add_middleware(GZipMiddleware, minimum_size=1000)

    # Session
    app.add_middleware(
        SessionMiddleware,
        secret_key=settings.SECRET_KEY,
        session_cookie="session",
        max_age=86400,
    )

    # 自定义中间件
    app.add_middleware(RequestIDMiddleware)
    app.add_middleware(LoggingMiddleware)
    app.add_middleware(SecurityHeadersMiddleware)

    # 限流中间件(需要 Redis)
    if settings.REDIS_URL:
        import redis.asyncio as redis
        redis_client = redis.from_url(settings.REDIS_URL)
        app.add_middleware(
            RateLimitMiddleware,
            redis_client=redis_client,
            limit=settings.RATE_LIMIT_PER_MINUTE,
            window=60
        )
        app.add_middleware(
            CacheMiddleware,
            redis_client=redis_client,
            default_ttl=60
        )

6.2 应用生命周期事件

# app/events.py
from fastapi import FastAPI
from contextlib import asynccontextmanager
import asyncio
import logging
from typing import AsyncGenerator

logger = logging.getLogger(__name__)

@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
    """
    应用生命周期管理

    启动时执行 yield 之前的代码
    关闭时执行 yield 之后的代码
    """
    # 启动时
    logger.info("应用启动中...")

    # 初始化数据库连接池
    from app.db.session import engine
    await engine.connect()
    logger.info("数据库连接池已初始化")

    # 初始化 Redis 连接
    if settings.REDIS_URL:
        import redis.asyncio as redis
        app.state.redis = redis.from_url(
            settings.REDIS_URL,
            encoding="utf-8",
            decode_responses=True
        )
        await app.state.redis.ping()
        logger.info("Redis 连接成功")

    # 初始化后台任务
    cleanup_task = asyncio.create_task(cleanup_expired_tokens())

    # 预热缓存
    if hasattr(app.state, 'redis'):
        await warmup_cache(app)

    logger.info("应用启动完成")

    yield  # 应用运行中

    # 关闭时
    logger.info("应用正在关闭...")

    # 取消后台任务
    cleanup_task.cancel()
    try:
        await cleanup_task
    except asyncio.CancelledError:
        pass

    # 关闭 Redis 连接
    if hasattr(app.state, 'redis'):
        await app.state.redis.close()
        logger.info("Redis 连接已关闭")

    # 关闭数据库连接池
    await engine.dispose()
    logger.info("数据库连接池已关闭")

    logger.info("应用已关闭")

async def cleanup_expired_tokens():
    """定期清理过期的令牌"""
    while True:
        try:
            await asyncio.sleep(3600)  # 每小时执行一次

            # 从数据库清理过期的刷新令牌
            from datetime import datetime
            from app.db.session import AsyncSessionLocal
            from sqlalchemy import delete

            async with AsyncSessionLocal() as db:
                await db.execute(
                    delete(RefreshToken).where(
                        RefreshToken.expires_at < datetime.utcnow()
                    )
                )
                await db.commit()

            logger.info("已清理过期的刷新令牌")

        except asyncio.CancelledError:
            break
        except Exception as e:
            logger.error(f"清理过期令牌失败: {e}")

async def warmup_cache(app: FastAPI):
    """预热缓存"""
    # 预加载热门数据到缓存
    from app.crud.post import post_crud
    from app.db.session import AsyncSessionLocal

    async with AsyncSessionLocal() as db:
        # 获取热门文章
        popular_posts = await post_crud.get_popular(db, limit=20)

        # 存入缓存
        for post in popular_posts:
            cache_key = f"post:{post.id}"
            await app.state.redis.setex(
                cache_key,
                3600,  # 1小时
                post.json()
            )

    logger.info(f"缓存预热完成,加载了 {len(popular_posts)} 篇文章")

# 在 main.py 中使用
app = FastAPI(
    title=settings.PROJECT_NAME,
    version=settings.VERSION,
    lifespan=lifespan,
)

# 也可以使用装饰器方式
@app.on_event("startup")
async def startup_event():
    """启动事件"""
    logger.info("执行启动事件")
    # 初始化操作

@app.on_event("shutdown")
async def shutdown_event():
    """关闭事件"""
    logger.info("执行关闭事件")
    # 清理操作

七、WebSocket 支持

7.1 WebSocket 实现实时通信

# app/api/v1/websocket.py
from fastapi import WebSocket, WebSocketDisconnect, Depends
from fastapi.routing import APIRouter
from typing import Dict, Set, Optional
import json
import asyncio
from datetime import datetime

router = APIRouter(prefix="/ws", tags=["websocket"])

# 连接管理器
class ConnectionManager:
    """WebSocket 连接管理器"""

    def __init__(self):
        # 活跃连接
        self.active_connections: Dict[str, Set[WebSocket]] = {}
        # 用户到连接的映射
        self.user_connections: Dict[int, Set[WebSocket]] = {}

    async def connect(
        self, 
        websocket: WebSocket, 
        room: str = "global",
        user_id: Optional[int] = None
    ):
        """接受连接"""
        await websocket.accept()

        # 添加到房间
        if room not in self.active_connections:
            self.active_connections[room] = set()
        self.active_connections[room].add(websocket)

        # 添加到用户映射
        if user_id:
            if user_id not in self.user_connections:
                self.user_connections[user_id] = set()
            self.user_connections[user_id].add(websocket)

    def disconnect(
        self, 
        websocket: WebSocket, 
        room: str = "global",
        user_id: Optional[int] = None
    ):
        """断开连接"""
        # 从房间移除
        if room in self.active_connections:
            self.active_connections[room].discard(websocket)
            if not self.active_connections[room]:
                del self.active_connections[room]

        # 从用户映射移除
        if user_id and user_id in self.user_connections:
            self.user_connections[user_id].discard(websocket)
            if not self.user_connections[user_id]:
                del self.user_connections[user_id]

    async def send_personal_message(
        self, 
        message: dict, 
        websocket: WebSocket
    ):
        """发送个人消息"""
        await websocket.send_json(message)

    async def broadcast(
        self, 
        message: dict, 
        room: str = "global",
        exclude: Optional[WebSocket] = None
    ):
        """广播消息到房间"""
        if room in self.active_connections:
            for connection in self.active_connections[room]:
                if connection != exclude:
                    await connection.send_json(message)

    async def send_to_user(
        self, 
        message: dict, 
        user_id: int
    ):
        """发送消息给指定用户"""
        if user_id in self.user_connections:
            for connection in self.user_connections[user_id]:
                await connection.send_json(message)

manager = ConnectionManager()

@router.websocket("/{room}")
async def websocket_endpoint(
    websocket: WebSocket,
    room: str,
    token: Optional[str] = None,
):
    """WebSocket 连接端点"""
    user_id = None

    # 验证 token(可选)
    if token:
        from app.core.security import decode_token
        payload = decode_token(token)
        if payload:
            user_id = int(payload.get("sub", 0))

    await manager.connect(websocket, room, user_id)

    try:
        # 发送欢迎消息
        await manager.send_personal_message({
            "type": "system",
            "message": f"已连接到房间: {room}",
            "timestamp": datetime.utcnow().isoformat()
        }, websocket)

        # 通知其他用户
        await manager.broadcast({
            "type": "system",
            "message": f"用户 {user_id or '匿名'} 加入了房间",
            "timestamp": datetime.utcnow().isoformat(),
            "user_id": user_id
        }, room, exclude=websocket)

        while True:
            # 接收消息
            data = await websocket.receive_text()

            try:
                message_data = json.loads(data)
            except json.JSONDecodeError:
                await manager.send_personal_message({
                    "type": "error",
                    "message": "无效的 JSON 格式"
                }, websocket)
                continue

            # 处理不同类型的消息
            msg_type = message_data.get("type")

            if msg_type == "chat":
                # 聊天消息
                broadcast_message = {
                    "type": "chat",
                    "content": message_data.get("content", ""),
                    "user_id": user_id,
                    "username": message_data.get("username", "匿名"),
                    "timestamp": datetime.utcnow().isoformat()
                }
                await manager.broadcast(broadcast_message, room)

            elif msg_type == "private":
                # 私聊消息
                target_user_id = message_data.get("target_user_id")
                if target_user_id:
                    await manager.send_to_user({
                        "type": "private",
                        "content": message_data.get("content", ""),
                        "from_user_id": user_id,
                        "timestamp": datetime.utcnow().isoformat()
                    }, target_user_id)

            elif msg_type == "typing":
                # 正在输入状态
                await manager.broadcast({
                    "type": "typing",
                    "user_id": user_id,
                    "is_typing": message_data.get("is_typing", False)
                }, room, exclude=websocket)

            elif msg_type == "ping":
                # 心跳
                await manager.send_personal_message({
                    "type": "pong",
                    "timestamp": datetime.utcnow().isoformat()
                }, websocket)

    except WebSocketDisconnect:
        manager.disconnect(websocket, room, user_id)

        # 通知其他用户
        await manager.broadcast({
            "type": "system",
            "message": f"用户 {user_id or '匿名'} 离开了房间",
            "timestamp": datetime.utcnow().isoformat(),
            "user_id": user_id
        }, room)

# 实时通知 WebSocket
@router.websocket("/notifications")
async def notification_websocket(
    websocket: WebSocket,
    token: str,
):
    """实时通知 WebSocket"""
    # 验证 token
    from app.core.security import decode_token
    payload = decode_token(token)

    if not payload:
        await websocket.close(code=4001, reason="无效的认证令牌")
        return

    user_id = int(payload.get("sub", 0))

    await manager.connect(websocket, f"notifications_{user_id}", user_id)

    try:
        # 发送未读通知
        from app.services.notification import get_unread_notifications
        unread = await get_unread_notifications(user_id)

        await manager.send_personal_message({
            "type": "unread",
            "notifications": unread,
            "count": len(unread)
        }, websocket)

        while True:
            data = await websocket.receive_text()
            message = json.loads(data)

            if message.get("type") == "mark_read":
                # 标记通知为已读
                notification_id = message.get("notification_id")
                if notification_id:
                    from app.services.notification import mark_as_read
                    await mark_as_read(notification_id, user_id)

    except WebSocketDisconnect:
        manager.disconnect(websocket, f"notifications_{user_id}", user_id)

八、后台任务与消息队列

8.1 BackgroundTasks 与 Celery 集成

# app/tasks/background.py
from fastapi import BackgroundTasks
from typing import List, Optional
import asyncio
from datetime import datetime

async def send_email_background(
    email_to: str,
    subject: str,
    body: str,
    template_name: Optional[str] = None,
    template_context: Optional[dict] = None,
):
    """
    后台发送邮件

    用法:
    background_tasks.add_task(
        send_email_background,
        email_to="user@example.com",
        subject="Welcome",
        body="Welcome to our platform!"
    )
    """
    # 模拟发送邮件
    await asyncio.sleep(2)

    # 实际发送邮件的代码
    from app.utils.email import send_email
    await send_email(
        email_to=email_to,
        subject=subject,
        body=body,
        template_name=template_name,
        template_context=template_context
    )

    # 记录日志
    from app.core.logging import logger
    logger.info(f"邮件已发送到 {email_to}: {subject}")

async def generate_thumbnail_background(image_path: str, sizes: List[tuple]):
    """后台生成缩略图"""
    from PIL import Image

    for width, height in sizes:
        # 生成缩略图
        img = Image.open(image_path)
        img.thumbnail((width, height))

        # 保存缩略图
        thumbnail_path = f"{image_path.rsplit('.', 1)[0]}_{width}x{height}.{image_path.rsplit('.', 1)[1]}"
        img.save(thumbnail_path)

        await asyncio.sleep(0.1)  # 让出控制权

    from app.core.logging import logger
    logger.info(f"缩略图生成完成: {image_path}")

async def process_uploaded_file_background(file_path: str, file_type: str):
    """后台处理上传的文件"""

    if file_type.startswith("image/"):
        # 图片处理
        await generate_thumbnail_background(
            file_path,
            [(100, 100), (300, 300), (800, 600)]
        )

    elif file_type == "application/pdf":
        # PDF 处理
        await extract_pdf_text_background(file_path)

    elif file_type.startswith("video/"):
        # 视频处理
        await transcode_video_background(file_path)

    # 更新文件状态
    from app.services.file import update_file_status
    await update_file_status(file_path, "processed")

# app/tasks/celery_app.py
from celery import Celery
from celery.schedules import crontab
from app.core.config import settings

# 创建 Celery 应用
celery_app = Celery(
    "app",
    broker=settings.REDIS_URL,
    backend=settings.REDIS_URL,
    include=["app.tasks.email", "app.tasks.report", "app.tasks.cleanup"]
)

# 配置
celery_app.conf.update(
    task_serializer="json",
    accept_content=["json"],
    result_serializer="json",
    timezone="Asia/Shanghai",
    enable_utc=True,

    # 任务路由
    task_routes={
        "app.tasks.email.*": {"queue": "email"},
        "app.tasks.report.*": {"queue": "report"},
        "app.tasks.cleanup.*": {"queue": "cleanup"},
    },

    # 定时任务
    beat_schedule={
        "cleanup-expired-tokens": {
            "task": "app.tasks.cleanup.cleanup_expired_tokens",
            "schedule": crontab(hour=3, minute=0),  # 每天凌晨3点
        },
        "send-weekly-report": {
            "task": "app.tasks.report.send_weekly_report",
            "schedule": crontab(day_of_week=1, hour=9, minute=0),  # 每周一早上9点
        },
        "update-popular-posts-cache": {
            "task": "app.tasks.cache.update_popular_posts_cache",
            "schedule": crontab(minute="*/30"),  # 每30分钟
        },
    },
)

# app/tasks/email.py
from app.tasks.celery_app import celery_app
from app.utils.email import send_email
import logging

logger = logging.getLogger(__name__)

@celery_app.task(bind=True, max_retries=3)
def send_welcome_email(self, user_id: int, email: str, username: str):
    """发送欢迎邮件"""
    try:
        subject = f"欢迎 {username}!"
        send_email(
            email_to=email,
            subject=subject,
            template_name="welcome",
            template_context={"username": username}
        )
        logger.info(f"欢迎邮件已发送到 {email}")
        return {"status": "success", "email": email}

    except Exception as e:
        logger.error(f"发送欢迎邮件失败: {e}")
        # 重试
        self.retry(exc=e, countdown=60 * 5)  # 5分钟后重试

@celery_app.task
def send_batch_emails(emails: list, subject: str, body: str):
    """批量发送邮件"""
    from app.utils.email import send_batch_email

    result = send_batch_email(emails, subject, body)

    logger.info(f"批量邮件发送完成: {result['sent']}/{result['total']}")
    return result

# app/tasks/report.py
from app.tasks.celery_app import celery_app
from app.services.report import generate_weekly_report
import logging

logger = logging.getLogger(__name__)

@celery_app.task
def send_weekly_report():
    """发送周报"""
    logger.info("开始生成周报...")

    # 生成报告
    report = generate_weekly_report()

    # 发送给管理员
    from app.services.user import get_admin_emails
    admin_emails = get_admin_emails()

    from app.utils.email import send_email
    for email in admin_emails:
        send_email(
            email_to=email,
            subject=f"周报 - {report['week']}",
            template_name="weekly_report",
            template_context=report
        )

    logger.info(f"周报已发送给 {len(admin_emails)} 位管理员")
    return {"status": "success", "recipients": len(admin_emails)}

# 在 FastAPI 中调用 Celery 任务
# app/api/v1/endpoints/users.py
from app.tasks.email import send_welcome_email

@router.post("/register")
async def register(user_in: UserCreate, db: AsyncSession = Depends(get_async_db)):
    # 创建用户
    user = await user_crud.create(db, obj_in=user_in)

    # 异步发送欢迎邮件
    send_welcome_email.delay(user.id, user.email, user.username)

    return user

九、测试与部署

9.1 单元测试

# tests/conftest.py
import pytest
from typing import AsyncGenerator, Generator
from httpx import AsyncClient, ASGITransport
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker

from app.main import app
from app.db.session import Base, get_async_db
from app.core.config import settings
from app.models.user import User
from app.core.security import get_password_hash

# 测试数据库 URL
TEST_DATABASE_URL = settings.DATABASE_URL.replace(
    settings.DB_NAME, 
    f"{settings.DB_NAME}_test"
)

# 创建测试引擎
test_engine = create_async_engine(TEST_DATABASE_URL, echo=False)
TestSessionLocal = async_sessionmaker(
    test_engine,
    class_=AsyncSession,
    expire_on_commit=False
)

async def override_get_db() -> AsyncGenerator[AsyncSession, None]:
    """覆盖依赖项,使用测试数据库"""
    async with TestSessionLocal() as session:
        yield session

app.dependency_overrides[get_async_db] = override_get_db

@pytest.fixture(scope="session")
def anyio_backend():
    """指定异步后端"""
    return "asyncio"

@pytest.fixture(scope="session")
async def setup_database():
    """设置测试数据库"""
    # 创建表
    async with test_engine.begin() as conn:
        await conn.run_sync(Base.metadata.drop_all)
        await conn.run_sync(Base.metadata.create_all)

    yield

    # 清理
    async with test_engine.begin() as conn:
        await conn.run_sync(Base.metadata.drop_all)

@pytest.fixture
async def client(setup_database) -> AsyncGenerator[AsyncClient, None]:
    """创建测试客户端"""
    transport = ASGITransport(app=app)
    async with AsyncClient(
        transport=transport,
        base_url="http://test"
    ) as ac:
        yield ac

@pytest.fixture
async def db_session() -> AsyncGenerator[AsyncSession, None]:
    """获取数据库会话"""
    async with TestSessionLocal() as session:
        yield session

@pytest.fixture
async def test_user(db_session: AsyncSession) -> User:
    """创建测试用户"""
    user = User(
        username="testuser",
        email="test@example.com",
        hashed_password=get_password_hash("testpass123"),
        full_name="Test User",
        is_active=True,
        email_verified=True
    )
    db_session.add(user)
    await db_session.commit()
    await db_session.refresh(user)
    return user

@pytest.fixture
async def auth_headers(client: AsyncClient, test_user: User) -> dict:
    """获取认证头"""
    response = await client.post("/api/v1/auth/login", data={
        "username": test_user.email,
        "password": "testpass123"
    })
    token = response.json()["access_token"]
    return {"Authorization": f"Bearer {token}"}

# tests/test_api/test_users.py
import pytest
from httpx import AsyncClient

@pytest.mark.asyncio
async def test_register_user(client: AsyncClient):
    """测试用户注册"""
    response = await client.post("/api/v1/auth/register", json={
        "username": "newuser",
        "email": "new@example.com",
        "password": "NewPass123",
        "password_confirm": "NewPass123",
        "full_name": "New User"
    })

    assert response.status_code == 201
    data = response.json()
    assert data["username"] == "newuser"
    assert data["email"] == "new@example.com"
    assert "password" not in data

@pytest.mark.asyncio
async def test_register_duplicate_email(
    client: AsyncClient, 
    test_user: User
):
    """测试重复邮箱注册"""
    response = await client.post("/api/v1/auth/register", json={
        "username": "another",
        "email": test_user.email,
        "password": "Pass12345",
        "password_confirm": "Pass12345"
    })

    assert response.status_code == 400
    assert "已被注册" in response.json()["detail"]

@pytest.mark.asyncio
async def test_login_success(client: AsyncClient, test_user: User):
    """测试登录成功"""
    response = await client.post("/api/v1/auth/login", data={
        "username": test_user.email,
        "password": "testpass123"
    })

    assert response.status_code == 200
    data = response.json()
    assert "access_token" in data
    assert data["token_type"] == "bearer"

@pytest.mark.asyncio
async def test_login_wrong_password(client: AsyncClient, test_user: User):
    """测试错误密码"""
    response = await client.post("/api/v1/auth/login", data={
        "username": test_user.email,
        "password": "wrongpassword"
    })

    assert response.status_code == 401

@pytest.mark.asyncio
async def test_get_current_user(
    client: AsyncClient, 
    test_user: User,
    auth_headers: dict
):
    """测试获取当前用户"""
    response = await client.get(
        "/api/v1/auth/me",
        headers=auth_headers
    )

    assert response.status_code == 200
    data = response.json()
    assert data["id"] == test_user.id
    assert data["email"] == test_user.email

@pytest.mark.asyncio
async def test_create_post(
    client: AsyncClient,
    auth_headers: dict,
    db_session
):
    """测试创建文章"""
    # 先创建分类
    from app.models.post import Category
    category = Category(name="Test", slug="test")
    db_session.add(category)
    await db_session.commit()

    response = await client.post(
        "/api/v1/posts/",
        headers=auth_headers,
        json={
            "title": "Test Post",
            "content": "This is a test post content.",
            "category_id": category.id,
            "is_published": True
        }
    )

    assert response.status_code == 201
    data = response.json()
    assert data["title"] == "Test Post"
    assert data["slug"] == "test-post"

# tests/test_services/test_user_service.py
import pytest
from sqlalchemy.ext.asyncio import AsyncSession
from app.crud.user import user_crud
from app.schemas.user import UserCreate

@pytest.mark.asyncio
async def test_create_user(db_session: AsyncSession):
    """测试创建用户服务"""
    user_in = UserCreate(
        username="servicetest",
        email="service@test.com",
        password="Service123",
        password_confirm="Service123"
    )

    user = await user_crud.create(db_session, obj_in=user_in)

    assert user.username == "servicetest"
    assert user.email == "service@test.com"
    assert user.hashed_password != "Service123"

@pytest.mark.asyncio
async def test_authenticate_user(db_session: AsyncSession):
    """测试用户认证"""
    # 创建用户
    user_in = UserCreate(
        username="authtest",
        email="auth@test.com",
        password="Auth12345",
        password_confirm="Auth12345"
    )
    await user_crud.create(db_session, obj_in=user_in)

    # 正确密码
    user = await user_crud.authenticate(
        db_session,
        email="auth@test.com",
        password="Auth12345"
    )
    assert user is not None

    # 错误密码
    user = await user_crud.authenticate(
        db_session,
        email="auth@test.com",
        password="wrong"
    )
    assert user is None

9.2 部署配置

# Dockerfile
"""
FROM python:3.11-slim

WORKDIR /app

# 安装系统依赖
RUN apt-get update && apt-get install -y \
    gcc \
    postgresql-client \
    && rm -rf /var/lib/apt/lists/*

# 复制依赖文件
COPY requirements.txt .

# 安装 Python 依赖
RUN pip install --no-cache-dir -r requirements.txt

# 复制应用代码
COPY . .

# 创建非 root 用户
RUN useradd -m -u 1000 fastapi && chown -R fastapi:fastapi /app
USER fastapi

# 暴露端口
EXPOSE 8000

# 启动命令
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]
"""

# docker-compose.yml
"""
version: '3.8'

services:
  api:
    build: .
    ports:
      - "8000:8000"
    environment:
      - DATABASE_URL=postgresql+asyncpg://postgres:password@db:5432/fastapi
      - REDIS_URL=redis://redis:6379/0
      - SECRET_KEY=${SECRET_KEY}
    depends_on:
      db:
        condition: service_healthy
      redis:
        condition: service_healthy
    volumes:
      - ./uploads:/app/uploads
      - ./logs:/app/logs
    restart: unless-stopped
    command: >
      sh -c "alembic upgrade head &&
             uvicorn app.main:app --host 0.0.0.0 --port 8000 --reload"

  db:
    image: postgres:15-alpine
    environment:
      - POSTGRES_USER=postgres
      - POSTGRES_PASSWORD=password
      - POSTGRES_DB=fastapi
    volumes:
      - postgres_data:/var/lib/postgresql/data
    healthcheck:
      test: ["CMD-SHELL", "pg_isready -U postgres"]
      interval: 10s
      timeout: 5s
      retries: 5
    restart: unless-stopped

  redis:
    image: redis:7-alpine
    volumes:
      - redis_data:/data
    healthcheck:
      test: ["CMD", "redis-cli", "ping"]
      interval: 10s
      timeout: 5s
      retries: 5
    restart: unless-stopped

  celery_worker:
    build: .
    environment:
      - DATABASE_URL=postgresql+asyncpg://postgres:password@db:5432/fastapi
      - REDIS_URL=redis://redis:6379/0
    depends_on:
      - db
      - redis
    volumes:
      - ./uploads:/app/uploads
      - ./logs:/app/logs
    command: celery -A app.tasks.celery_app worker --loglevel=info
    restart: unless-stopped

  celery_beat:
    build: .
    environment:
      - DATABASE_URL=postgresql+asyncpg://postgres:password@db:5432/fastapi
      - REDIS_URL=redis://redis:6379/0
    depends_on:
      - db
      - redis
    command: celery -A app.tasks.celery_app beat --loglevel=info
    restart: unless-stopped

  nginx:
    image: nginx:alpine
    ports:
      - "80:80"
      - "443:443"
    volumes:
      - ./nginx.conf:/etc/nginx/nginx.conf:ro
      - ./ssl:/etc/nginx/ssl:ro
      - ./static:/app/static:ro
    depends_on:
      - api
    restart: unless-stopped

volumes:
  postgres_data:
  redis_data:
"""

十、总结

FastAPI 核心优势

1. 性能卓越 - 基于 Starlette(ASGI)和 Pydantic - 异步支持,高并发处理能力 - 性能媲美 NodeJS 和 Go

2. 开发效率 - 自动生成 OpenAPI 文档(Swagger UI 和 ReDoc) - 强大的依赖注入系统 - Pydantic 自动数据验证和序列化 - 出色的编辑器支持和类型提示

3. 现代化特性 - 原生异步支持(async/await) - WebSocket 支持 - 后台任务 - 中间件系统 - 依赖注入

4. 适用场景 - RESTful API 服务 - 微服务架构 - 实时应用(WebSocket) - 机器学习模型服务 - 高性能 Web 应用

最佳实践总结

"""
FastAPI 最佳实践清单:

1. 项目结构
   ✓ 使用模块化结构分离关注点
   ✓ 将配置、模型、路由、服务分开
   ✓ 使用 Pydantic 进行数据验证

2. 性能优化
   ✓ 使用异步数据库驱动(asyncpg, aiomysql)
   ✓ 实现 Redis 缓存
   ✓ 使用连接池
   ✓ 合理设置并发限制

3. 安全性
   ✓ 使用环境变量管理敏感信息
   ✓ 实现 JWT 认证
   ✓ 启用 CORS 保护
   ✓ 添加安全头中间件
   ✓ 实现请求限流

4. 错误处理
   ✓ 统一异常处理
   ✓ 使用 HTTPException
   ✓ 记录错误日志
   ✓ 返回友好的错误信息

5. 测试
   ✓ 编写单元测试和集成测试
   ✓ 使用测试数据库
   ✓ 模拟外部依赖
   ✓ 达到 80% 以上覆盖率

6. 部署
   ✓ 使用 Gunicorn + Uvicorn workers
   ✓ 配置 Nginx 反向代理
   ✓ 使用 Docker 容器化
   ✓ 实现健康检查端点
   ✓ 配置日志聚合

7. API 设计
   ✓ 遵循 RESTful 规范
   ✓ 使用版本控制(/api/v1/)
   ✓ 实现分页、过滤、排序
   ✓ 返回一致的响应格式
   ✓ 编写 API 文档
"""

FastAPI 代表了 Python Web 框架的现代化发展方向。它将类型提示、异步编程和自动文档生成完美结合,为开发者提供了极佳的开发体验。无论是构建微服务、REST API 还是实时应用,FastAPI 都是一个值得深入学习和使用的优秀框架。


本文由 尚先生 原创,转载请注明出处。

📖相关推荐

评论

0
暂无评论,来发表第一条评论吧

发表评论

登录 后发表评论