Files
PythonLearn/04_数据库/4_2_Python数据库事务与数据访问层/transaction_repository_example.py

215 lines
7.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""第 4-2 课示例:使用事务和 Repository 完成安全转账。"""
from decimal import Decimal
from pathlib import Path
import tomllib
import psycopg
from psycopg import Connection
from psycopg.rows import dict_row
CONFIG_PATH = Path(__file__).with_name("config.toml")
class TransferError(Exception):
"""表示转账过程中可以预期的业务失败。"""
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
class AccountRepository:
"""封装账户表 SQL但不自行提交或回滚事务。"""
def __init__(self, connection: Connection) -> None:
self.connection = connection
def create_table(self) -> None:
"""创建本课专用表;表已存在时保持不变。"""
self.connection.execute(
"""
CREATE TABLE IF NOT EXISTS course_bank_account (
account_no VARCHAR(30) PRIMARY KEY,
owner_name VARCHAR(50) NOT NULL,
balance NUMERIC(12, 2) NOT NULL CHECK (balance >= 0)
)
"""
)
def reset_course_accounts(self) -> None:
"""只清理 COURSE- 前缀的课程数据,避免影响其他记录。"""
self.connection.execute(
"DELETE FROM course_bank_account WHERE account_no LIKE %s",
("COURSE-%",),
)
def add_accounts(self, accounts: list[tuple[str, str, Decimal]]) -> None:
"""使用 executemany() 批量新增课程账户。"""
# executemany() 是 Cursor 的方法,因此显式创建并关闭游标。
with self.connection.cursor() as cursor:
cursor.executemany(
"""
INSERT INTO course_bank_account (account_no, owner_name, balance)
VALUES (%s, %s, %s)
""",
accounts,
)
def get_balance_for_update(self, account_no: str) -> Decimal:
"""查询并锁定账户,防止并发事务同时修改同一余额。"""
result = self.connection.execute(
"""
SELECT balance
FROM course_bank_account
WHERE account_no = %s
FOR UPDATE
""",
(account_no,),
).fetchone()
if result is None:
raise TransferError(f"账户不存在:{account_no}")
return result[0]
def change_balance(self, account_no: str, amount: Decimal) -> None:
"""使用数据库加法更新余额,并检查目标账户是否存在。"""
cursor = self.connection.execute(
"""
UPDATE course_bank_account
SET balance = balance + %s
WHERE account_no = %s
""",
(amount, account_no),
)
if cursor.rowcount != 1:
raise TransferError(f"账户不存在:{account_no}")
def find_course_accounts(self) -> list[dict[str, object]]:
"""按账号查询课程账户,并以字典行返回。"""
cursor = self.connection.cursor(row_factory=dict_row)
try:
cursor.execute(
"""
SELECT account_no, owner_name, balance
FROM course_bank_account
WHERE account_no LIKE %s
ORDER BY account_no
""",
("COURSE-%",),
)
return list(cursor.fetchall())
finally:
cursor.close()
class TransferService:
"""组织转账业务规则;事务由调用它的连接上下文统一管理。"""
def __init__(self, repository: AccountRepository) -> None:
self.repository = repository
def transfer(self, source_no: str, target_no: str, amount: Decimal) -> None:
"""在同一事务中完成扣款与入账。"""
if amount <= 0:
raise TransferError("转账金额必须大于 0。")
source_balance = self.repository.get_balance_for_update(source_no)
self.repository.get_balance_for_update(target_no)
if source_balance < amount:
raise TransferError("账户余额不足。")
self.repository.change_balance(source_no, -amount)
self.repository.change_balance(target_no, amount)
def print_accounts(title: str, accounts: list[dict[str, object]]) -> None:
"""输出当前课程账户余额。"""
print(title)
for account in accounts:
print(
f"{account['account_no']}{account['owner_name']}"
f"余额:{account['balance']}"
)
def prepare_data(database_config: dict[str, str | int]) -> None:
"""创建专用表并重置本课固定数据。"""
with psycopg.connect(**database_config) as connection:
repository = AccountRepository(connection)
repository.create_table()
repository.reset_course_accounts()
repository.add_accounts(
[
("COURSE-A001", "小明", Decimal("1000.00")),
("COURSE-A002", "小红", Decimal("500.00")),
]
)
def run_successful_transfer(database_config: dict[str, str | int]) -> None:
"""演示正常离开连接上下文时自动提交事务。"""
with psycopg.connect(**database_config) as connection:
service = TransferService(AccountRepository(connection))
service.transfer("COURSE-A001", "COURSE-A002", Decimal("200.00"))
def run_failed_transfer(database_config: dict[str, str | int]) -> None:
"""演示异常离开连接上下文时自动回滚整个事务。"""
try:
with psycopg.connect(**database_config) as connection:
repository = AccountRepository(connection)
# 先执行一条成功更新,再主动触发业务异常。
# 外层 with 会回滚,因此这 50 元扣款不会保留下来。
repository.change_balance("COURSE-A001", Decimal("-50.00"))
raise TransferError("模拟第二步失败,验证前一步更新会被回滚。")
except TransferError as error:
print(f"失败事务已回滚:{error}")
def query_accounts(
database_config: dict[str, str | int],
) -> list[dict[str, object]]:
"""使用独立连接回查已经提交的数据。"""
with psycopg.connect(**database_config) as connection:
return AccountRepository(connection).find_course_accounts()
def main() -> None:
"""依次演示初始化、提交、回滚和回查。"""
try:
database_config = load_database_config(CONFIG_PATH)
prepare_data(database_config)
print_accounts("初始余额:", query_accounts(database_config))
run_successful_transfer(database_config)
print_accounts("成功转账 200 元后:", query_accounts(database_config))
run_failed_transfer(database_config)
print_accounts("失败事务回滚后:", query_accounts(database_config))
except (OSError, tomllib.TOMLDecodeError, RuntimeError) as error:
print(f"配置读取失败:{error}")
except psycopg.Error as error:
print(f"数据库访问失败:{error}")
if __name__ == "__main__":
main()