ForcePilot/backend/test/storage/transactions/test_sqlalchemy_adapter.py

290 lines
10 KiB
Python
Raw Normal View History

"""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 / Assertsavepoint 内抛异常,被捕获后外层仍能正常退出
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()`` 返回构造时传入的 sessionC-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",
)
# AssertUnifiedError 协议字段
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:
# Assertyield 的是 factory 创建的 session
assert s is session
# Assertfactory 调用一次session 关闭一次
factory.assert_called_once_with()
session.close.assert_awaited_once()