Files
PythonLearn/04_数据库/4_3_SQLAlchemy基础/sqlalchemy_crud_example.py
T

194 lines
6.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""第4-3课示例:使用SQLAlchemy 2.x完成ORM增删改查和事务回滚。"""
from decimal import Decimal
from pathlib import Path
import tomllib
from sqlalchemy import Numeric, String, URL, create_engine, delete, select
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.orm import DeclarativeBase, Mapped, Session, mapped_column, sessionmaker
CONFIG_PATH = Path(__file__).with_name("config.toml")
class Base(DeclarativeBase):
"""保存本课所有ORM模型共享的映射元数据。"""
class Book(Base):
"""把Python图书对象映射到course_orm_book表。"""
__tablename__ = "course_orm_book"
id: Mapped[int] = mapped_column(primary_key=True)
isbn: Mapped[str] = mapped_column(String(30), unique=True, nullable=False)
title: Mapped[str] = mapped_column(String(100), nullable=False)
author: Mapped[str] = mapped_column(String(50), nullable=False)
price: Mapped[Decimal] = mapped_column(Numeric(10, 2), nullable=False)
def __repr__(self) -> str:
"""提供适合开发调试的对象显示。"""
return (
f"Book(id={self.id!r}, isbn={self.isbn!r}, "
f"title={self.title!r}, price={self.price!r})"
)
def load_database_config(config_path: Path) -> dict[str, str | int]:
"""读取并检查本地TOML数据库配置。"""
if not config_path.exists():
raise RuntimeError(
"未找到 config.toml,请复制 config.example.toml 并填写练习数据库配置。"
)
with config_path.open("rb") as config_file:
config_data = tomllib.load(config_file)
database_config = config_data.get("postgresql")
if not isinstance(database_config, dict):
raise RuntimeError("config.toml 缺少 [postgresql] 配置节。")
return database_config
def create_database_url(database_config: dict[str, str | int]) -> URL:
"""用结构化参数创建URL,避免手工拼接和处理密码特殊字符。"""
return URL.create(
drivername="postgresql+psycopg",
username=str(database_config["user"]),
password=str(database_config["password"]),
host=str(database_config["host"]),
port=int(database_config["port"]),
database=str(database_config["dbname"]),
)
def reset_example_data(session: Session) -> None:
"""只删除ORM-前缀的课程示例数据。"""
session.execute(delete(Book).where(Book.isbn.like("ORM-%")))
def add_books(session: Session) -> None:
"""创建Python对象并交给Session持久化。"""
books = [
Book(
isbn="ORM-001",
title="Python数据库编程",
author="小明",
price=Decimal("68.00"),
),
Book(
isbn="ORM-002",
title="SQLAlchemy实践",
author="小红",
price=Decimal("88.00"),
),
]
session.add_all(books)
# flush把待处理INSERT发送到数据库,但当前事务尚未提交。
session.flush()
print(f"flush后第一本书ID:{books[0].id}")
def find_books(session: Session) -> list[Book]:
"""使用SQLAlchemy 2.x的select()查询课程图书。"""
statement = (
select(Book)
.where(Book.isbn.like("ORM-%"))
.order_by(Book.isbn)
)
return list(session.scalars(statement).all())
def update_book(session: Session) -> None:
"""查询ORM对象并修改属性,由Session跟踪变化。"""
book = session.scalar(select(Book).where(Book.isbn == "ORM-001"))
if book is None:
raise RuntimeError("没有找到待修改图书ORM-001。")
book.price = Decimal("72.00")
def delete_book(session: Session) -> None:
"""查询ORM对象并标记删除。"""
book = session.scalar(select(Book).where(Book.isbn == "ORM-002"))
if book is None:
raise RuntimeError("没有找到待删除图书ORM-002。")
session.delete(book)
def demonstrate_rollback(session_factory: sessionmaker[Session]) -> None:
"""演示异常离开Session.begin()时自动回滚。"""
try:
with session_factory.begin() as session:
book = session.scalar(select(Book).where(Book.isbn == "ORM-001"))
if book is None:
raise RuntimeError("没有找到回滚演示图书ORM-001。")
book.price = Decimal("1.00")
session.flush()
raise RuntimeError("模拟后续业务失败。")
except RuntimeError as error:
print(f"失败事务已回滚:{error}")
def print_books(title: str, books: list[Book]) -> None:
"""输出一个查询阶段的图书结果。"""
print(title)
for book in books:
print(f"{book.isbn}|{book.title}|作者:{book.author}|价格:{book.price}")
def main() -> None:
"""创建Engine和Session工厂,依次演示ORM增删改查。"""
try:
database_config = load_database_config(CONFIG_PATH)
database_url = create_database_url(database_config)
# Engine通常在应用启动时创建一次;它管理数据库方言和连接池。
engine = create_engine(
database_url,
connect_args={
"connect_timeout": int(database_config.get("connect_timeout", 10))
},
pool_size=5,
max_overflow=5,
pool_pre_ping=True,
echo=False,
)
session_factory = sessionmaker(engine, expire_on_commit=False)
# 根据模型元数据创建缺失的课程表,不会迁移已有表结构。
Base.metadata.create_all(engine)
with session_factory.begin() as session:
reset_example_data(session)
add_books(session)
with session_factory() as session:
print_books("新增后:", find_books(session))
with session_factory.begin() as session:
update_book(session)
delete_book(session)
with session_factory() as session:
print_books("修改并删除后:", find_books(session))
demonstrate_rollback(session_factory)
with session_factory() as session:
print_books("失败事务回滚后:", find_books(session))
except (OSError, tomllib.TOMLDecodeError, KeyError, RuntimeError) as error:
print(f"配置或课程数据错误:{error}")
except SQLAlchemyError as error:
# SQLAlchemyError是SQLAlchemy数据库访问异常的共同基础类型。
print(f"数据库访问失败:{error}")
if __name__ == "__main__":
main()