02. 使用SQLAlchemy操作数据库
什么是SQLAlchemy
SQLAlchemy 是 Python 生态中最著名、功能最强大的 SQL 工具包和对象关系映射(ORM)框架
可以让你很方便地使用代码操作数据库
创建数据库会话工厂
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.core.config import settings
engine = create_engine(
settings.DATABASE_URL,
pool_pre_ping=True,
echo=settings.DEBUG
)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
数据库 orm 对象
Base 基类
from datetime import datetime
from sqlalchemy import DateTime
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
class Base(DeclarativeBase):
"""所有表共用:创建时间、更新时间(统一用本地时间)"""
create_time: Mapped[datetime] = mapped_column(
DateTime, default=datetime.now, comment="创建时间"
)
update_time: Mapped[datetime] = mapped_column(
DateTime, default=datetime.now, onupdate=datetime.now, comment="更新时间"
)
创建数据库 fastapi-db

商品模型和分类模型
from decimal import Decimal
from sqlalchemy import Integer, String, Numeric, ForeignKey
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base
class Goods(Base):
__tablename__ = 'goods'
__table_args__ = {"comment": "商品信息"}
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True, comment="主键ID")
name: Mapped[str | None] = mapped_column(String(255), comment="商品名称")
price: Mapped[Decimal | None] = mapped_column(Numeric(10, 2), default=0, comment="价格")
stock: Mapped[int | None] = mapped_column(Integer, default=0, comment="库存")
category_id: Mapped[int | None] = mapped_column(Integer, ForeignKey("category.id"), comment="分类ID")
from sqlalchemy import Integer, String
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base
class Category(Base):
__tablename__ = 'category'
__table_args__ = {"comment": "分类信息"}
id: Mapped[int | None] = mapped_column(Integer, primary_key=True, autoincrement=True, comment="主键ID")
name: Mapped[str | None] = mapped_column(String(255), comment="分类名称")
SQLAlchemy orm 框架
if __name__ == '__main__':
Base.metadata.create_all(bind=engine)
db = SessionLocal()
try:
category1 = Category(name="饮料")
category2 = Category(name="零食")
category3 = Category(name="日用品")
db.add_all([category1, category2, category3])
db.flush()
goods1 = Goods(
name="可口可乐",
price=Decimal(3),
stock=100,
category_id=category1.id
)
goods2 = Goods(
name="乐事薯片",
price=Decimal(5),
stock=100,
category_id=category2.id
)
db.add(goods1)
db.add(goods2)
db.flush()
print(f"新增商品ID: {goods1.id}")
db.execute(update(Goods).where(Goods.id == goods1.id).values(name="百事可乐"))
print(f"更新商品名称: {goods1.name}")
dbGoods = db.scalar(select(Goods).where(Goods.id == goods1.id))
if dbGoods:
print(f"单个查询---ID: {goods2.id}, 名称: {goods2.name}")
goodsList = db.scalars(select(Goods)).all()
for g in goodsList:
print(f"批量查询---ID: {g.id}, 名称: {g.name}")
stmt = (
select(Goods, Category.name)
.outerjoin(Category, Goods.category_id == Category.id)
)
rows = db.execute(stmt).all()
for goods, category_name in rows:
print(goods.id, goods.name, goods.price, category_name)
db.execute(delete(Goods).where(Goods.id == goods1.id))
print(f"删除商品ID: {goods1.id}")
db.commit()
except Exception:
db.rollback()
raise
finally:
db.close()