feat:更新架构
This commit is contained in:
820
docs/Web-Backend/FastAPI/FastAPI-Guide.md
Normal file
820
docs/Web-Backend/FastAPI/FastAPI-Guide.md
Normal file
@@ -0,0 +1,820 @@
|
||||
# 一、简介
|
||||
  [FastAPI](https://fastapi.tiangolo.com/zh/) 是一个用于构建 API 的现代、快速(高性能)的 web 框架,使用 Python 并基于标准的 Python 类型提示。
|
||||
  安装:
|
||||
```cmd
|
||||
pip install fastapi
|
||||
pip install uvicorn[standard] # ASGI服务器
|
||||
```
|
||||
|
||||
# 二、数据层
|
||||
## 2.1 数据库ORM SQLAlchemy
|
||||
### 2.1.1 配置与连接
|
||||
```python
|
||||
# database.py - 数据库配置
|
||||
from sqlalchemy import create_engine, MetaData
|
||||
from sqlalchemy.ext.declarative import declarative_base
|
||||
from sqlalchemy.orm import sessionmaker, Session
|
||||
from typing import Generator
|
||||
import os
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# 数据库配置
|
||||
DATABASE_URL = f"mysql+pymysql://{os.getenv("DATABASE_URL", "root:123456@127.0.0.1:3306/test")}"
|
||||
|
||||
# 创建引擎
|
||||
engine = create_engine(
|
||||
DATABASE_URL,
|
||||
echo=True, # 显示 SQL 语句(开发环境)
|
||||
pool_size=20, # 连接池大小
|
||||
max_overflow=40, # 最大溢出连接数
|
||||
pool_pre_ping=True, # 连接前 ping
|
||||
pool_recycle=3600, # 连接回收时间(秒)
|
||||
)
|
||||
|
||||
# 创建会话工厂
|
||||
SessionLocal = sessionmaker(
|
||||
autocommit=False,
|
||||
autoflush=False,
|
||||
bind=engine,
|
||||
expire_on_commit=False, # 提交后不使实例过期
|
||||
)
|
||||
|
||||
# 声明基类
|
||||
Base = declarative_base()
|
||||
|
||||
# 依赖注入:获取数据库会话
|
||||
def get_db() -> Generator[Session, None, None]:
|
||||
"""
|
||||
获取数据库会话
|
||||
使用 yield 确保会话正确关闭
|
||||
"""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
db.close()
|
||||
```
|
||||
|
||||
>[!WARNING]
|
||||
>需要把pool_pre_ping=True打开,否则会出现超过pool_recycle时间后,数据库断线。
|
||||
|
||||
### 2.1.2 数据库基类
|
||||
```python
|
||||
from sqlalchemy import Column, BigInteger, String, DateTime, event
|
||||
from sqlalchemy.ext.declarative import declared_attr
|
||||
|
||||
from config.auth import context_sub
|
||||
from config.database import Base
|
||||
from datetime import datetime
|
||||
|
||||
from utils.common import camel_to_snake
|
||||
from id_generator import options, generator
|
||||
|
||||
# https://github.com/yitter/IdGenerator/tree/master/Python
|
||||
options = options.IdGeneratorOptions(worker_id=23)
|
||||
idgen = generator.DefaultIdGenerator()
|
||||
idgen.set_id_generator(options)
|
||||
|
||||
|
||||
# 第二层基类:包含ID
|
||||
class IdBase(Base):
|
||||
__abstract__ = True
|
||||
|
||||
id = Column(BigInteger, primary_key=True, index=True)
|
||||
|
||||
@declared_attr
|
||||
def __tablename__(cls):
|
||||
# 自动把数据库实体类名驼峰转为数据库表名下划线
|
||||
return camel_to_snake(cls.__name__)
|
||||
|
||||
|
||||
# 自动填充id
|
||||
@event.listens_for(IdBase, 'before_insert', propagate=True)
|
||||
def before_insert_listener(mapper, connection, target):
|
||||
if target.id is None:
|
||||
target.id = idgen.next_id()
|
||||
|
||||
|
||||
# 第二层基类:包含ID和审计字段
|
||||
class AuditBase(IdBase):
|
||||
__abstract__ = True
|
||||
|
||||
create_time = Column(DateTime, nullable=True, default=datetime.now)
|
||||
create_by = Column(String(255), nullable=True)
|
||||
update_time = Column(DateTime, nullable=True, default=datetime.now, onupdate=datetime.now)
|
||||
update_by = Column(String(255), nullable=True)
|
||||
|
||||
|
||||
@event.listens_for(AuditBase, 'before_insert', propagate=True)
|
||||
def before_insert(mapper, connection, target):
|
||||
value = context_sub.get(None)
|
||||
if value is not None:
|
||||
target.create_by = value
|
||||
|
||||
|
||||
@event.listens_for(AuditBase, 'before_update', propagate=True)
|
||||
def before_update(mapper, connection, target):
|
||||
value = context_sub.get(None)
|
||||
if value is not None:
|
||||
target.update_by = value
|
||||
```
|
||||
|
||||
  id采用yitter雪花id。
|
||||
  审计字段中的create_time和update_time通过SQLAlchemy参数配置。
|
||||
  create_by和update_by字段为登录认证中存储的账号信息。
|
||||
|
||||
### 2.1.3 数据库模型
|
||||
```python
|
||||
from sqlalchemy import (BigInteger, Boolean, Column, Integer, DECIMAL, String, LargeBinary)
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from models.base import AuditBase, IdBase
|
||||
|
||||
|
||||
class Blog(AuditBase):
|
||||
title = Column(String(255), nullable=False, comment="博客标题")
|
||||
top_value = Column(Integer, nullable=False, default=0, comment="置顶值 越大越靠前")
|
||||
is_great = Column(Boolean, nullable=False, default=False, comment="是否是精品")
|
||||
category_id = Column(BigInteger, nullable=False, comment="博客类别")
|
||||
summary = Column(String(255), nullable=False, comment="博客内容概要")
|
||||
content_id = Column(BigInteger, nullable=False, comment="博客内容")
|
||||
word_count = Column(Integer, nullable=False, default=0, comment="字数统计")
|
||||
read_duration = Column(DECIMAL(10, 2), nullable=False, default=0.00, comment="阅读时长")
|
||||
is_approved = Column(Boolean, nullable=False, default=False, comment="是否发布")
|
||||
|
||||
# 分类关系(多对一)
|
||||
category = relationship(
|
||||
"BlogCategory",
|
||||
# 与BlogCategory的blogs属性建立双向关系
|
||||
back_populates="blogs",
|
||||
# 只级联保存和合并操作,不级联删除(删除博客不应删除分类
|
||||
cascade="save-update, merge",
|
||||
# 明确指定连接条件
|
||||
# 如果数据库设置了外键可以省略
|
||||
primaryjoin="foreign(Blog.category_id) == BlogCategory.id"
|
||||
)
|
||||
|
||||
# 内容关系(一对一)
|
||||
content = relationship(
|
||||
"BlogContent",
|
||||
# 与BlogContent的blog属性建立双向关系
|
||||
back_populates="blog",
|
||||
# 完全级联操作:保存、合并、刷新、删除等所有操作都会级联
|
||||
cascade="all, delete-orphan",
|
||||
# 设置为False表示一对一关系,返回单个对象而不是列表
|
||||
uselist=False,
|
||||
# 确保内容只有一个父博客,与delete-orphan配合使用
|
||||
single_parent=True,
|
||||
# 明确指定连接条件
|
||||
primaryjoin="foreign(Blog.content_id) == BlogContent.id"
|
||||
)
|
||||
|
||||
# 评论(一对多)
|
||||
comments = relationship(
|
||||
"BlogComment",
|
||||
# 与BlogComment的blog属性建立双向关系
|
||||
back_populates="blog",
|
||||
# 完全级联操作:博客删除时自动删除所有评论
|
||||
cascade="all, delete-orphan",
|
||||
# 明确指定连接条件
|
||||
primaryjoin="Blog.id == foreign(BlogComment.blog_id)"
|
||||
)
|
||||
|
||||
|
||||
class BlogCategory(AuditBase):
|
||||
name = Column(String(45), nullable=False, comment="类别名称")
|
||||
|
||||
blogs = relationship(
|
||||
"Blog",
|
||||
# 与Blog的category属性建立双向关系
|
||||
back_populates="category",
|
||||
# 完全级联操作:分类删除时自动删除所有关联的博客
|
||||
# 警告:这会级联删除分类下的所有博客,包括博客的内容、访问记录和评论
|
||||
cascade="all, delete-orphan",
|
||||
# 明确指定连接条件
|
||||
primaryjoin="BlogCategory.id == foreign(Blog.category_id)"
|
||||
)
|
||||
|
||||
|
||||
class BlogContent(IdBase):
|
||||
content = Column(LargeBinary, nullable=False, comment="博客内容")
|
||||
|
||||
blog = relationship(
|
||||
"Blog",
|
||||
# 与Blog的content属性建立双向关系
|
||||
back_populates="content",
|
||||
# 设置为False表示一对一关系
|
||||
uselist=False,
|
||||
# 明确指定连接条件
|
||||
primaryjoin="BlogContent.id == foreign(Blog.content_id)"
|
||||
)
|
||||
|
||||
|
||||
class BlogComment(AuditBase):
|
||||
blog_id = Column(BigInteger, nullable=False, comment="博客ID")
|
||||
parent_id = Column(BigInteger, nullable=False, comment="父评论ID")
|
||||
name = Column(String(255), nullable=False, comment="评论人昵称")
|
||||
website = Column(String(255), nullable=True, comment="评论人网站")
|
||||
ip_address = Column(String(45), nullable=False, comment="评论人IP")
|
||||
user_agent = Column(String(255), nullable=False, comment="评论人浏览器信息")
|
||||
content = Column(String(255), nullable=False, comment="评论内容")
|
||||
is_approved = Column(Boolean, nullable=False, default=False, comment="是否通过")
|
||||
|
||||
blog = relationship(
|
||||
"Blog",
|
||||
# 与Blog的comments属性建立双向关系
|
||||
back_populates="comments",
|
||||
# 明确指定连接条件
|
||||
primaryjoin="foreign(BlogComment.blog_id) == Blog.id"
|
||||
)
|
||||
```
|
||||
|
||||
  这里没有通过在数据库建立外键,而是通过SQLAlchemy中的relationship来建立。
|
||||
  其中primaryjoin明确指定了连接条件。如果为一对一,需要设置uselist=False和single_parent=True,返回单个对象而不是列表。
|
||||
  cascade表示级联操作:
|
||||
```
|
||||
# cascade 的完整可选值列表:
|
||||
|
||||
# 基本级联操作
|
||||
"save-update" # 保存/更新时级联
|
||||
"merge" # 合并会话时级联
|
||||
"refresh-expire" # 刷新过期对象时级联
|
||||
"expunge" # 从会话中移除时级联
|
||||
"delete" # 删除时级联
|
||||
"delete-orphan" # 成为孤儿时删除
|
||||
|
||||
# 快捷组合
|
||||
"all" # 包含除 delete-orphan 外的所有操作
|
||||
"all, delete-orphan" # 包含所有操作
|
||||
"none" # 禁用所有级联(默认)
|
||||
|
||||
# 其他组合
|
||||
"save-update, merge" # 常用组合
|
||||
"save-update, merge, delete" # 包含删除
|
||||
"save-update, merge, refresh-expire"
|
||||
```
|
||||
|
||||
## 2.2 数据验证pydantic
|
||||
```python
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, ConfigDict
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
class BlogQuery(BaseModel):
|
||||
"""博客查询参数"""
|
||||
category: Optional[str] = Field(None, description="分类名称")
|
||||
title: Optional[str] = Field(None, description="标题关键词")
|
||||
year: Optional[int] = Field(None, description="发布年份", ge=2000, le=datetime.now().year)
|
||||
|
||||
@field_validator('year')
|
||||
def validate_year(cls, v):
|
||||
if v is not None and v > datetime.now().year:
|
||||
raise ValueError('年份不能超过当前年份')
|
||||
return v
|
||||
|
||||
|
||||
class BlogBase(BaseModel):
|
||||
"""博客基础模型"""
|
||||
title: str = Field(..., min_length=1, max_length=255, description="博客标题")
|
||||
top_value: int = Field(default=0, ge=0, description="置顶值,越大越靠前", alias="topValue")
|
||||
is_great: bool = Field(default=False, description="是否是精品", alias="isGreat")
|
||||
category: str = Field(..., min_length=1, max_length=45, description="分类名称")
|
||||
content: Optional[str] = Field(None, description="博客内容")
|
||||
is_approved: bool = Field(default=False, description="是否已发布", alias="isApproved")
|
||||
|
||||
@field_validator('title')
|
||||
def title_not_empty(cls, v):
|
||||
if not v or not v.strip():
|
||||
raise ValueError('标题不能为空')
|
||||
return v.strip()
|
||||
|
||||
@field_validator('category')
|
||||
def category_not_empty(cls, v):
|
||||
if not v or not v.strip():
|
||||
raise ValueError('分类不能为空')
|
||||
return v.strip()
|
||||
|
||||
|
||||
class BlogCreate(BlogBase):
|
||||
pass
|
||||
|
||||
|
||||
class BlogUpdate(BlogBase):
|
||||
pass
|
||||
|
||||
|
||||
class BlogResponse(BlogBase):
|
||||
id: int = Field(..., description="博客ID")
|
||||
summary: Optional[str] = Field(None, description="内容摘要")
|
||||
word_count: Optional[int] = Field(None, description="字数统计", alias="wordCount")
|
||||
read_duration: Optional[float] = Field(None, description="阅读时长", alias="readDuration")
|
||||
visit_count: Optional[int] = Field(0, description="访问次数", alias="visitCount")
|
||||
create_time: datetime = Field(..., description="创建时间", alias="createTime")
|
||||
update_time: datetime = Field(..., description="更新时间", alias="updateTime")
|
||||
|
||||
model_config = ConfigDict(
|
||||
from_attributes=True,
|
||||
populate_by_name=True,
|
||||
json_encoders={
|
||||
datetime: lambda dt: dt.strftime('%Y-%m-%d %H:%M:%S')
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class BlogCategoryResponse(BaseModel):
|
||||
name: str = Field(..., description="分类名称")
|
||||
count: int = Field(..., description="博客数量")
|
||||
|
||||
model_config = ConfigDict(
|
||||
from_attributes=True,
|
||||
populate_by_name=True
|
||||
)
|
||||
|
||||
|
||||
class BlogStatsResponse(BaseModel):
|
||||
blog_count: int = Field(..., description="博客总数", alias="blogCount")
|
||||
category_count: int = Field(..., description="分类总数", alias="categoryCount")
|
||||
word_count: int = Field(..., description="总字数", alias="wordCount")
|
||||
|
||||
|
||||
class BlogVisitResponse(BaseModel):
|
||||
ip: str = Field(..., description="IP地址")
|
||||
os: str = Field(..., description="操作系统")
|
||||
browser: str = Field(..., description="浏览器")
|
||||
uri: str = Field(..., description="访问路径")
|
||||
title: str = Field(None, description="博客标题")
|
||||
visit_time: datetime = Field(..., description="访问时间", alias="visitTime")
|
||||
|
||||
model_config = ConfigDict(
|
||||
from_attributes=True,
|
||||
populate_by_name=True,
|
||||
json_encoders={
|
||||
datetime: lambda dt: dt.strftime('%Y-%m-%d %H:%M:%S')
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class BlogLatestResponse(BaseModel):
|
||||
id: int = Field(..., description="博客ID")
|
||||
title: str = Field(..., description="博客标题")
|
||||
|
||||
model_config = ConfigDict(
|
||||
from_attributes=True,
|
||||
populate_by_name=True
|
||||
)
|
||||
|
||||
|
||||
class BlogAdjacentResponse(BaseModel):
|
||||
id: int = Field(..., description="博客ID")
|
||||
title: str = Field(..., description="博客标题")
|
||||
|
||||
model_config = ConfigDict(
|
||||
from_attributes=True,
|
||||
populate_by_name=True
|
||||
)
|
||||
|
||||
|
||||
class BlogCommentCreate(BaseModel):
|
||||
parent_id: int = Field(default=0, ge=0, description="父评论ID,0表示顶级评论", alias="parentId")
|
||||
name: str = Field(..., min_length=1, max_length=50, description="评论人昵称")
|
||||
website: Optional[str] = Field(None, description="评论人网站")
|
||||
content: str = Field(..., min_length=1, max_length=1000, description="评论内容")
|
||||
|
||||
@field_validator('name')
|
||||
def name_not_empty(cls, v):
|
||||
if not v or not v.strip():
|
||||
raise ValueError('昵称不能为空')
|
||||
return v.strip()
|
||||
|
||||
@field_validator('content')
|
||||
def content_not_empty(cls, v):
|
||||
if not v or not v.strip():
|
||||
raise ValueError('评论内容不能为空')
|
||||
return v.strip()
|
||||
|
||||
|
||||
class BlogCommentResponse(BlogCommentCreate):
|
||||
id: int = Field(..., description="评论ID")
|
||||
ip_address: str = Field(..., description="IP地址", alias="ipAddress")
|
||||
user_agent: str = Field(..., description="浏览器信息", alias="userAgent")
|
||||
create_time: datetime = Field(..., description="创建时间", alias="createTime")
|
||||
|
||||
model_config = ConfigDict(
|
||||
from_attributes=True,
|
||||
populate_by_name=True,
|
||||
json_encoders={
|
||||
datetime: lambda dt: dt.strftime('%Y-%m-%d %H:%M:%S')
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
  `@field_validator`为Pydantic V2的验证器装饰器。
|
||||
  `ConfigDict`为全局配置类,`from_attributes=True`表示允许从SQLAlchemy等ORM对象创建。`populate_by_name=True`表示允许通过字段的别名(alias)来赋值 `json_encoders`表示自定义特定类型的 JSON 序列化方式。
|
||||
|
||||
# 三、服务层
|
||||
## 3.1 查询操作
|
||||
```python
|
||||
def query_blog_by_id(db: Session, blog_id: int) -> BlogResponse:
|
||||
check_blog_exist(db, blog_id)
|
||||
|
||||
stmt = (select(
|
||||
Blog.id,
|
||||
Blog.title,
|
||||
Blog.top_value,
|
||||
Blog.is_great,
|
||||
BlogCategory.name.label("category"),
|
||||
Blog.summary,
|
||||
BlogContent.content,
|
||||
Blog.word_count,
|
||||
Blog.read_duration,
|
||||
func.count(BlogVisit.id).label("visitCount"),
|
||||
Blog.is_approved,
|
||||
Blog.create_time,
|
||||
Blog.update_time
|
||||
).where(Blog.id == blog_id)
|
||||
.outerjoin(BlogCategory, Blog.category_id == BlogCategory.id)
|
||||
.outerjoin(BlogContent, Blog.content_id == BlogContent.id)
|
||||
.outerjoin(BlogVisit, Blog.id == BlogVisit.blog_id))
|
||||
|
||||
blog = db.execute(stmt).first()
|
||||
|
||||
return BlogResponse.model_validate(blog)
|
||||
```
|
||||
|
||||
  条件和分页查询
|
||||
```python
|
||||
def query_blog_by_condition(db: Session, blog_query: BlogQuery) -> List[BlogResponse]:
|
||||
stmt = get_query_blog_by_condition_stmt(blog_query)
|
||||
results = db.execute(stmt).fetchall()
|
||||
return [BlogResponse.model_validate(result) for result in results]
|
||||
|
||||
|
||||
def get_query_blog_by_condition_stmt(blog_query: BlogQuery):
|
||||
stmt = (select(
|
||||
Blog.id,
|
||||
Blog.title,
|
||||
Blog.top_value,
|
||||
Blog.is_great,
|
||||
BlogCategory.name.label("category"),
|
||||
Blog.summary,
|
||||
BlogContent.content,
|
||||
Blog.word_count,
|
||||
Blog.read_duration,
|
||||
func.count(BlogVisit.id).label("visitCount"),
|
||||
Blog.is_approved,
|
||||
Blog.create_time,
|
||||
Blog.update_time
|
||||
).where(Blog.is_approved == 1)
|
||||
.outerjoin(BlogCategory, Blog.category_id == BlogCategory.id)
|
||||
.outerjoin(BlogContent, Blog.content_id == BlogContent.id)
|
||||
.outerjoin(BlogVisit, Blog.id == BlogVisit.blog_id)
|
||||
.group_by(Blog.id))
|
||||
|
||||
conditions = []
|
||||
|
||||
if blog_query.category:
|
||||
conditions.append(BlogCategory.name == blog_query.category)
|
||||
|
||||
if blog_query.title:
|
||||
conditions.append(Blog.title.like(f"%{blog_query.title}%"))
|
||||
|
||||
if blog_query.year:
|
||||
conditions.append(func.extract('year', Blog.create_time) == blog_query.year)
|
||||
|
||||
if conditions:
|
||||
stmt = stmt.where(and_(*conditions))
|
||||
|
||||
stmt = stmt.order_by(desc(Blog.is_great), desc(Blog.update_time))
|
||||
|
||||
return stmt
|
||||
|
||||
|
||||
def query_blog_by_condition_page(db: Session, blog_query: BlogQuery,
|
||||
current_page: int = 1, page_size: int = 10) -> PageResult[BlogResponse]:
|
||||
stmt = get_query_blog_by_condition_stmt(blog_query)
|
||||
return paginate_query(db, stmt, current_page, page_size)
|
||||
|
||||
def paginate_query(db: Session, query, current_page: int = 1, page_size: int = 10) -> PageResult:
|
||||
"""通用分页查询函数"""
|
||||
# 计算总记录数
|
||||
total = db.execute(select(func.count()).select_from(query.subquery())).scalar_one_or_none() or 0
|
||||
|
||||
# 计算总页数
|
||||
total_pages = (total + page_size - 1) // page_size if page_size != 0 else 0
|
||||
|
||||
# 执行分页查询
|
||||
results = db.execute(query.offset((current_page - 1) * page_size).limit(page_size)).all()
|
||||
|
||||
# 转换为字典列表
|
||||
records = [row._asdict() if hasattr(row, "_asdict") else dict(row) for row in results]
|
||||
|
||||
return PageResult(
|
||||
current=current_page,
|
||||
size=page_size,
|
||||
total=total,
|
||||
pages=total_pages,
|
||||
records=records
|
||||
)
|
||||
```
|
||||
|
||||
## 3.2 新增操作
|
||||
```python
|
||||
def add_blog(db: Session, blog: BlogCreate) -> bool:
|
||||
word_count = get_word_count(blog.content)
|
||||
|
||||
db_blog = Blog(
|
||||
title=blog.title,
|
||||
top_value=blog.top_value,
|
||||
is_great=blog.is_great,
|
||||
category_id=add_blog_category(db, blog.category),
|
||||
summary=get_blog_summary(blog.content),
|
||||
word_count=word_count,
|
||||
read_duration=get_read_duration(word_count),
|
||||
is_approved=blog.is_approved
|
||||
)
|
||||
|
||||
# 级联新增
|
||||
db_blog.content = BlogContent(content=blog.content.encode('utf-8'))
|
||||
|
||||
db.add(db_blog)
|
||||
db.commit()
|
||||
db.refresh(db_blog)
|
||||
|
||||
return True
|
||||
```
|
||||
|
||||
  这里的`db_blog.content`可以直接赋值`BlogContent`对象实现级联新增,不用传`content_id`。
|
||||
>[!WARNING]
|
||||
>这里的category和blog是多对一的关系,且需要一定的逻辑处理,不能直接赋值`BlogCategory`对象,否则每新增一条博客都会新增一个`category`。
|
||||
|
||||
## 3.3 更新操作
|
||||
```python
|
||||
def update_blog(db: Session, blog_id: int, blog: BlogUpdate) -> bool:
|
||||
db_blog = check_blog_exist(db, blog_id)
|
||||
|
||||
db_blog.title = blog.title
|
||||
db_blog.top_value = blog.topValue
|
||||
db_blog.is_great = blog.isGreat
|
||||
db_blog.category_id = update_blog_category(db, blog.category)
|
||||
|
||||
# 级联更新
|
||||
db_blog.content = BlogContent(content=blog.content.encode('utf-8'))
|
||||
|
||||
word_count = get_word_count(blog.content)
|
||||
db_blog.summary = get_blog_summary(blog.content),
|
||||
db_blog.word_count = word_count
|
||||
db_blog.read_duration = get_read_duration(word_count)
|
||||
db_blog.is_approved = blog.isApproved
|
||||
|
||||
db.commit()
|
||||
db.refresh(db_blog)
|
||||
|
||||
return True
|
||||
```
|
||||
|
||||
## 3.4 删除操作
|
||||
```python
|
||||
def delete_blog(db: Session, blog_id: int) -> bool:
|
||||
db_blog = check_blog_exist(db, blog_id)
|
||||
|
||||
# 级联删除
|
||||
db.delete(db_blog)
|
||||
db.commit()
|
||||
|
||||
return True
|
||||
```
|
||||
|
||||
# 四、路由层
|
||||
```python
|
||||
router = APIRouter(
|
||||
prefix="/blog",
|
||||
tags=["博客管理"],
|
||||
responses={404: {"description": "Not found"}}
|
||||
)
|
||||
|
||||
@router.get("/page", summary="分页查询博客", response_model=PageResult[BlogResponse])
|
||||
def query_blog_by_page(
|
||||
current_page: int = Query(1, ge=1, alias="currentPage", description="当前页码,从1开始"),
|
||||
page_size: int = Query(10, ge=1, le=100, alias="pageSize", description="每页显示数量,1-100之间"),
|
||||
db: Session = Depends(get_db)
|
||||
):
|
||||
return blog_service.query_blog_by_page(db, current_page, page_size)
|
||||
|
||||
@router.get("/condition", summary="条件查询博客", response_model=List[BlogResponse])
|
||||
def query_blog_by_condition(query: BlogQuery = Depends(), db: Session = Depends(get_db)):
|
||||
return blog_service.query_blog_by_condition(db, query)
|
||||
|
||||
@router.post("", summary="新增博客内容", response_model=bool)
|
||||
def add_blog(
|
||||
blog: BlogCreate = Body(..., description="博客创建数据"),
|
||||
db: Session = Depends(get_db),
|
||||
_=Depends(verify_token)
|
||||
):
|
||||
return blog_service.add_blog(db, blog)
|
||||
|
||||
@router.put("/{blog_id}", summary="更新博客内容", response_model=bool)
|
||||
def update_blog(
|
||||
blog_id: int = Path(..., ge=1, description="博客ID"),
|
||||
blog: BlogUpdate = Body(..., description="博客更新数据"),
|
||||
db: Session = Depends(get_db),
|
||||
_=Depends(verify_token)
|
||||
):
|
||||
return blog_service.update_blog(db, blog_id, blog)
|
||||
|
||||
@router.delete("/{blog_id}", summary="删除博客内容", response_model=bool)
|
||||
def delete_blog(
|
||||
blog_id: int = Path(..., ge=1, description="博客ID"),
|
||||
db: Session = Depends(get_db),
|
||||
_=Depends(verify_token)
|
||||
):
|
||||
return blog_service.delete_blog(db, blog_id)
|
||||
|
||||
@router.put("/{blog_id}/comment", summary="新增博客评论", response_model=bool)
|
||||
def add_blog_comment(
|
||||
request: Request,
|
||||
blog_id: int = Path(..., ge=1, description="博客ID"),
|
||||
blog_comment: BlogCommentCreate = Body(..., description="博客评论数据"),
|
||||
db: Session = Depends(get_db)
|
||||
):
|
||||
return blog_service.add_blog_comment(db, request, blog_id, blog_comment)
|
||||
```
|
||||
|
||||
  路由装饰器参数中的`response_model`表示定义接口返回的数据模型。
|
||||
  参数注解中的`Query()`表示参数来自URL查询字符串,`Path/Query/Body`分别对应路径参数、查询参数、请求体参数。
|
||||
  `Depends()`表示依赖注入,自动解析参数或执行依赖函数。
|
||||
  `db: Session = Depends(get_db)`表示获取数据库连接,`_=Depends(verify_token)`表示验证用户身份。
|
||||
  `request: Request`可以获取到HTTP请求的完整上下文信息。
|
||||
|
||||
>[!TIP]
|
||||
>这里的`query: BlogQuery = Depends()`会从查询参数中自动实例化`BlogQuery`对象,并进行数据验证和类型转换。
|
||||
|
||||
  注册路由:
|
||||
```python
|
||||
from fastapi import FastAPI
|
||||
from .blog import router as blog_router
|
||||
|
||||
def register_routers(app: FastAPI):
|
||||
app.include_router(blog_router, prefix="")
|
||||
|
||||
|
||||
# main.py
|
||||
app = FastAPI(title="Blog Service")
|
||||
register_routers(app)
|
||||
```
|
||||
|
||||
# 五、中间件
|
||||
## 5.1 全局异常处理器
|
||||
```python
|
||||
from fastapi import Request, HTTPException, status
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
from config.logging import logger
|
||||
|
||||
|
||||
# 自定义异常类
|
||||
class AppException(Exception):
|
||||
def __init__(self, message: str, details=None):
|
||||
self.message = message
|
||||
self.details = details
|
||||
|
||||
|
||||
# 全局异常处理中间件
|
||||
async def global_exception_handler(request: Request, call_next):
|
||||
try:
|
||||
# 记录请求信息(可选)
|
||||
logger.info(f"请求: {request.method} {request.url}")
|
||||
if request.query_params:
|
||||
logger.info(f"查询参数: {dict(request.query_params)}")
|
||||
|
||||
response = await call_next(request)
|
||||
|
||||
# 记录响应信息(可选)
|
||||
if response.status_code >= 400:
|
||||
logger.warning(f"响应: {response.status_code}")
|
||||
|
||||
return response
|
||||
|
||||
except AppException as e:
|
||||
# 记录业务异常
|
||||
logger.error(f"业务异常: {e.message} - 详情: {e.details}")
|
||||
return JSONResponse(status_code=status.HTTP_400_BAD_REQUEST,
|
||||
content={"message": e.message, "details": e.details})
|
||||
|
||||
except SQLAlchemyError as e:
|
||||
# 记录数据库异常
|
||||
logger.critical(f"数据库异常: {str(e)}")
|
||||
return JSONResponse(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
content={"message": "数据库操作失败", "details": str(e)})
|
||||
|
||||
except Exception as e:
|
||||
# 记录未知异常(带堆栈信息)
|
||||
logger.critical(f"未知异常: {str(e)}")
|
||||
return JSONResponse(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
content={"code": 500, "message": "服务器内部错误", "details": str(e)})
|
||||
|
||||
|
||||
def get_credentials_exception() -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Could not validate credentials",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
```
|
||||
|
||||
## 5.2 日志处理器
|
||||
```python
|
||||
import atexit
|
||||
import sys
|
||||
|
||||
from fluent import sender
|
||||
from loguru import logger
|
||||
|
||||
from config.setting import settings
|
||||
|
||||
FLUENTD_HOST = settings.FLUENTD_HOST
|
||||
FLUENTD_PORT = 24224
|
||||
TOPIC_TAG = 'blog-service'
|
||||
|
||||
# 日志级别
|
||||
LOG_LEVEL = settings.LOG_LEVEL.upper()
|
||||
|
||||
# 日志格式
|
||||
STDOUT_FORMAT = (
|
||||
"<green>{time:YYYY-MM-DD HH:mm:ss.SSS}</green> | "
|
||||
"<level>{level: <8}</level> | "
|
||||
"<cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - "
|
||||
"<level>{message}</level>"
|
||||
)
|
||||
|
||||
FILE_FORMAT = (
|
||||
"{time:YYYY-MM-DD HH:mm:ss.SSS} | "
|
||||
"{level: <8} | "
|
||||
"{name}:{function}:{line} - {message}"
|
||||
)
|
||||
|
||||
fluent_sender = sender.FluentSender(
|
||||
tag=TOPIC_TAG,
|
||||
host=FLUENTD_HOST,
|
||||
port=FLUENTD_PORT,
|
||||
buffer_max_size=8 * 1024 * 1024,
|
||||
timeout=3.0,
|
||||
retry_timeout=60
|
||||
)
|
||||
|
||||
|
||||
def log_to_fluent(message):
|
||||
try:
|
||||
record = message.record
|
||||
|
||||
# 构建结构化日志数据
|
||||
log_data = {
|
||||
'topic': TOPIC_TAG,
|
||||
'timestamp': record['time'].timestamp(),
|
||||
'level': record['level'].name.lower(),
|
||||
'message': record['message'],
|
||||
'source': f"{record['file'].path}:{record['line']}",
|
||||
'module': record['module'],
|
||||
'function': record['function'],
|
||||
'process_id': record['process'].id,
|
||||
'thread_id': record['thread'].id,
|
||||
**record['extra']
|
||||
}
|
||||
|
||||
if not fluent_sender.emit(TOPIC_TAG, log_data):
|
||||
print(f"Fluentd 发送失败: {fluent_sender.last_error}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"日志处理异常: {str(e)}")
|
||||
|
||||
|
||||
# 移除默认处理器
|
||||
logger.remove()
|
||||
|
||||
# 添加控制台处理器
|
||||
logger.add(
|
||||
sink=sys.stdout,
|
||||
level=LOG_LEVEL,
|
||||
format=STDOUT_FORMAT,
|
||||
colorize=True,
|
||||
backtrace=True, # 显示完整异常堆栈
|
||||
diagnose=True, # 显示详细异常信息
|
||||
)
|
||||
|
||||
if settings.ENVIRONMENT == 'docker':
|
||||
logger.add(
|
||||
log_to_fluent,
|
||||
level=LOG_LEVEL, # 处理 INFO 及以上级别
|
||||
format="{message}", # 原始消息(实际使用结构化数据)
|
||||
backtrace=True, # 启用堆栈回溯
|
||||
diagnose=True # 显示诊断信息
|
||||
)
|
||||
|
||||
atexit.register(fluent_sender.close)
|
||||
|
||||
# 导出配置好的logger
|
||||
__all__ = ["logger"]
|
||||
```
|
||||
165
docs/Web-Backend/FastAPI/OAuth2.md
Normal file
165
docs/Web-Backend/FastAPI/OAuth2.md
Normal file
@@ -0,0 +1,165 @@
|
||||
# 一、基础概念
|
||||
## 1.1 OAuth2 Password Bearer 模式
|
||||
  用于**用户名+密码**登录,获取**access_token**。
|
||||
|
||||
## 1.2 FastAPI 的 OAuth2PasswordBearer
|
||||
  从请求头**Authorization**中提取Token,Token格式为 **Bearer Token**,必须是这个格式,如果不是则会提示**401 Unauthorized**错误。
|
||||
|
||||
# 二、核心流程
|
||||
## 2.1 获取Token
|
||||
```python
|
||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="session")
|
||||
|
||||
@app.post("/session")
|
||||
async def login(form_data: OAuth2PasswordRequestForm = Depends()):
|
||||
# 1. 验证用户名和密码
|
||||
# 2. 返回 Token
|
||||
return Token
|
||||
```
|
||||
|
||||
   这里的**tokenUrl="session"** 对应的是fastapi中的url路径(@app.post("/session")),用于Swagger文档中的Token认证。
|
||||
   前端必须要以FormData表单的形式传递username和password,且必须是username和password字段。
|
||||
```ts
|
||||
const data = new FormData()
|
||||
data.append('username', username)
|
||||
data.append('password', password)
|
||||
```
|
||||
   在获取Token前,一般还需要进行验证用户名和密码是否和数据库中的信息一致。
|
||||
   可以采用JWT格式封装TokenValue:
|
||||
```python
|
||||
def create_token(payload: dict, expires_delta: Optional[timedelta] = None):
|
||||
# 复制一份
|
||||
payload_copy = payload.copy()
|
||||
|
||||
# 加上有效时间
|
||||
if expires_delta:
|
||||
expire = datetime.now() + expires_delta
|
||||
else:
|
||||
expire = datetime.now() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||
payload_copy.update({"exp": expire})
|
||||
|
||||
# 生成jwt Token
|
||||
return jwt.encode(payload_copy, SECRET_KEY, algorithm=ALGORITHM)
|
||||
```
|
||||
|
||||
## 2.2 访问接口
|
||||
```python
|
||||
@router.delete("/{blog_id}", summary="删除博客内容", response_model=bool)
|
||||
def delete_blog(blog_id: int, db: Session = Depends(get_db), _ = Depends(verify_token)):
|
||||
return blog_service.delete_blog(db, blog_id)
|
||||
```
|
||||
   例如访问这个删除接口,在参数中添加验证token的依赖:
|
||||
```python
|
||||
async def verify_token(token: str = Depends(oauth2_scheme)):
|
||||
try:
|
||||
# 1. 检验token信息
|
||||
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
||||
# 2. 校验payload信息
|
||||
if not verify_payload(payload):
|
||||
raise get_credentials_exception()
|
||||
# 3. 校验数据库中是否存在 payload中的账户信息
|
||||
sub = get_payload_sub(payload)
|
||||
|
||||
# 4. 存储账户信息
|
||||
context_sub.set(sub)
|
||||
except jwt.exceptions.InvalidTokenError:
|
||||
raise get_credentials_exception()
|
||||
```
|
||||
   该依赖又依赖于子依赖oauth2_scheme,通过调用OAuth2PasswordBearer方法从请求头Authorization中获取token值。
|
||||
   校验Token通常包含校验格式是否正确和Token包含的账户信息是否正确。
|
||||
|
||||
# 三、参考代码
|
||||
```python
|
||||
from contextvars import ContextVar
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import Depends
|
||||
from fastapi.security import OAuth2PasswordBearer
|
||||
import jwt
|
||||
from passlib.context import CryptContext
|
||||
|
||||
from middleware.exceptions import get_credentials_exception
|
||||
|
||||
# 密钥和算法配置
|
||||
SECRET_KEY = "sjdi@!#3ksj2780se1283"
|
||||
ALGORITHM = "HS256"
|
||||
ACCESS_TOKEN_EXPIRE_MINUTES = 30
|
||||
|
||||
# 密码哈希上下文
|
||||
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
||||
|
||||
# OAuth2 方案
|
||||
# 设置默认登录接口
|
||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="session")
|
||||
|
||||
# 请求上下文
|
||||
context_sub: ContextVar[str] = ContextVar('sub')
|
||||
|
||||
|
||||
# 获取数据库 加密密码
|
||||
def get_password_hash(password: str):
|
||||
return pwd_context.hash(password)
|
||||
|
||||
|
||||
# 验证数据库密码
|
||||
def verify_password(plain_password: str, hashed_password: str):
|
||||
return pwd_context.verify(plain_password, hashed_password)
|
||||
|
||||
|
||||
# 生成payload
|
||||
def create_payload(sub: str) -> dict:
|
||||
return {
|
||||
"sub": sub
|
||||
}
|
||||
|
||||
|
||||
# 验证payload
|
||||
def verify_payload(payload: dict) -> bool:
|
||||
return "sub" in payload
|
||||
|
||||
|
||||
# 从payload中获取用户
|
||||
def get_payload_sub(payload: dict) -> str:
|
||||
return payload["sub"]
|
||||
|
||||
|
||||
# 创建token
|
||||
def create_token(payload: dict, expires_delta: Optional[timedelta] = None):
|
||||
# 复制一份
|
||||
payload_copy = payload.copy()
|
||||
|
||||
# 加上有效时间
|
||||
if expires_delta:
|
||||
expire = datetime.now() + expires_delta
|
||||
else:
|
||||
expire = datetime.now() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||
payload_copy.update({"exp": expire})
|
||||
|
||||
# 生成jwt Token
|
||||
return jwt.encode(payload_copy, SECRET_KEY, algorithm=ALGORITHM)
|
||||
|
||||
|
||||
# 验证token
|
||||
async def verify_token(token: str = Depends(oauth2_scheme)):
|
||||
try:
|
||||
# 1. 检验token信息
|
||||
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
||||
# 2. 校验payload信息
|
||||
if not verify_payload(payload):
|
||||
raise get_credentials_exception()
|
||||
# 3. 校验数据库中是否存在 payload中的账户信息
|
||||
sub = get_payload_sub(payload)
|
||||
|
||||
# 4. 存储账户信息
|
||||
context_sub.set(sub)
|
||||
except jwt.exceptions.InvalidTokenError:
|
||||
raise get_credentials_exception()
|
||||
```
|
||||
|
||||
# 四、注意事项
|
||||
1. 如果整个路由模块都需要Token验证,可以在APIRouter中添加依赖:
|
||||
|
||||
```python
|
||||
protected_router = APIRouter(dependencies=[Depends(verify_token)])
|
||||
```
|
||||
Reference in New Issue
Block a user