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