194 lines
6.4 KiB
Python
194 lines
6.4 KiB
Python
"""第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()
|