Files
PythonLearn/04_数据库/4_4_SQLAlchemy关系映射与工程实践/relationship_query_example.py
T

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