- 新增多个业务域的__init__.py模块文件,规范包导出结构 - 调整多个DTO文件的导入路径,统一模块组织方式 - 移除测试文件中多余的空行与导入语句 - 优化部分业务模块的包层级划分
290 lines
10 KiB
Python
290 lines
10 KiB
Python
"""yuxi.storage.transactions 单元测试。
|
||
|
||
覆盖 ``SqlAlchemyTransactionContext`` / ``SqlAlchemyTransactionAdapter`` /
|
||
``with_session`` 与四个事务异常类,使用 ``MagicMock`` / ``AsyncMock`` 模拟
|
||
``AsyncSession`` 与事务对象,不连接真实 DB。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
import pytest
|
||
from yuxi.storage.transactions import (
|
||
SqlAlchemyTransactionContext,
|
||
TransactionBeginError,
|
||
TransactionCommitError,
|
||
TransactionInactiveError,
|
||
TransactionRollbackError,
|
||
with_session,
|
||
)
|
||
|
||
pytestmark = pytest.mark.unit
|
||
|
||
|
||
def _make_session() -> MagicMock:
|
||
"""构造 AsyncSession 桩。
|
||
|
||
默认 ``in_transaction`` 返回 False(无隐式事务),模拟干净 session。
|
||
``begin_nested`` 返回一个 ``AsyncMock``,可作为 async context manager。
|
||
"""
|
||
session = MagicMock()
|
||
session.begin = AsyncMock()
|
||
session.in_transaction = MagicMock(return_value=False)
|
||
session.commit = AsyncMock()
|
||
session.close = AsyncMock()
|
||
# AsyncMock 实例自身支持 async context manager 协议(__aenter__/__aexit__)
|
||
session.begin_nested = MagicMock(return_value=AsyncMock())
|
||
return session
|
||
|
||
|
||
def _make_txn() -> MagicMock:
|
||
"""构造事务对象桩。"""
|
||
txn = MagicMock()
|
||
txn.commit = AsyncMock()
|
||
txn.rollback = AsyncMock()
|
||
txn.is_active = True
|
||
return txn
|
||
|
||
|
||
# ─── __aenter__:autobegin 隐式事务处理 ─────────────────────────────────────
|
||
|
||
|
||
class TestAenterHandlesAutobegin:
|
||
async def test_begin_handles_autobegin_implicit_transaction(self):
|
||
"""SQLAlchemy 2.0 autobegin 语义:前置 SELECT 触发隐式事务,
|
||
``__aenter__`` 应先提交隐式事务再开启显式事务,避免
|
||
``InvalidRequestError: A transaction is already begun``。
|
||
"""
|
||
# Arrange
|
||
session = _make_session()
|
||
session.in_transaction = MagicMock(return_value=True)
|
||
txn = _make_txn()
|
||
session.begin = AsyncMock(return_value=txn)
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
|
||
# Act
|
||
result = await ctx.__aenter__()
|
||
|
||
# Assert
|
||
assert result is ctx
|
||
session.commit.assert_awaited_once() # 先提交隐式事务
|
||
session.begin.assert_awaited_once() # 再开启显式事务
|
||
assert ctx._txn is txn
|
||
|
||
|
||
# ─── __aexit__:is_active 检查避免 ResourceClosedError ─────────────────────
|
||
|
||
|
||
class TestAexitChecksIsActive:
|
||
async def test_aexit_checks_is_active_avoids_resource_closed_error(self):
|
||
"""事务已被 SQLAlchemy 内部关闭(``is_active=False``,如 flush 失败)
|
||
时,``__aexit__`` 应跳过 commit/rollback,避免 ``ResourceClosedError``。
|
||
"""
|
||
# Arrange
|
||
session = _make_session()
|
||
txn = _make_txn()
|
||
txn.is_active = False
|
||
session.begin = AsyncMock(return_value=txn)
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
await ctx.__aenter__()
|
||
|
||
# Act
|
||
await ctx.__aexit__(None, None, None)
|
||
|
||
# Assert
|
||
txn.commit.assert_not_awaited()
|
||
txn.rollback.assert_not_awaited()
|
||
assert ctx._txn is None
|
||
|
||
|
||
# ─── commit / rollback 幂等 ────────────────────────────────────────────────
|
||
|
||
|
||
class TestCommitIdempotent:
|
||
async def test_commit_is_idempotent(self):
|
||
"""多次调用 ``commit()`` 仅首次生效,后续 no-op。"""
|
||
# Arrange
|
||
session = _make_session()
|
||
txn = _make_txn()
|
||
session.begin = AsyncMock(return_value=txn)
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
await ctx.__aenter__()
|
||
|
||
# Act
|
||
await ctx.commit()
|
||
await ctx.commit()
|
||
|
||
# Assert
|
||
txn.commit.assert_awaited_once()
|
||
assert ctx._txn is None
|
||
|
||
|
||
class TestRollbackIdempotent:
|
||
async def test_rollback_is_idempotent(self):
|
||
"""多次调用 ``rollback()`` 仅首次生效,后续 no-op。"""
|
||
# Arrange
|
||
session = _make_session()
|
||
txn = _make_txn()
|
||
session.begin = AsyncMock(return_value=txn)
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
await ctx.__aenter__()
|
||
|
||
# Act
|
||
await ctx.rollback()
|
||
await ctx.rollback()
|
||
|
||
# Assert
|
||
txn.rollback.assert_awaited_once()
|
||
assert ctx._txn is None
|
||
|
||
|
||
# ─── savepoint ─────────────────────────────────────────────────────────────
|
||
|
||
|
||
class TestSavepoint:
|
||
async def test_savepoint_rolls_back_on_exception(self):
|
||
"""``savepoint()`` 在异常时回滚到 SAVEPOINT,外层事务不受污染:
|
||
SAVEPOINT 的 ``__aexit__`` 收到异常,外层 ``_txn`` 仍可正常提交。
|
||
"""
|
||
# Arrange
|
||
session = _make_session()
|
||
txn = _make_txn()
|
||
session.begin = AsyncMock(return_value=txn)
|
||
sp_cm = session.begin_nested.return_value # AsyncMock
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
await ctx.__aenter__()
|
||
|
||
# Act / Assert:savepoint 内抛异常,被捕获后外层仍能正常退出
|
||
with pytest.raises(ValueError, match="boom"):
|
||
async with ctx.savepoint():
|
||
raise ValueError("boom")
|
||
|
||
# SAVEPOINT 的 __aexit__ 收到异常
|
||
sp_cm.__aexit__.assert_awaited_once()
|
||
exc_type_arg = sp_cm.__aexit__.await_args.args[0]
|
||
assert exc_type_arg is ValueError
|
||
# 外层事务未被污染:rollback 未触发
|
||
txn.rollback.assert_not_awaited()
|
||
|
||
# 外层正常退出 → commit
|
||
await ctx.__aexit__(None, None, None)
|
||
txn.commit.assert_awaited_once()
|
||
|
||
async def test_savepoint_commits_on_success(self):
|
||
"""``savepoint()`` 在正常退出时释放 SAVEPOINT(``__aexit__`` 收到 None)。"""
|
||
# Arrange
|
||
session = _make_session()
|
||
sp_cm = session.begin_nested.return_value
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
|
||
# Act
|
||
async with ctx.savepoint():
|
||
pass
|
||
|
||
# Assert
|
||
sp_cm.__aexit__.assert_awaited_once()
|
||
args = sp_cm.__aexit__.await_args.args
|
||
assert args == (None, None, None)
|
||
|
||
|
||
# ─── get_session / transaction_id ──────────────────────────────────────────
|
||
|
||
|
||
class TestGetSession:
|
||
def test_get_session_returns_shared_session(self):
|
||
"""``get_session()`` 返回构造时传入的 session(C-I1 透传)。"""
|
||
# Arrange
|
||
session = _make_session()
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
|
||
# Act / Assert
|
||
assert ctx.get_session() is session
|
||
|
||
|
||
class TestTransactionId:
|
||
def test_transaction_id_generated(self):
|
||
"""``__init__`` 生成 ``tx_`` 前缀的唯一 transaction_id。"""
|
||
# Arrange
|
||
session = _make_session()
|
||
|
||
# Act
|
||
ctx1 = SqlAlchemyTransactionContext(session)
|
||
ctx2 = SqlAlchemyTransactionContext(session)
|
||
|
||
# Assert
|
||
assert ctx1.transaction_id.startswith("tx_")
|
||
assert ctx2.transaction_id.startswith("tx_")
|
||
assert ctx1.transaction_id != ctx2.transaction_id
|
||
|
||
|
||
# ─── 异常链与 UnifiedError 协议 ─────────────────────────────────────────────
|
||
|
||
|
||
class TestExceptionTraceback:
|
||
async def test_exception_preserves_original_traceback(self):
|
||
"""``raise ... from exc`` 保留原始异常链(``__cause__``)。"""
|
||
# Arrange
|
||
session = _make_session()
|
||
txn = _make_txn()
|
||
original = RuntimeError("db connection lost")
|
||
txn.commit = AsyncMock(side_effect=original)
|
||
session.begin = AsyncMock(return_value=txn)
|
||
ctx = SqlAlchemyTransactionContext(session)
|
||
await ctx.__aenter__()
|
||
|
||
# Act / Assert
|
||
with pytest.raises(TransactionCommitError) as exc_info:
|
||
await ctx.commit()
|
||
assert exc_info.value.__cause__ is original
|
||
|
||
|
||
class TestUnifiedErrorProtocol:
|
||
def test_exception_implements_unified_error_protocol(self):
|
||
"""四个异常类提供 status_code / error_code / message / details / trace_id。"""
|
||
cases = [
|
||
(TransactionBeginError, "TX_BEGIN_ERROR"),
|
||
(TransactionCommitError, "TX_COMMIT_ERROR"),
|
||
(TransactionRollbackError, "TX_ROLLBACK_ERROR"),
|
||
(TransactionInactiveError, "TX_INACTIVE_ERROR"),
|
||
]
|
||
for exc_cls, code in cases:
|
||
# Act
|
||
exc = exc_cls(
|
||
"msg",
|
||
transaction_id="tx1",
|
||
session_state="active",
|
||
trace_id="t1",
|
||
)
|
||
|
||
# Assert:UnifiedError 协议字段
|
||
assert exc.error_code == code
|
||
assert exc.message == "msg"
|
||
assert exc.trace_id == "t1"
|
||
# 四个 TX_* 错误码未注册到状态码映射表,基类查表返回 500
|
||
assert exc.status_code == 500
|
||
# details 自动包含诊断字段(to_dict 排除 error_code/message/trace_id)
|
||
details = exc.details
|
||
assert details["transaction_id"] == "tx1"
|
||
assert details["session_state"] == "active"
|
||
|
||
|
||
# ─── with_session 辅助方法 ─────────────────────────────────────────────────
|
||
|
||
|
||
class TestWithSession:
|
||
async def test_with_session_creates_and_closes_temporary_session(self):
|
||
"""``with_session`` 创建临时 session 并在退出时关闭。"""
|
||
# Arrange
|
||
session = _make_session()
|
||
factory = MagicMock(return_value=session)
|
||
|
||
# Act
|
||
async with with_session(factory) as s:
|
||
# Assert:yield 的是 factory 创建的 session
|
||
assert s is session
|
||
|
||
# Assert:factory 调用一次,session 关闭一次
|
||
factory.assert_called_once_with()
|
||
session.close.assert_awaited_once()
|