Python的FastAPI+MySQL实现简单CRUD

一、创建项目

PyCharm新建FastAPI项目:

image

二、技术实现

1. 数据库操作

对实际数据库CRUD操作封装,使用SQLAlchemy实现,代码base.py:

from typing import Type, TypeVar, List, Generic, Optional, Any, Union, Iterable

from sqlalchemy.orm import Session

# 定义泛型变量
ModelType = TypeVar("ModelType")

class BaseCRUD(Generic[ModelType]):
    def __init__(self, model: Type[ModelType]):
        """
        :param model: 传入 SQLAlchemy 的模型类 (例如 User)
        """
        self.model = model

    # --- 查 ---
    def get(self, db: Session, id: Any) -> Optional[ModelType]:
        """根据 ID 获取单条记录"""
        return db.get(self.model, id)

    def get_multi(
            self,
            db: Session,
            *,
            skip: int = 0,
            limit: int = 100,
            filters: list = None, # 接收 [User.age > 18, User.name.like("%A%")]
            order_by: Optional[Union[Any, Iterable[Any]]] = None, # 支持单个或多个排序字段
    ) -> List[ModelType]:
        """获取多条记录(支持分页)"""
        query = db.query(self.model)
        if filters:
            query = query.filter(*filters)
        if order_by is not None:
            # 如果传入的是单个字段(如 User.id),转成列表处理
            if not isinstance(order_by, (list, tuple)):
                order_by = [order_by]
            # 使用 * 运算符解包列表
            query = query.order_by(*order_by)
        return query.offset(skip).limit(limit).all()

    # --- 增 ---
    def create(
            self,
            db: Session,
            *,
            obj_in: Union[dict, ModelType]  # 允许传入字典或已实例化的模型
    ) -> ModelType:
        """创建记录"""
        # 1. 统一转换为数据字典
        if isinstance(obj_in, dict):
            create_data = obj_in
        else:
            # 如果传入的是个模型实例,提取其有效字段
            create_data = {
                c.name: getattr(obj_in, c.name)
                for c in obj_in.__table__.columns
                if getattr(obj_in, c.name) is not None  # 过滤掉未赋值的字段
            }

        # 2. 核心逻辑:使用解包 (unpacking) 实例化模型
        # 这样就不需要手动循环 setattr 了,性能更好
        db_obj = self.model(**create_data)

        db.add(db_obj)

        try:
            db.commit()
            db.refresh(db_obj)  # 刷新以获取数据库自动生成的 ID 或默认值
        except Exception as e:
            db.rollback()  # 发生冲突(如唯一索引报错)时必须回滚
            raise e

        return db_obj

    # --- 改 ---
    def update(
            self,
            db: Session,
            *,
            db_obj: ModelType = None,
            obj_in: Union[dict, ModelType] # 允许传入字典或模型实例
    ) -> ModelType:
        """创建记录"""
        # 1. 统一转换为字典格式
        if isinstance(obj_in, dict):
            update_data = obj_in
        else:
            # 如果是 SQLAlchemy 模型,将其转为字典(排除内部状态字段)
            update_data = {
                c.name: getattr(obj_in, c.name)
                for c in obj_in.__table__.columns
            }

        # 2. 执行更新逻辑
        for field in update_data:
            if hasattr(db_obj, field) and update_data[field] is not None:
                setattr(db_obj, field, update_data[field])

        db.add(db_obj)
        db.commit()
        db.refresh(db_obj)
        return db_obj

    # --- 删 ---
    def remove(self, db: Session, *, id: int) -> ModelType:
        """删除记录"""
        obj = db.query(self.model).get(id)
        db.delete(obj)
        db.commit()
        return obj

2. 数据库连接

使用pymysql连接MySQL数据库,实现数据库操作database.py:

from sqlalchemy import create_engine
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker

# 1. 定义数据库连接地址 (以 SQLite 为例,MySQL/Postgres 替换 URL 即可)
MYSQL_USER = "root"
MYSQL_PASSWORD = "your_password"
MYSQL_HOST = "127.0.0.1"
MYSQL_PORT = "3306"
MYSQL_DB = "test_db"
SQLALCHEMY_DATABASE_URL = f"mysql+pymysql://{MYSQL_USER}:{MYSQL_PASSWORD}@{MYSQL_HOST}:{MYSQL_PORT}/{MYSQL_DB}?charset=utf8mb4"

