217 lines
7.0 KiB
Python
217 lines
7.0 KiB
Python
"""第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()
|