"""第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()