# 2. 创建引擎
# pool_recycle: 自动回收连接,防止 MySQL 默认 8 小时断开连接导致的 "MySQL server has gone away"
# pool_size: 连接池大小
engine = create_engine(
    SQLALCHEMY_DATABASE_URL,
    pool_size=10,
    max_overflow=20,
    pool_recycle=3600,
    pool_pre_ping=True
)

# 3. 创建会话工厂
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)

# 4. 创建基本映射类
Base = declarative_base()

# 获取数据库连接的工具函数(依赖注入常用)
def get_db():
    db = SessionLocal()
    try:
        yield db
    finally:
        db.close()

3. 简单模型

使用员工模型employee.py:

from datetime import datetime
from typing import Annotated, Optional

from pydantic import BaseModel, ConfigDict
from pydantic import PlainSerializer
from sqlalchemy import Column, String, DateTime, DOUBLE
from sqlalchemy.ext.declarative import declarative_base

Base = declarative_base()

class Employee(Base):
    __tablename__ = "employee"
    id = Column(String, primary_key=True)
    name = Column(String)
    salary = Column(DOUBLE)
    create_time = Column(DateTime)

# 1. 定义一个全局通用的日期类型
CustomDatetime = Annotated[
    datetime,
    PlainSerializer(lambda v: v.strftime("%Y-%m-%d %H:%M:%S"), return_type=str)
]

# --- 响应模型 (Pydantic v2) ---
class EmployeeSchema(BaseModel):
    id: Optional[str] = None
    name: Optional[str] = None
    salary: Optional[float] = None
    create_time: Optional[CustomDatetime] = None

    model_config = ConfigDict(from_attributes=True)

4. CRUD实现

实现员工表的CRUDmain.py:

from datetime import datetime
from typing import Any, Optional

from fastapi import FastAPI, Depends, HTTPException
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session

from db.base import BaseCRUD
from db.employee import Employee, EmployeeSchema
from db.database import get_db
import json

class CustomJSONResponse(JSONResponse):
    def render(self, content: any) -> bytes:
        return json.dumps(
            content,
            ensure_ascii=False,
            allow_nan=False,
            indent=None,
            separators=(",", ":"),
            # 在这里处理所有的日期对象
            default=lambda obj: obj.strftime("%Y-%m-%d %H:%M:%S") if isinstance(obj, datetime) else str(obj),
        ).encode("utf-8")

app = FastAPI(default_response_class=CustomJSONResponse)

# 实例化
employee_service = BaseCRUD(Employee)

@app.get("/")
async def root():
    return {"message": "Hello World"}


@app.get("/employee/get", response_model=EmployeeSchema)
async def say_hello(id: Any, db: Session = Depends(get_db)):
    return employee_service.get(db, id=id)

@app.post("/employee/list", response_model=list[EmployeeSchema])
async def list_employee(user: Optional[dict] = None, db: Session = Depends(get_db)):
    if user is not None and "name" in user:
        return employee_service.get_multi(db, filters=[Employee.name.like("%" + user["name"] + "%")])
    return employee_service.get_multi(db)

@app.post("/employee/save", response_model=EmployeeSchema)
async def save_employee(user: dict, db: Session = Depends(get_db)):
    user["create_time"] = datetime.now()
    employee = Employee(**user)
    return employee_service.create(db, obj_in=employee)

@app.post("/employee/modify", response_model=EmployeeSchema)
async def modify_employee(user: dict, db: Session = Depends(get_db)):
    employee = Employee(**user)
    db_obj = employee_service.get(db, id=employee.id)
    if db_obj is None:
        raise HTTPException(status_code=404, detail="Employee not found")
    return employee_service.update(db, db_obj=db_obj, obj_in=employee)

@app.post("/employee/delete", response_model=EmployeeSchema)
async def delete_employee(user: dict, db: Session = Depends(get_db)):
    return employee_service.remove(db, id=user["id"])

三、测试

1、启动

指定8080端口运行:

--host 0.0.0.0 --port 8080

image

2、新增

调用/employee/save接口新增数据:

image

结果:新增成功

3、修改

调用/employee/modify接口修改员工数据

image

结果:修改成功

4、列表

调用/employee/list接口查看数据

image

结果:员工列表

5、删除

调用/employee/delete删除员工数据:

image

结果:删除成功

posted @ 2026-01-07 11:26  旧色染新烟  阅读(171)  评论(0)    收藏  举报