"""第4-4课标准示例:SQLAlchemy关系映射、联表查询与DTO。""" from dataclasses import dataclass from decimal import Decimal from pathlib import Path import tomllib from sqlalchemy import ForeignKey, Numeric, String, URL, create_engine, delete, func, select from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.orm import ( DeclarativeBase, Mapped, Session, mapped_column, relationship, selectinload, sessionmaker, ) class Base(DeclarativeBase): """所有ORM模型共同继承的声明式基类。""" class Customer(Base): """客户模型:一名客户可以拥有多张订单。""" __tablename__ = "course_orm_customer" id: Mapped[int] = mapped_column(primary_key=True) customer_code: Mapped[str] = mapped_column(String(30), unique=True, nullable=False) customer_name: Mapped[str] = mapped_column(String(100), nullable=False) # relationship描述Python对象之间的关系,本身不是数据库中的字段。 # back_populates让Customer.orders和Order.customer成为双向关系。 orders: Mapped[list["Order"]] = relationship( back_populates="customer", cascade="all, delete-orphan", ) class Order(Base): """订单模型:每张订单通过外键归属于一名客户。""" __tablename__ = "course_orm_order" id: Mapped[int] = mapped_column(primary_key=True) order_no: Mapped[str] = mapped_column(String(30), unique=True, nullable=False) amount: Mapped[Decimal] = mapped_column(Numeric(12, 2), nullable=False) customer_id: Mapped[int] = mapped_column( ForeignKey("course_orm_customer.id"), nullable=False, ) customer: Mapped[Customer] = relationship(back_populates="orders") @dataclass(frozen=True) class OrderSummaryDTO: """联表查询结果对象,作用类似Java中专门承载查询结果的DTO。""" order_no: str customer_name: str amount: Decimal def load_database_config() -> dict: """从本课目录的本地TOML文件读取数据库配置。""" config_path = Path(__file__).with_name("config.toml") if not config_path.exists(): raise FileNotFoundError( "没有找到config.toml,请复制config.example.toml并填写本地数据库信息。" ) with config_path.open("rb") as config_file: config = tomllib.load(config_file) if "postgresql" not in config: raise KeyError("config.toml中缺少[postgresql]配置段。") return config["postgresql"] def create_database_url(database_config: dict) -> URL: """使用URL.create安全构造连接地址,避免手工拼接密码。""" return URL.create( drivername="postgresql+psycopg", username=database_config["user"], password=database_config["password"], host=database_config["host"], port=database_config["port"], database=database_config["dbname"], ) def reset_and_add_data(session: Session) -> None: """清理并重新创建本课专用数据,保证示例可以重复运行。""" # 先删子表再删父表,满足数据库外键约束。 customer_ids = select(Customer.id).where(Customer.customer_code.like("ORM-C-%")) session.execute(delete(Order).where(Order.customer_id.in_(customer_ids))) session.execute(delete(Customer).where(Customer.customer_code.like("ORM-C-%"))) alice = Customer( customer_code="ORM-C-001", customer_name="张三", orders=[ Order(order_no="ORM-O-001", amount=Decimal("299.00")), Order(order_no="ORM-O-002", amount=Decimal("99.00")), ], ) bob = Customer( customer_code="ORM-C-002", customer_name="李四", orders=[Order(order_no="ORM-O-003", amount=Decimal("599.00"))], ) # cascade配置使新增Customer时能够同时新增其orders集合中的订单。 session.add_all([alice, bob]) def find_customers_with_orders(session: Session) -> list[Customer]: """使用预加载一次取得客户及其订单,避免N+1查询。""" statement = ( select(Customer) .where(Customer.customer_code.like("ORM-C-%")) .options(selectinload(Customer.orders)) .order_by(Customer.customer_code) ) return list(session.scalars(statement)) def find_order_summaries(session: Session) -> list[OrderSummaryDTO]: """显式联表并只查询DTO所需列。""" statement = ( select(Order.order_no, Customer.customer_name, Order.amount) .join(Customer, Order.customer_id == Customer.id) .where(Customer.customer_code.like("ORM-C-%")) .order_by(Order.order_no) ) rows = session.execute(statement).all() return [ OrderSummaryDTO( order_no=row.order_no, customer_name=row.customer_name, amount=row.amount, ) for row in rows ] def count_orders_by_customer(session: Session) -> list[tuple[str, int]]: """让数据库按照客户分组并统计订单数量。""" statement = ( select(Customer.customer_name, func.count(Order.id)) .join(Order, Customer.id == Order.customer_id) .where(Customer.customer_code.like("ORM-C-%")) .group_by(Customer.id, Customer.customer_name) .order_by(Customer.customer_code) ) return [(name, order_count) for name, order_count in session.execute(statement)] def main() -> None: """按事务写入数据,再分别演示三种多表查询。""" engine = None try: database_config = load_database_config() engine = create_engine( create_database_url(database_config), pool_size=5, max_overflow=10, pool_pre_ping=True, connect_args={ "connect_timeout": database_config.get("connect_timeout", 10) }, ) session_factory = sessionmaker(engine, expire_on_commit=False) Base.metadata.create_all(engine) with session_factory.begin() as session: reset_and_add_data(session) with session_factory() as session: print("关系对象查询:") for customer in find_customers_with_orders(session): print(customer.customer_name) for order in customer.orders: print(f" {order.order_no}|金额:{order.amount}") print("DTO联表查询:") for summary in find_order_summaries(session): print( f"{summary.order_no}|{summary.customer_name}|" f"金额:{summary.amount}" ) print("聚合查询:") for customer_name, order_count in count_orders_by_customer(session): print(f"{customer_name}|订单数量:{order_count}") except (OSError, KeyError, tomllib.TOMLDecodeError) as error: print(f"配置读取失败:{error}") except SQLAlchemyError as error: print(f"数据库访问失败:{error}") finally: if engine is not None: engine.dispose() if __name__ == "__main__": main()