From 713fcdcd8ec6f8954f19a72cc78cdb07bc8fb484 Mon Sep 17 00:00:00 2001 From: Kris <2893855659@qq.com> Date: Fri, 10 Jul 2026 14:29:23 +0800 Subject: [PATCH] =?UTF-8?q?refactor(test):=20=E6=89=B9=E9=87=8F=E9=87=8D?= =?UTF-8?q?=E6=9E=84=E6=B5=8B=E8=AF=95=E4=BB=A3=E7=A0=81=EF=BC=8C=E7=AE=80?= =?UTF-8?q?=E5=8C=96=E8=B0=83=E5=BA=A6=E5=99=A8handler=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. 统一替换多个调度器handler测试用例,移除冗余的session_factory相关代码和fake session上下文 2. 调整ChannelPersistenceAdapter测试,使用模块级patcher管理并统一初始化方式 3. 修正OutboxEntry聚合根操作的版本号逻辑,移除不必要的版本递增 4. 新增UTC时区转换相关测试用例,完善datetime类型映射测试 5. 更新微信WOC适配器测试,适配新的session_type枚举值 6. 优化测试代码的可读性和一致性,统一测试辅助函数的实现方式 --- .../adapters/test_agent_run_adapter.py | 10 +- .../test_channel_persistence_adapter.py | 70 +-- .../test_content_review_repository_adapter.py | 495 ++++++++---------- .../unit/channels/adapters/test_mappers.py | 52 +- .../adapters/test_route_binding_repository.py | 29 +- ...est_channel_audit_log_retention_handler.py | 35 +- ...hannel_content_review_retention_handler.py | 35 +- ...est_channel_idempotency_cleanup_handler.py | 36 +- .../test_channel_outbox_recovery_handler.py | 40 +- ...channel_outbox_terminal_cleanup_handler.py | 59 +-- ...test_channel_pairing_expiration_handler.py | 35 +- ...hannel_pairing_terminal_cleanup_handler.py | 50 +- ...hannel_session_inactive_cleanup_handler.py | 56 +- .../outbound/test_load_build_stage.py | 197 +++++-- .../pipeline/outbound/test_prefix_stage.py | 8 +- .../outbound/test_stream_chunk_stage.py | 367 ++++++++++++- .../usecase/test_identity_merge_service.py | 13 +- .../channels/core/model/test_outbox_entry.py | 16 +- .../core/model/test_outbox_entry_aggregate.py | 4 +- .../infrastructure/test_channel_use_cases.py | 248 ++++----- .../channels/infrastructure/test_factory.py | 46 +- .../infrastructure/test_host_shutdown.py | 55 +- .../channels/infrastructure/test_scheduler.py | 24 +- .../adapters/test_inbound_adapter.py | 38 +- 24 files changed, 1105 insertions(+), 913 deletions(-) diff --git a/backend/test/unit/channels/adapters/test_agent_run_adapter.py b/backend/test/unit/channels/adapters/test_agent_run_adapter.py index 45a280e2..8e7918a7 100644 --- a/backend/test/unit/channels/adapters/test_agent_run_adapter.py +++ b/backend/test/unit/channels/adapters/test_agent_run_adapter.py @@ -549,8 +549,9 @@ class TestAgentRunAdapterStream: service_account_port = MagicMock() adapter = AgentRunAdapter(execution_port, logger, service_account_port) - # Act - results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))] + # Act — streamAgentRun 通过 _session_scope 获取会话调用 getRunUid + with _patch_pg_session(): + results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))] # Assert assert len(results) == 2 @@ -577,7 +578,8 @@ class TestAgentRunAdapterStream: adapter = AgentRunAdapter(execution_port, logger, service_account_port) # Act - results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))] + with _patch_pg_session(): + results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))] # Assert assert results[0].payload == {} @@ -594,7 +596,7 @@ class TestAgentRunAdapterStream: adapter = AgentRunAdapter(execution_port, logger, service_account_port) # Act / Assert - with pytest.raises(DependencyError): + with _patch_pg_session(), pytest.raises(DependencyError): async for _ in adapter.streamAgentRun(AgentRunId("run-1")): pass logger.error.assert_awaited_once() diff --git a/backend/test/unit/channels/adapters/test_channel_persistence_adapter.py b/backend/test/unit/channels/adapters/test_channel_persistence_adapter.py index ec380be9..461a0d23 100644 --- a/backend/test/unit/channels/adapters/test_channel_persistence_adapter.py +++ b/backend/test/unit/channels/adapters/test_channel_persistence_adapter.py @@ -125,13 +125,34 @@ def _make_logger() -> MagicMock: return logger +_active_patchers: list = [] + + def _build_adapter(db: MagicMock, repos: MagicMock) -> ChannelPersistenceAdapter: - """构造 adapter,patch ``create_repositories`` 返回桩 repos。""" - with patch( + """构造 adapter,patch ``create_repositories`` 返回桩 repos。 + + patch 通过 ``start()`` 启动并在模块级 ``_active_patchers`` 列表中注册, + 由 autouse fixture ``_cleanup_patchers`` 在测试方法结束后统一 ``stop()``。 + 这是因为改造后 ``create_repositories`` 在 ``_session_scope(tx)`` 方法调用 + 时执行(而非 ``__init__`` 构造时),patch 需要跨越 ``_build_adapter`` + 返回后继续生效。 + """ + patcher = patch( "yuxi.channels.adapters.channel_persistence_adapter.create_repositories", return_value=repos, - ): - return ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=_make_logger()) + ) + patcher.start() + _active_patchers.append(patcher) + return ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=_make_logger()) + + +@pytest.fixture(autouse=True) +def _cleanup_patchers(): + """每个测试结束后停止所有活跃的 patcher。""" + yield + for p in _active_patchers: + p.stop() + _active_patchers.clear() @pytest.mark.unit @@ -476,11 +497,10 @@ class TestChannelPersistenceAdapterSession: "yuxi.channels.adapters.channel_persistence_adapter.create_repositories", return_value=repos, ): - adapter = ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=logger) - cmd = UpdateChannelSessionCmd(session_id="sess-1", is_temporary=True) - - # Act - await adapter.updateChannelSession(cmd) + adapter = ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=logger) + cmd = UpdateChannelSessionCmd(session_id="sess-1", is_temporary=True) + # Act + await adapter.updateChannelSession(cmd) # Assert: expected_version=None 时应记录 WARN 暴露乐观锁缺口 logger.warn.assert_awaited_once() @@ -500,15 +520,14 @@ class TestChannelPersistenceAdapterSession: "yuxi.channels.adapters.channel_persistence_adapter.create_repositories", return_value=repos, ): - adapter = ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=logger) - cmd = UpdateChannelSessionCmd( - session_id="sess-1", - is_temporary=True, - expected_version=1, - ) - - # Act - await adapter.updateChannelSession(cmd) + adapter = ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=logger) + cmd = UpdateChannelSessionCmd( + session_id="sess-1", + is_temporary=True, + expected_version=1, + ) + # Act + await adapter.updateChannelSession(cmd) # Assert: 传入 expected_version 时不应触发 WARN logger.warn.assert_not_awaited() @@ -835,21 +854,6 @@ class TestChannelPersistenceAdapterCleanup: await adapter.cleanupInactiveSessions(["sess-1"]) -@pytest.mark.unit -class TestChannelPersistenceAdapterClose: - @pytest.mark.asyncio - async def test_aclose_closes_shared_session(self): - # Arrange - db = _make_db() - adapter = _build_adapter(db, _make_repos()) - - # Act - await adapter.aclose() - - # Assert - db.close.assert_awaited_once() - - def _make_pairing_orm(*, version: int = 1, status: str = "pending") -> MagicMock: """构造 ChannelPairing ORM 桩。""" orm = MagicMock() diff --git a/backend/test/unit/channels/adapters/test_content_review_repository_adapter.py b/backend/test/unit/channels/adapters/test_content_review_repository_adapter.py index fbb5a6a7..d588b8cd 100644 --- a/backend/test/unit/channels/adapters/test_content_review_repository_adapter.py +++ b/backend/test/unit/channels/adapters/test_content_review_repository_adapter.py @@ -2,15 +2,15 @@ 覆盖 ``ContentReviewRepositoryAdapter`` 的 ``saveReviewResult`` / ``queryReviewHistory`` / ``countReviewHistory`` / ``getReviewDetail`` / -``getReviewAnalytics`` / ``getReviewStats`` / ``updateReviewVerdict`` 方法, -使用 ``MagicMock`` 模拟 ``AsyncSession``,不连接真实 DB。 +``getReviewAnalytics`` / ``getReviewStats`` / ``updateReviewVerdict`` / +``deleteOldReviewRecords`` 方法,patch +``ChannelContentReviewRecordRepository`` 返回 mock 仓储,不连接真实 DB。 """ from __future__ import annotations from datetime import datetime -from typing import Any -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch from zoneinfo import ZoneInfo import pytest @@ -44,18 +44,28 @@ pytestmark = pytest.mark.unit def _make_db() -> MagicMock: - """构造 AsyncSession 桩,常用方法均为 AsyncMock。""" + """构造 AsyncSession 桩,仅需 close/rollback 供 ``_session_scope`` 使用。""" db = MagicMock() - db.add = MagicMock() - db.flush = AsyncMock() db.commit = AsyncMock() db.rollback = AsyncMock() - db.scalar = AsyncMock() - db.execute = AsyncMock() - db.refresh = AsyncMock() + db.close = AsyncMock() return db +def _make_repo() -> MagicMock: + """构造 ``ChannelContentReviewRecordRepository`` 桩。""" + repo = MagicMock() + repo.create = AsyncMock() + repo.list = AsyncMock(return_value=[]) + repo.count = AsyncMock(return_value=0) + repo.get_by_review_id = AsyncMock(return_value=None) + repo.update_verdict = AsyncMock(return_value=None) + repo.query_analytics = AsyncMock() + repo.query_stats = AsyncMock() + repo.delete_old_records = AsyncMock(return_value=0) + return repo + + def _make_outcome() -> ContentModerationOutcome: """构造审核结论 DTO。""" return ContentModerationOutcome( @@ -134,32 +144,6 @@ def _make_analytics_query() -> ContentReviewAnalyticsQuery: ) -def _make_result_with_one(one_value: Any) -> MagicMock: - """构造 execute 返回值桩,使 ``.one()`` 返回指定聚合行。 - - analytics / stats 源码用 ``(await db.execute(stmt)).one()`` 访问聚合行, - execute 返回的 MagicMock 默认 ``.one()`` 返回新 MagicMock,无法读到 - 预置字段。本 helper 显式绑定 ``one`` 返回值。 - """ - result = MagicMock() - result.one.return_value = one_value - return result - - -def _make_scalar_result(scalar_value: Any) -> MagicMock: - """构造 execute 返回值桩,使 ``.scalar_one_or_none()`` 返回指定值。 - - get_by_review_id / update_verdict 源码用 - ``(await db.execute(stmt)).scalar_one_or_none()`` 访问单行, - execute 返回的 MagicMock 默认 ``.scalar_one_or_none()`` 返回新 - MagicMock(非 None),导致 not-found 分支无法触发。本 helper 显式 - 绑定 ``scalar_one_or_none`` 返回值。 - """ - result = MagicMock() - result.scalar_one_or_none.return_value = scalar_value - return result - - def _make_stats_query() -> ContentReviewStatsQuery: """构造审核统计查询条件。""" return ContentReviewStatsQuery( @@ -168,50 +152,80 @@ def _make_stats_query() -> ContentReviewStatsQuery: ) +_active_patchers: list = [] + + +def _build_adapter( + db: MagicMock, + repo: MagicMock, + logger: AsyncMock | None = None, +) -> ContentReviewRepositoryAdapter: + """构造 adapter,patch ``ChannelContentReviewRecordRepository`` 返回桩 repo。 + + patch 通过 ``start()`` 启动并在模块级 ``_active_patchers`` 列表中注册, + 由 autouse fixture ``_cleanup_patchers`` 在测试方法结束后统一 ``stop()``。 + 这是因为改造后 ``ChannelContentReviewRecordRepository(session)`` 在 + ``_session_scope(tx)`` 方法调用时执行(而非 ``__init__`` 构造时), + patch 需要跨越 ``_build_adapter`` 返回后继续生效。 + """ + patcher = patch( + "yuxi.channels.adapters.content_review_repository_adapter.ChannelContentReviewRecordRepository", + return_value=repo, + ) + patcher.start() + _active_patchers.append(patcher) + return ContentReviewRepositoryAdapter(lambda: db, logger or AsyncMock()) + + +@pytest.fixture(autouse=True) +def _cleanup_patchers(): + """每个测试结束后停止所有活跃的 patcher。""" + yield + for p in _active_patchers: + p.stop() + _active_patchers.clear() + + @pytest.mark.unit class TestContentReviewSaveReviewResult: @pytest.mark.asyncio async def test_save_review_result_commits_when_no_tx(self): # Arrange - # repo.create 在 commit=True 时调用 add + commit + refresh(不调用 flush) db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + adapter = _build_adapter(db, repo) record = _make_record() # Act await adapter.saveReviewResult(record) - # Assert - db.add.assert_called_once() - db.commit.assert_awaited_once() - db.refresh.assert_awaited_once() - db.flush.assert_not_awaited() + # Assert: tx 为 None 时 commit=True(自主提交) + repo.create.assert_awaited_once() + assert repo.create.call_args.kwargs["commit"] is True @pytest.mark.asyncio async def test_save_review_result_skips_commit_when_tx_provided(self): # Arrange db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + adapter = _build_adapter(db, repo) record = _make_record() tx = MagicMock() + tx.get_session.return_value = db # Act await adapter.saveReviewResult(record, tx=tx) - # Assert - db.flush.assert_awaited_once() - db.commit.assert_not_awaited() + # Assert: tx 非空时 commit=False(加入应用层事务) + assert repo.create.call_args.kwargs["commit"] is False @pytest.mark.asyncio async def test_save_review_result_translates_integrity_error_to_conflict(self): # Arrange - # repo.create 在 commit=True 时调用 commit(非 flush),IntegrityError 从 commit 抛出 db = _make_db() - db.commit = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig"))) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.create = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig"))) + adapter = _build_adapter(db, repo) record = _make_record() # Act / Assert @@ -222,11 +236,10 @@ class TestContentReviewSaveReviewResult: @pytest.mark.asyncio async def test_save_review_result_translates_sqlalchemy_error_to_dependency(self): # Arrange - # repo.create 在 commit=True 时调用 commit(非 flush),SQLAlchemyError 从 commit 抛出 db = _make_db() - db.commit = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.create = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) record = _make_record() # Act / Assert @@ -241,11 +254,9 @@ class TestContentReviewQueryHistory: async def test_query_history_returns_items_tuple(self): # Arrange db = _make_db() - result_mock = MagicMock() - result_mock.scalars.return_value.all.return_value = [_make_orm(review_id="rev-1"), _make_orm(review_id="rev-2")] - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.list = AsyncMock(return_value=[_make_orm(review_id="rev-1"), _make_orm(review_id="rev-2")]) + adapter = _build_adapter(db, repo) # Act items = await adapter.queryReviewHistory(_make_filter(), limit=10, offset=0) @@ -259,11 +270,9 @@ class TestContentReviewQueryHistory: async def test_query_history_returns_empty_tuple_when_no_records(self): # Arrange db = _make_db() - result_mock = MagicMock() - result_mock.scalars.return_value.all.return_value = [] - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.list = AsyncMock(return_value=[]) + adapter = _build_adapter(db, repo) # Act items = await adapter.queryReviewHistory(_make_filter(), limit=10, offset=0) @@ -275,9 +284,9 @@ class TestContentReviewQueryHistory: async def test_query_history_translates_sqlalchemy_error(self): # Arrange db = _make_db() - db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.list = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -285,19 +294,14 @@ class TestContentReviewQueryHistory: @pytest.mark.asyncio async def test_query_history_translates_invalid_enum_to_dependency(self): - # Arrange - # ORM 持有非法 verdict 值(迁移残留 / 数据损坏), + # Arrange: ORM 持有非法 verdict 值(迁移残留 / 数据损坏), # ``_enum`` 应翻译 ValueError 为 DependencyError,禁止穿透核心层(INV-7) - # 注:``ChannelType`` 是 ``str`` 子类而非 Enum,``ChannelType("unknown")`` - # 不抛异常;真正受 ``_enum`` 保护的是 verdict / source / resource_type db = _make_db() + repo = _make_repo() orm = _make_orm(review_id="rev-1") orm.verdict = "unknown_verdict" - result_mock = MagicMock() - result_mock.scalars.return_value.all.return_value = [orm] - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo.list = AsyncMock(return_value=[orm]) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -310,11 +314,9 @@ class TestContentReviewCountHistory: async def test_count_history_returns_int(self): # Arrange db = _make_db() - result_mock = MagicMock() - result_mock.scalar.return_value = 42 - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.count = AsyncMock(return_value=42) + adapter = _build_adapter(db, repo) # Act count = await adapter.countReviewHistory(_make_filter()) @@ -326,11 +328,9 @@ class TestContentReviewCountHistory: async def test_count_history_returns_zero_when_no_records(self): # Arrange db = _make_db() - result_mock = MagicMock() - result_mock.scalar.return_value = None - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.count = AsyncMock(return_value=0) + adapter = _build_adapter(db, repo) # Act count = await adapter.countReviewHistory(_make_filter()) @@ -342,9 +342,9 @@ class TestContentReviewCountHistory: async def test_count_history_translates_sqlalchemy_error(self): # Arrange db = _make_db() - db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.count = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -356,13 +356,10 @@ class TestContentReviewGetDetail: @pytest.mark.asyncio async def test_get_detail_returns_detail_when_found(self): # Arrange - # repo.get_by_review_id 用 ``(await db.execute(stmt)).scalar_one_or_none()`` - # 访问单行,故桩 db.execute 返回带 scalar_one_or_none 的结果对象 db = _make_db() - orm = _make_orm(review_id="rev-1") - db.execute = AsyncMock(return_value=_make_scalar_result(orm)) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.get_by_review_id = AsyncMock(return_value=_make_orm(review_id="rev-1")) + adapter = _build_adapter(db, repo) # Act detail = await adapter.getReviewDetail("rev-1") @@ -375,9 +372,9 @@ class TestContentReviewGetDetail: async def test_get_detail_returns_none_when_not_found(self): # Arrange db = _make_db() - db.execute = AsyncMock(return_value=_make_scalar_result(None)) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.get_by_review_id = AsyncMock(return_value=None) + adapter = _build_adapter(db, repo) # Act detail = await adapter.getReviewDetail("missing") @@ -389,9 +386,9 @@ class TestContentReviewGetDetail: async def test_get_detail_translates_sqlalchemy_error(self): # Arrange db = _make_db() - db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.get_by_review_id = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -399,15 +396,14 @@ class TestContentReviewGetDetail: @pytest.mark.asyncio async def test_get_detail_translates_invalid_verdict_to_dependency(self): - # Arrange - # ORM 持有非法 verdict 值(数据损坏),``_enum`` 应翻译 ValueError + # Arrange: ORM 持有非法 verdict 值(数据损坏),``_enum`` 应翻译 ValueError # 为 DependencyError,禁止穿透核心层(INV-7) db = _make_db() + repo = _make_repo() orm = _make_orm(review_id="rev-1") orm.verdict = "unknown_verdict" - db.execute = AsyncMock(return_value=_make_scalar_result(orm)) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo.get_by_review_id = AsyncMock(return_value=orm) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -418,31 +414,22 @@ class TestContentReviewGetDetail: class TestContentReviewGetAnalytics: @pytest.mark.asyncio async def test_get_analytics_returns_aggregated_result(self): - # Arrange - # 源码用 ``(await db.execute(agg_stmt)).one()`` 访问聚合行字段, - # 故 execute 须返回带 ``one()`` 方法的桩,``one()`` 返回预置聚合行 + # Arrange: repo.query_analytics 返回聚合 dict,适配器映射为 DTO db = _make_db() - agg_row = MagicMock() - agg_row.total_reviews = 10 - agg_row.block_count = 4 - category_result = MagicMock() - category_result.all.return_value = [ - MagicMock(category="politics", count=3), - MagicMock(category="violence", count=1), - ] - trend_result = MagicMock() - trend_result.all.return_value = [ - (datetime(2026, 1, 1), 5, 2), - ] - db.execute = AsyncMock( - side_effect=[ - _make_result_with_one(agg_row), - category_result, - trend_result, - ] + repo = _make_repo() + repo.query_analytics = AsyncMock( + return_value={ + "total_reviews": 10, + "block_count": 4, + "block_rate": 0.4, + "by_category": [ + {"category": "politics", "count": 3}, + {"category": "violence", "count": 1}, + ], + "trend": [(datetime(2026, 1, 1), 5, 2)], + } ) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + adapter = _build_adapter(db, repo) # Act result = await adapter.getReviewAnalytics(_make_analytics_query()) @@ -458,22 +445,17 @@ class TestContentReviewGetAnalytics: async def test_get_analytics_block_rate_zero_when_no_reviews(self): # Arrange db = _make_db() - agg_row = MagicMock() - agg_row.total_reviews = 0 - agg_row.block_count = 0 - category_result = MagicMock() - category_result.all.return_value = [] - trend_result = MagicMock() - trend_result.all.return_value = [] - db.execute = AsyncMock( - side_effect=[ - _make_result_with_one(agg_row), - category_result, - trend_result, - ] + repo = _make_repo() + repo.query_analytics = AsyncMock( + return_value={ + "total_reviews": 0, + "block_count": 0, + "block_rate": 0.0, + "by_category": [], + "trend": [], + } ) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + adapter = _build_adapter(db, repo) # Act result = await adapter.getReviewAnalytics(_make_analytics_query()) @@ -486,9 +468,9 @@ class TestContentReviewGetAnalytics: async def test_get_analytics_translates_sqlalchemy_error(self): # Arrange db = _make_db() - db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.query_analytics = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -499,30 +481,24 @@ class TestContentReviewGetAnalytics: class TestContentReviewGetStats: @pytest.mark.asyncio async def test_get_stats_returns_aggregated_result(self): - # Arrange - # 源码用 ``(await db.execute(agg_stmt)).one()`` 访问聚合行字段, - # 故 execute 须返回带 ``one()`` 方法的桩,``one()`` 返回预置聚合行 + # Arrange: repo.query_stats 返回聚合 dict,适配器映射为 DTO db = _make_db() - agg_row = MagicMock() - agg_row.total_reviews = 20 - agg_row.pass_count = 12 - agg_row.review_count = 3 - agg_row.block_count = 5 - agg_row.manual_count = 2 - agg_row.avg_decision_seconds = 42.5 - category_result = MagicMock() - category_result.all.return_value = [MagicMock(category="spam", count=4)] - trend_result = MagicMock() - trend_result.all.return_value = [(datetime(2026, 1, 1), 12, 5)] - db.execute = AsyncMock( - side_effect=[ - _make_result_with_one(agg_row), - category_result, - trend_result, - ] + repo = _make_repo() + repo.query_stats = AsyncMock( + return_value={ + "total_reviews": 20, + "pass_count": 12, + "review_count": 3, + "block_count": 5, + "pass_rate": 0.6, + "block_rate": 0.25, + "manual_intervention_rate": 0.1, + "avg_decision_seconds": 42.5, + "by_category": [{"category": "spam", "count": 4}], + "trend": [{"timestamp": datetime(2026, 1, 1), "pass_count": 12, "block_count": 5}], + } ) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + adapter = _build_adapter(db, repo) # Act result = await adapter.getReviewStats(_make_stats_query()) @@ -540,26 +516,22 @@ class TestContentReviewGetStats: async def test_get_stats_rates_zero_when_no_reviews(self): # Arrange db = _make_db() - agg_row = MagicMock() - agg_row.total_reviews = 0 - agg_row.pass_count = 0 - agg_row.review_count = 0 - agg_row.block_count = 0 - agg_row.manual_count = 0 - agg_row.avg_decision_seconds = None - category_result = MagicMock() - category_result.all.return_value = [] - trend_result = MagicMock() - trend_result.all.return_value = [] - db.execute = AsyncMock( - side_effect=[ - _make_result_with_one(agg_row), - category_result, - trend_result, - ] + repo = _make_repo() + repo.query_stats = AsyncMock( + return_value={ + "total_reviews": 0, + "pass_count": 0, + "review_count": 0, + "block_count": 0, + "pass_rate": 0.0, + "block_rate": 0.0, + "manual_intervention_rate": 0.0, + "avg_decision_seconds": None, + "by_category": [], + "trend": [], + } ) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + adapter = _build_adapter(db, repo) # Act result = await adapter.getReviewStats(_make_stats_query()) @@ -573,9 +545,9 @@ class TestContentReviewGetStats: async def test_get_stats_translates_sqlalchemy_error(self): # Arrange db = _make_db() - db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.query_stats = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -588,69 +560,56 @@ class TestContentReviewUpdateVerdict: async def test_update_verdict_returns_detail_when_found(self): # Arrange db = _make_db() - orm = _make_orm(review_id="rev-1") - result_mock = MagicMock() - result_mock.scalar_one_or_none.return_value = orm - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.update_verdict = AsyncMock(return_value=_make_orm(review_id="rev-1")) + adapter = _build_adapter(db, repo) # Act detail = await adapter.updateReviewVerdict("rev-1", "pass", "admin-1") # Assert - # repo.update_verdict 在 commit=True 时调用 commit + refresh(不调用 flush) assert detail is not None assert detail.review_id == "rev-1" - assert orm.verdict == "pass" - assert orm.reviewer == "admin-1" - db.commit.assert_awaited_once() - db.refresh.assert_awaited_once() - db.flush.assert_not_awaited() + repo.update_verdict.assert_awaited_once() + assert repo.update_verdict.call_args.kwargs["commit"] is True @pytest.mark.asyncio async def test_update_verdict_returns_none_when_not_found(self): # Arrange db = _make_db() - result_mock = MagicMock() - result_mock.scalar_one_or_none.return_value = None - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.update_verdict = AsyncMock(return_value=None) + adapter = _build_adapter(db, repo) # Act detail = await adapter.updateReviewVerdict("missing", "pass", "admin-1") # Assert assert detail is None - db.rollback.assert_awaited_once() @pytest.mark.asyncio async def test_update_verdict_skips_commit_when_tx_provided(self): # Arrange db = _make_db() - orm = _make_orm(review_id="rev-1") - result_mock = MagicMock() - result_mock.scalar_one_or_none.return_value = orm - db.execute = AsyncMock(return_value=result_mock) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.update_verdict = AsyncMock(return_value=_make_orm(review_id="rev-1")) + adapter = _build_adapter(db, repo) tx = MagicMock() + tx.get_session.return_value = db # Act await adapter.updateReviewVerdict("rev-1", "pass", "admin-1", tx=tx) # Assert - db.flush.assert_awaited_once() - db.commit.assert_not_awaited() + assert repo.update_verdict.call_args.kwargs["commit"] is False @pytest.mark.asyncio async def test_update_verdict_translates_sqlalchemy_error(self): # Arrange db = _make_db() - db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) + repo = _make_repo() + repo.update_verdict = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) # Act / Assert with pytest.raises(DependencyError): @@ -666,66 +625,63 @@ class TestContentReviewRepositoryDeleteOldRecords: 与各类异常翻译(IntegrityError / SQLAlchemyError / ConflictError / DependencyError / 通用 Exception)。 - 使用 AsyncMock 替换 ``adapter._repo.delete_old_records``,不连接真实 - DB,聚焦适配器的参数归一化、委托与错误翻译职责。 + 通过 ``_build_adapter`` patch ``ChannelContentReviewRecordRepository`` + 返回 mock 仓储,不连接真实 DB,聚焦适配器的参数归一化、委托与 + 错误翻译职责。 """ @pytest.mark.asyncio async def test_delete_old_records_returns_count_on_success(self): # 正常路径:repo.delete_old_records 返回删除数量,适配器原样透传 db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(return_value=42) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(return_value=42) + adapter = _build_adapter(db, repo) result = await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0)) assert result == 42 - delete_mock.assert_awaited_once() + repo.delete_old_records.assert_awaited_once() db.rollback.assert_not_awaited() @pytest.mark.asyncio async def test_delete_old_records_uses_default_limit_1000(self): # limit 未传时默认 1000,应原样传递给 repo db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(return_value=0) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(return_value=0) + adapter = _build_adapter(db, repo) before = datetime(2026, 1, 1, 0, 0, 0) await adapter.deleteOldReviewRecords(before) - delete_mock.assert_awaited_once_with(before, limit=1000, commit=True) + repo.delete_old_records.assert_awaited_once_with(before, limit=1000, commit=True) @pytest.mark.asyncio async def test_delete_old_records_passes_custom_limit(self): # 自定义 limit 应透传给 repo db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(return_value=0) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(return_value=0) + adapter = _build_adapter(db, repo) before = datetime(2026, 1, 1, 0, 0, 0) await adapter.deleteOldReviewRecords(before, limit=500) - delete_mock.assert_awaited_once_with(before, limit=500, commit=True) + repo.delete_old_records.assert_awaited_once_with(before, limit=500, commit=True) @pytest.mark.asyncio async def test_delete_old_records_falls_back_when_before_is_none(self): # before 为 None 时,_to_naive_utc(None) or None = None, # 兜底传递 None 给 repo(适配器不因 None 崩溃) db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(return_value=0) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(return_value=0) + adapter = _build_adapter(db, repo) await adapter.deleteOldReviewRecords(None) # type: ignore[arg-type] - delete_mock.assert_awaited_once_with(None, limit=1000, commit=True) + repo.delete_old_records.assert_awaited_once_with(None, limit=1000, commit=True) @pytest.mark.asyncio async def test_delete_old_records_normalizes_aware_before_to_naive_utc(self): @@ -733,26 +689,22 @@ class TestContentReviewRepositoryDeleteOldRecords: # _to_naive_utc 应归一化为 naive UTC 后再传递给 repo # Shanghai 2026-01-01 08:00:00 (+08:00) → UTC 2026-01-01 00:00:00 (naive) db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(return_value=0) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(return_value=0) + adapter = _build_adapter(db, repo) aware = datetime(2026, 1, 1, 8, 0, 0, tzinfo=ZoneInfo("Asia/Shanghai")) await adapter.deleteOldReviewRecords(aware) - delete_mock.assert_awaited_once_with( - datetime(2026, 1, 1, 0, 0, 0), limit=1000, commit=True - ) + repo.delete_old_records.assert_awaited_once_with(datetime(2026, 1, 1, 0, 0, 0), limit=1000, commit=True) @pytest.mark.asyncio async def test_delete_old_records_translates_integrity_error_to_conflict(self): # IntegrityError → ConflictError,并回滚事务 db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig"))) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig"))) + adapter = _build_adapter(db, repo) with pytest.raises(ConflictError): await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0)) @@ -762,10 +714,9 @@ class TestContentReviewRepositoryDeleteOldRecords: async def test_delete_old_records_translates_sqlalchemy_error_to_dependency(self): # SQLAlchemyError → DependencyError,并回滚事务 db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(side_effect=SQLAlchemyError("db failure")) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(side_effect=SQLAlchemyError("db failure")) + adapter = _build_adapter(db, repo) with pytest.raises(DependencyError): await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0)) @@ -775,10 +726,9 @@ class TestContentReviewRepositoryDeleteOldRecords: async def test_delete_old_records_reraises_conflict_error(self): # 契约层 ConflictError 原样重抛,不二次翻译,并回滚事务 db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(side_effect=ConflictError("content_review_record")) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(side_effect=ConflictError("content_review_record")) + adapter = _build_adapter(db, repo) with pytest.raises(ConflictError): await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0)) @@ -788,14 +738,11 @@ class TestContentReviewRepositoryDeleteOldRecords: async def test_delete_old_records_reraises_dependency_error(self): # 契约层 DependencyError 原样重抛,不二次翻译,并回滚事务 db = _make_db() - logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock( - side_effect=DependencyError( - "content_review_repository", Error("upstream failure") - ) + repo = _make_repo() + repo.delete_old_records = AsyncMock( + side_effect=DependencyError("content_review_repository", Error("upstream failure")) ) - adapter._repo.delete_old_records = delete_mock + adapter = _build_adapter(db, repo) with pytest.raises(DependencyError): await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0)) @@ -806,9 +753,9 @@ class TestContentReviewRepositoryDeleteOldRecords: # 通用 Exception → DependencyError 包装,记录日志并回滚事务 db = _make_db() logger = AsyncMock() - adapter = ContentReviewRepositoryAdapter(db, logger) - delete_mock = AsyncMock(side_effect=RuntimeError("unexpected boom")) - adapter._repo.delete_old_records = delete_mock + repo = _make_repo() + repo.delete_old_records = AsyncMock(side_effect=RuntimeError("unexpected boom")) + adapter = _build_adapter(db, repo, logger=logger) with pytest.raises(DependencyError): await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0)) diff --git a/backend/test/unit/channels/adapters/test_mappers.py b/backend/test/unit/channels/adapters/test_mappers.py index bc8ce931..946be4e3 100644 --- a/backend/test/unit/channels/adapters/test_mappers.py +++ b/backend/test/unit/channels/adapters/test_mappers.py @@ -9,7 +9,7 @@ from __future__ import annotations -from datetime import date, datetime +from datetime import UTC, date, datetime from unittest.mock import MagicMock, patch import pytest @@ -314,6 +314,22 @@ class TestOrmToAuditLog: assert result.target_channel == "feishu" +@pytest.mark.unit +class TestToAwareUtc: + def test_converts_naive_datetime_to_aware_utc(self): + naive = datetime(2026, 1, 1, 12, 0, 0) + result = mappers._to_aware_utc(naive) + assert result == datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC) + assert result.tzinfo is UTC + + def test_returns_aware_datetime_unchanged(self): + aware = datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC) + assert mappers._to_aware_utc(aware) is aware + + def test_returns_none_unchanged(self): + assert mappers._to_aware_utc(None) is None + + @pytest.mark.unit class TestOrmToOutboxEntry: def test_returns_outbox_entry_with_str_ids(self): @@ -332,6 +348,25 @@ class TestOrmToOutboxEntry: assert result.status == OutboxStatus.PENDING assert result.durability_policy == MessageDurabilityPolicy.REQUIRED + def test_datetime_fields_are_aware_utc(self): + """DB 读出的 naive datetime 必须转为 aware UTC,避免领域层时区混用。""" + orm = _make_outbox_orm() + orm.created_at = datetime(2026, 1, 1) + orm.updated_at = datetime(2026, 1, 2) + orm.expires_at = datetime(2026, 1, 3) + orm.sent_at = datetime(2026, 1, 4) + orm.last_retry_at = datetime(2026, 1, 5) + orm.next_retry_at = datetime(2026, 1, 6) + + result = mappers.orm_to_outbox_entry(orm, _make_account_orm()) + + assert result.created_at.tzinfo is UTC + assert result.updated_at.tzinfo is UTC + assert result.expires_at.tzinfo is UTC + assert result.sent_at.tzinfo is UTC + assert result.last_retry_at.tzinfo is UTC + assert result.next_retry_at.tzinfo is UTC + def test_handles_none_channel_session_id(self): # Arrange orm = _make_outbox_orm() @@ -343,6 +378,21 @@ class TestOrmToOutboxEntry: # Assert assert result.channel_session_id is None + def test_handles_none_datetime_fields(self): + """可空 datetime 字段为 None 时不应报错。""" + orm = _make_outbox_orm() + orm.expires_at = None + orm.sent_at = None + orm.last_retry_at = None + orm.next_retry_at = None + + result = mappers.orm_to_outbox_entry(orm, _make_account_orm()) + + assert result.expires_at is None + assert result.sent_at is None + assert result.last_retry_at is None + assert result.next_retry_at is None + @pytest.mark.unit class TestOrmToUserIdentity: diff --git a/backend/test/unit/channels/adapters/test_route_binding_repository.py b/backend/test/unit/channels/adapters/test_route_binding_repository.py index 43d7af47..068d388d 100644 --- a/backend/test/unit/channels/adapters/test_route_binding_repository.py +++ b/backend/test/unit/channels/adapters/test_route_binding_repository.py @@ -90,13 +90,34 @@ def _make_binding_orm( return orm +_active_patchers: list = [] + + def _build_adapter(db: MagicMock, repos: MagicMock) -> ChannelPersistenceAdapter: - """构造 adapter,patch ``create_repositories`` 返回桩 repos。""" - with patch( + """构造 adapter,patch ``create_repositories`` 返回桩 repos。 + + patch 通过 ``start()`` 启动并在模块级 ``_active_patchers`` 列表中注册, + 由 autouse fixture ``_cleanup_patchers`` 在测试方法结束后统一 ``stop()``。 + 这是因为改造后 ``create_repositories`` 在 ``_session_scope(tx)`` 方法调用 + 时执行(而非 ``__init__`` 构造时),patch 需要跨越 ``_build_adapter`` + 返回后继续生效。 + """ + patcher = patch( "yuxi.channels.adapters.channel_persistence_adapter.create_repositories", return_value=repos, - ): - return ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=MagicMock()) + ) + patcher.start() + _active_patchers.append(patcher) + return ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=MagicMock()) + + +@pytest.fixture(autouse=True) +def _cleanup_patchers(): + """每个测试结束后停止所有活跃的 patcher。""" + yield + for p in _active_patchers: + p.stop() + _active_patchers.clear() def _make_operator() -> Operator: diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_audit_log_retention_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_audit_log_retention_handler.py index c5402e02..fb65f59e 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_audit_log_retention_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_audit_log_retention_handler.py @@ -3,14 +3,13 @@ 覆盖 ``yuxi.channels.application.extension.scheduler_handlers.channel_audit_log_retention_handler``: - execute 正常路径:无记录 / 有记录删除 - payload 覆盖 retention_days 与 batch_size 默认值 -- session / repository 异常转换为 TaskResult(success=False) +- repository 异常转换为 TaskResult(success=False) """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime, timedelta -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_audit_log_retention_handler import ( @@ -38,15 +37,8 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - def _make_handler(*, audit_log_repo=None, logger=None, cache_port=None): """构造 ChannelAuditLogRetentionHandler 及其依赖桩。""" - db = MagicMock() if audit_log_repo is None: audit_log_repo = AsyncMock() audit_log_repo.deleteOldAuditLogs.return_value = 0 @@ -58,13 +50,11 @@ def _make_handler(*, audit_log_repo=None, logger=None, cache_port=None): cache_port.releaseAdvisoryLock.return_value = True handler = ChannelAuditLogRetentionHandler( - session_factory=lambda: _fake_session_factory(db), - audit_log_repo_factory=lambda _db: audit_log_repo, + audit_log_repo=audit_log_repo, cache_port=cache_port, logger=logger, ) return handler, { - "db": db, "audit_log_repo": audit_log_repo, "cache_port": cache_port, "logger": logger, @@ -131,25 +121,6 @@ class TestHandlerExecute: @pytest.mark.unit class TestHandlerErrors: - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelAuditLogRetentionHandler( - session_factory=_broken_session_factory, - audit_log_repo_factory=lambda _db: AsyncMock(), - cache_port=fake_cache, - logger=fake_logger, - ) - - result = await handler.execute(_make_ctx()) - - assert result.success is False - assert "session creation failed" in (result.error or "") - @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger): handler, deps = _make_handler(logger=fake_logger) diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_content_review_retention_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_content_review_retention_handler.py index 22485dd5..55cc9eb3 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_content_review_retention_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_content_review_retention_handler.py @@ -3,14 +3,13 @@ 覆盖 ``yuxi.channels.application.extension.scheduler_handlers.channel_content_review_retention_handler``: - execute 正常路径:无记录 / 有记录删除 - payload 覆盖 retention_days 与 batch_size 默认值 -- session / repository 异常转换为 TaskResult(success=False) +- repository 异常转换为 TaskResult(success=False) """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime, timedelta -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_content_review_retention_handler import ( @@ -38,15 +37,8 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - def _make_handler(*, content_review_repo=None, logger=None, cache_port=None): """构造 ChannelContentReviewRetentionHandler 及其依赖桩。""" - db = MagicMock() if content_review_repo is None: content_review_repo = AsyncMock() content_review_repo.deleteOldReviewRecords.return_value = 0 @@ -58,13 +50,11 @@ def _make_handler(*, content_review_repo=None, logger=None, cache_port=None): cache_port.releaseAdvisoryLock.return_value = True handler = ChannelContentReviewRetentionHandler( - session_factory=lambda: _fake_session_factory(db), - content_review_repo_factory=lambda _db: content_review_repo, + content_review_repo=content_review_repo, cache_port=cache_port, logger=logger, ) return handler, { - "db": db, "content_review_repo": content_review_repo, "cache_port": cache_port, "logger": logger, @@ -132,25 +122,6 @@ class TestHandlerExecute: @pytest.mark.unit class TestHandlerErrors: - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelContentReviewRetentionHandler( - session_factory=_broken_session_factory, - content_review_repo_factory=lambda _db: AsyncMock(), - cache_port=fake_cache, - logger=fake_logger, - ) - - result = await handler.execute(_make_ctx()) - - assert result.success is False - assert "session creation failed" in (result.error or "") - @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger): handler, deps = _make_handler(logger=fake_logger) diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_idempotency_cleanup_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_idempotency_cleanup_handler.py index 57efbe39..d12c60f1 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_idempotency_cleanup_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_idempotency_cleanup_handler.py @@ -3,14 +3,13 @@ 覆盖 ``yuxi.channels.application.extension.scheduler_handlers.channel_idempotency_cleanup_handler``: - execute 正常路径:无记录 / 有记录删除 - deleteExpiredRecords 以 ``utc_now_naive`` 为截止时间 -- session / repository 异常转换为 TaskResult(success=False) +- repository 异常转换为 TaskResult(success=False) """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_idempotency_cleanup_handler import ( @@ -38,15 +37,8 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - def _make_handler(*, idempotency_repo=None, logger=None, cache_port=None): """构造 ChannelIdempotencyCleanupHandler 及其依赖桩。""" - db = MagicMock() if idempotency_repo is None: idempotency_repo = AsyncMock() idempotency_repo.deleteExpiredRecords.return_value = 0 @@ -58,13 +50,11 @@ def _make_handler(*, idempotency_repo=None, logger=None, cache_port=None): cache_port.releaseAdvisoryLock.return_value = True handler = ChannelIdempotencyCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - idempotency_repo_factory=lambda _db: idempotency_repo, + idempotency_repo=idempotency_repo, cache_port=cache_port, logger=logger, ) return handler, { - "db": db, "idempotency_repo": idempotency_repo, "cache_port": cache_port, "logger": logger, @@ -134,26 +124,6 @@ class TestHandlerExecute: @pytest.mark.unit class TestHandlerErrors: - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelIdempotencyCleanupHandler( - session_factory=_broken_session_factory, - idempotency_repo_factory=lambda _db: AsyncMock(), - cache_port=fake_cache, - logger=fake_logger, - ) - - result = await handler.execute(_make_ctx()) - - assert result.success is False - assert "session creation failed" in (result.error or "") - fake_logger.exception.assert_awaited_once() - @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger): handler, deps = _make_handler(logger=fake_logger) diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_recovery_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_recovery_handler.py index 525c54cd..42ebb93c 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_recovery_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_recovery_handler.py @@ -6,14 +6,13 @@ - 死信审计日志写入及写入失败不阻塞 - ARQ 入队失败仅记录日志 - 单条异常隔离 -- session / repository 异常转换为 TaskResult(success=False) +- repository 异常转换为 TaskResult(success=False) """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime, timedelta -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_outbox_recovery_handler import ( @@ -49,12 +48,6 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - def _make_outbox_config( ttl_seconds: int = 86400, max_retry: int = 5, @@ -110,7 +103,6 @@ def _make_handler( cache_port=None, ): """构造 ChannelOutboxRecoveryHandler 及其依赖桩。""" - db = MagicMock() if outbox_repo is None: outbox_repo = AsyncMock() outbox_repo.listPendingOutboxEntries.return_value = [] @@ -130,16 +122,14 @@ def _make_handler( cache_port.releaseAdvisoryLock.return_value = True handler = ChannelOutboxRecoveryHandler( - session_factory=lambda: _fake_session_factory(db), - outbox_repo_factory=lambda _db: outbox_repo, - audit_log_repo_factory=lambda _db: audit_log_repo, + outbox_repo=outbox_repo, + audit_log_repo=audit_log_repo, queue_port=queue_port, outbox_config=outbox_config, cache_port=cache_port, logger=logger, ) return handler, { - "db": db, "outbox_repo": outbox_repo, "audit_log_repo": audit_log_repo, "queue_port": queue_port, @@ -377,28 +367,6 @@ class TestHandlerErrors: assert result.success is True assert result.output["processed_count"] == 1 - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelOutboxRecoveryHandler( - session_factory=_broken_session_factory, - outbox_repo_factory=lambda _db: AsyncMock(), - audit_log_repo_factory=lambda _db: AsyncMock(), - queue_port=AsyncMock(), - outbox_config=_make_outbox_config(), - cache_port=fake_cache, - logger=fake_logger, - ) - - result = await handler.execute(_make_ctx()) - - assert result.success is False - assert "session creation failed" in (result.error or "") - @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger): handler, deps = _make_handler(logger=fake_logger) diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_terminal_cleanup_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_terminal_cleanup_handler.py index 05fbf3ff..c84565da 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_terminal_cleanup_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_outbox_terminal_cleanup_handler.py @@ -6,14 +6,13 @@ - payload 覆盖默认保留小时数与批量大小 - event_publisher 为 None 时跳过事件发布 - event_publisher 发布异常仅记录 WARN,不影响主流程 -- 异常隔离:仓储异常 / session_factory 异常转换为 TaskResult(success=False) +- 异常隔离:仓储异常转换为 TaskResult(success=False) """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_outbox_terminal_cleanup_handler import ( @@ -40,12 +39,6 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - # ─── 类属性 ─────────────────────────────────────────────────────────────── @@ -67,15 +60,13 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_successful_cleanup_returns_success(self, fake_logger, fake_cache): # Arrange - db = MagicMock() outbox_repo = AsyncMock() outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1", "dead-2"]) outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=["sent-1"]) event_publisher = AsyncMock() handler = ChannelOutboxTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - outbox_repo_factory=lambda _db: outbox_repo, + outbox_repo=outbox_repo, event_publisher=event_publisher, cache_port=fake_cache, logger=fake_logger, @@ -92,14 +83,12 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_payload_overrides_defaults(self, fake_logger, fake_cache): # Arrange - payload 指定 dead_retention_hours=24, sent_retention_hours=168, batch_size=50 - db = MagicMock() outbox_repo = AsyncMock() outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=[]) outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=[]) handler = ChannelOutboxTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - outbox_repo_factory=lambda _db: outbox_repo, + outbox_repo=outbox_repo, event_publisher=None, cache_port=fake_cache, logger=fake_logger, @@ -123,14 +112,12 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_event_publisher_none_skips_publish(self, fake_logger, fake_cache): # Arrange - event_publisher 为 None - db = MagicMock() outbox_repo = AsyncMock() outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1"]) outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=[]) handler = ChannelOutboxTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - outbox_repo_factory=lambda _db: outbox_repo, + outbox_repo=outbox_repo, event_publisher=None, cache_port=fake_cache, logger=fake_logger, @@ -146,7 +133,6 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_event_publish_failure_continues(self, fake_logger, fake_cache): # Arrange - event_publisher 发布异常仅记录 WARN,不影响主流程 - db = MagicMock() outbox_repo = AsyncMock() outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1"]) outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=["sent-1"]) @@ -154,8 +140,7 @@ class TestHandlerExecute: event_publisher.publishOutboxEntryPurged = AsyncMock(side_effect=RuntimeError("publish failed")) handler = ChannelOutboxTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - outbox_repo_factory=lambda _db: outbox_repo, + outbox_repo=outbox_repo, event_publisher=event_publisher, cache_port=fake_cache, logger=fake_logger, @@ -173,15 +158,13 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_event_published_with_purged_entry_ids(self, fake_logger, fake_cache): # Arrange - 验证发布的事件包含被清理条目的完整 ID 列表 - db = MagicMock() outbox_repo = AsyncMock() outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1"]) outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=["sent-1"]) event_publisher = AsyncMock() handler = ChannelOutboxTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - outbox_repo_factory=lambda _db: outbox_repo, + outbox_repo=outbox_repo, event_publisher=event_publisher, cache_port=fake_cache, logger=fake_logger, @@ -206,13 +189,11 @@ class TestHandlerErrors: @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger, fake_cache): # Arrange - 仓储抛异常 - db = MagicMock() outbox_repo = AsyncMock() outbox_repo.cleanupOldDeadEntries = AsyncMock(side_effect=RuntimeError("db error")) handler = ChannelOutboxTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - outbox_repo_factory=lambda _db: outbox_repo, + outbox_repo=outbox_repo, event_publisher=None, cache_port=fake_cache, logger=fake_logger, @@ -224,27 +205,3 @@ class TestHandlerErrors: # Assert assert result.success is False assert "db error" in (result.error or "") - - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - # Arrange - session_factory 自身抛异常 - - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelOutboxTerminalCleanupHandler( - session_factory=_broken_session_factory, - outbox_repo_factory=lambda _db: AsyncMock(), - event_publisher=None, - cache_port=fake_cache, - logger=fake_logger, - ) - - # Act - result = await handler.execute(_make_ctx()) - - # Assert - assert result.success is False - assert "session creation failed" in (result.error or "") diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_expiration_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_expiration_handler.py index d86a5d71..95c2b5ef 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_expiration_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_expiration_handler.py @@ -4,15 +4,14 @@ - execute 正常路径:空扫描 / 过期标记 / 未过期跳过 / 混合记录 - payload 覆盖 batch_size 默认值 - 单条记录异常隔离 -- session / repository 异常转换为 TaskResult(success=False) +- repository 异常转换为 TaskResult(success=False) - fromRecord 使用记录自身 version """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime, timedelta -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_pairing_expiration_handler import ( @@ -43,12 +42,6 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - def _make_pairing_record( *, pairing_id: str = "pr-001", @@ -73,7 +66,6 @@ def _make_pairing_record( def _make_handler(*, pairing_repo=None, logger=None, cache_port=None): """构造 ChannelPairingExpirationHandler 及其依赖桩。""" - db = MagicMock() if pairing_repo is None: pairing_repo = AsyncMock() pairing_repo.listExpiredPendingPairings.return_value = () @@ -85,13 +77,11 @@ def _make_handler(*, pairing_repo=None, logger=None, cache_port=None): cache_port.releaseAdvisoryLock.return_value = True handler = ChannelPairingExpirationHandler( - session_factory=lambda: _fake_session_factory(db), - pairing_repo_factory=lambda _db: pairing_repo, + pairing_repo=pairing_repo, cache_port=cache_port, logger=logger, ) return handler, { - "db": db, "pairing_repo": pairing_repo, "cache_port": cache_port, "logger": logger, @@ -222,25 +212,6 @@ class TestHandlerErrors: assert result.output["scanned_count"] == 2 deps["logger"].exception.assert_awaited_once() - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelPairingExpirationHandler( - session_factory=_broken_session_factory, - pairing_repo_factory=lambda _db: AsyncMock(), - cache_port=fake_cache, - logger=fake_logger, - ) - - result = await handler.execute(_make_ctx()) - - assert result.success is False - assert "session creation failed" in (result.error or "") - @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger): handler, deps = _make_handler(logger=fake_logger) diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_terminal_cleanup_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_terminal_cleanup_handler.py index fd974e27..4a21dbed 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_terminal_cleanup_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_pairing_terminal_cleanup_handler.py @@ -4,14 +4,13 @@ 模块的 ``ChannelPairingTerminalCleanupHandler``: - execute 正常路径:调用 cleanupOldTerminalPairings 并返回成功结果 - payload 覆盖默认保留天数与批量大小 -- 异常隔离:仓储异常 / session_factory 异常转换为 TaskResult(success=False) +- 异常隔离:仓储异常转换为 TaskResult(success=False) """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_pairing_terminal_cleanup_handler import ( @@ -37,12 +36,6 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - # ─── 类属性 ─────────────────────────────────────────────────────────────── @@ -64,13 +57,11 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_successful_cleanup_returns_success(self, fake_logger, fake_cache): # Arrange - db = MagicMock() pairing_repo = AsyncMock() pairing_repo.cleanupOldTerminalPairings = AsyncMock(return_value=42) handler = ChannelPairingTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - pairing_repo_factory=lambda _db: pairing_repo, + pairing_repo=pairing_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -86,13 +77,11 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_payload_overrides_defaults(self, fake_logger, fake_cache): # Arrange - payload 指定 retention_days=7, batch_size=10 - db = MagicMock() pairing_repo = AsyncMock() pairing_repo.cleanupOldTerminalPairings = AsyncMock(return_value=5) handler = ChannelPairingTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - pairing_repo_factory=lambda _db: pairing_repo, + pairing_repo=pairing_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -109,13 +98,11 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_default_retention_days_used_when_payload_missing(self, fake_logger, fake_cache): # Arrange - payload 不含 retention_days 与 batch_size - db = MagicMock() pairing_repo = AsyncMock() pairing_repo.cleanupOldTerminalPairings = AsyncMock(return_value=0) handler = ChannelPairingTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - pairing_repo_factory=lambda _db: pairing_repo, + pairing_repo=pairing_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -137,13 +124,11 @@ class TestHandlerErrors: @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger, fake_cache): # Arrange - 仓储抛异常 - db = MagicMock() pairing_repo = AsyncMock() pairing_repo.cleanupOldTerminalPairings = AsyncMock(side_effect=RuntimeError("db error")) handler = ChannelPairingTerminalCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - pairing_repo_factory=lambda _db: pairing_repo, + pairing_repo=pairing_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -154,26 +139,3 @@ class TestHandlerErrors: # Assert - 异常被捕获并转换为 TaskResult(success=False) assert result.success is False assert "db error" in (result.error or "") - - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - # Arrange - session_factory 自身抛异常 - - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelPairingTerminalCleanupHandler( - session_factory=_broken_session_factory, - pairing_repo_factory=lambda _db: AsyncMock(), - cache_port=fake_cache, - logger=fake_logger, - ) - - # Act - result = await handler.execute(_make_ctx()) - - # Assert - assert result.success is False - assert "session creation failed" in (result.error or "") diff --git a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_session_inactive_cleanup_handler.py b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_session_inactive_cleanup_handler.py index ac48454a..2ae790f1 100644 --- a/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_session_inactive_cleanup_handler.py +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_session_inactive_cleanup_handler.py @@ -7,12 +7,11 @@ - payload 覆盖默认阈值与批量大小 - realtime_metrics 非空时调用 releaseSession - realtime_metrics 为 None 时跳过释放 -- 异常隔离:仓储异常 / session_factory 异常转换为 TaskResult(success=False) +- 异常隔离:仓储异常转换为 TaskResult(success=False) """ from __future__ import annotations -from contextlib import asynccontextmanager from datetime import datetime from unittest.mock import AsyncMock, MagicMock @@ -40,12 +39,6 @@ def _make_ctx(**payload_overrides) -> TaskContext: ) -@asynccontextmanager -async def _fake_session_factory(db_mock): - """构造假的 session_factory,yield db_mock。""" - yield db_mock - - def _make_session(session_id: str = "sess-001"): """构造带 session_id 属性的假会话对象。""" session = MagicMock() @@ -74,14 +67,12 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_successful_cleanup_returns_success(self, fake_logger, fake_cache): # Arrange - db = MagicMock() session_repo = AsyncMock() session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[_make_session("s1"), _make_session("s2")]) session_repo.cleanupInactiveSessions = AsyncMock(return_value=2) handler = ChannelSessionInactiveCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - session_repo_factory=lambda _db: session_repo, + session_repo=session_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -98,13 +89,11 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_no_inactive_sessions_returns_zero(self, fake_logger, fake_cache): # Arrange - 无非活跃会话时直接返回 - db = MagicMock() session_repo = AsyncMock() session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[]) handler = ChannelSessionInactiveCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - session_repo_factory=lambda _db: session_repo, + session_repo=session_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -120,14 +109,12 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_payload_overrides_defaults(self, fake_logger, fake_cache): # Arrange - payload 指定 inactive_threshold_minutes=30, batch_size=50 - db = MagicMock() session_repo = AsyncMock() session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[]) session_repo.cleanupInactiveSessions = AsyncMock(return_value=0) handler = ChannelSessionInactiveCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - session_repo_factory=lambda _db: session_repo, + session_repo=session_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -143,15 +130,13 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_realtime_metrics_release_called(self, fake_logger, fake_cache): # Arrange - realtime_metrics 非空时对每个 session_id 调用 releaseSession - db = MagicMock() session_repo = AsyncMock() session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[_make_session("s1"), _make_session("s2")]) session_repo.cleanupInactiveSessions = AsyncMock(return_value=2) realtime_metrics = AsyncMock() handler = ChannelSessionInactiveCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - session_repo_factory=lambda _db: session_repo, + session_repo=session_repo, cache_port=fake_cache, logger=fake_logger, realtime_metrics=realtime_metrics, @@ -169,14 +154,12 @@ class TestHandlerExecute: @pytest.mark.asyncio async def test_realtime_metrics_none_skips_release(self, fake_logger, fake_cache): # Arrange - realtime_metrics 为 None 时不调用 releaseSession - db = MagicMock() session_repo = AsyncMock() session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[_make_session("s1")]) session_repo.cleanupInactiveSessions = AsyncMock(return_value=1) handler = ChannelSessionInactiveCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - session_repo_factory=lambda _db: session_repo, + session_repo=session_repo, cache_port=fake_cache, logger=fake_logger, realtime_metrics=None, @@ -198,13 +181,11 @@ class TestHandlerErrors: @pytest.mark.asyncio async def test_repository_exception_returns_failure(self, fake_logger, fake_cache): # Arrange - 仓储抛异常 - db = MagicMock() session_repo = AsyncMock() session_repo.listInactiveTemporarySessions = AsyncMock(side_effect=RuntimeError("db error")) handler = ChannelSessionInactiveCleanupHandler( - session_factory=lambda: _fake_session_factory(db), - session_repo_factory=lambda _db: session_repo, + session_repo=session_repo, cache_port=fake_cache, logger=fake_logger, ) @@ -215,26 +196,3 @@ class TestHandlerErrors: # Assert assert result.success is False assert "db error" in (result.error or "") - - @pytest.mark.asyncio - async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache): - # Arrange - session_factory 自身抛异常 - - @asynccontextmanager - async def _broken_session_factory(): - raise RuntimeError("session creation failed") - yield # pragma: no cover - - handler = ChannelSessionInactiveCleanupHandler( - session_factory=_broken_session_factory, - session_repo_factory=lambda _db: AsyncMock(), - cache_port=fake_cache, - logger=fake_logger, - ) - - # Act - result = await handler.execute(_make_ctx()) - - # Assert - assert result.success is False - assert "session creation failed" in (result.error or "") diff --git a/backend/test/unit/channels/application/pipeline/outbound/test_load_build_stage.py b/backend/test/unit/channels/application/pipeline/outbound/test_load_build_stage.py index bd2a230e..2d670d34 100644 --- a/backend/test/unit/channels/application/pipeline/outbound/test_load_build_stage.py +++ b/backend/test/unit/channels/application/pipeline/outbound/test_load_build_stage.py @@ -3,15 +3,21 @@ 覆盖 ``LoadBuildStage``: - ``process(context)``:正常路径(含富消息/附件)、无富消息、空分块、 image_url/video_url 转换为附件 +- 持久化模式下 AgentRun 内容加载(streamAgentRun 阻塞消费 + getAgentRunFinalOutput) """ from __future__ import annotations +from collections.abc import AsyncIterator +from unittest.mock import AsyncMock, MagicMock + import pytest from yuxi.channels.application.context.outbound_context import OutboundContext from yuxi.channels.application.pipeline.outbound.load_build_stage import LoadBuildStage +from yuxi.channels.contract.dtos.agent_run import AgentRunId from yuxi.channels.contract.dtos.channel import ChannelType from yuxi.channels.contract.dtos.common import Attachment +from yuxi.channels.contract.dtos.option import Nothing, Some from yuxi.channels.contract.dtos.outbound import RichMessage from yuxi.channels.contract.plugin.extension_point import FailureStrategy @@ -33,12 +39,36 @@ def _make_ctx(**overrides) -> OutboundContext: return OutboundContext(**defaults) +def _make_agent_run_port( + *, + stream_events: list | None = None, + final_output: Some | Nothing = None, +) -> MagicMock: + """构造 mock AgentRunPort。 + + stream_events:streamAgentRun 产出的事件列表(按序 yield 后结束迭代)。 + final_output:getAgentRunFinalOutput 返回的 Option。 + """ + port = MagicMock() + + async def _stream(_run_id: AgentRunId) -> AsyncIterator: + for event in stream_events or []: + yield event + + port.streamAgentRun = MagicMock(side_effect=_stream) + port.getAgentRunFinalOutput = AsyncMock(return_value=final_output or Nothing()) + return port + + +# ─── process 正常路径 ──────────────────────────────────────────────────────── + + @pytest.mark.unit class TestLoadBuildStageProcess: @pytest.mark.asyncio async def test_normal_path_builds_payload_with_rich_message(self): # Arrange - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) rich = RichMessage(text="hello", image_url="http://img", video_url="http://vid") ctx = _make_ctx( stream_chunks=["chunk-a", "chunk-b"], @@ -65,7 +95,7 @@ class TestLoadBuildStageProcess: @pytest.mark.asyncio async def test_no_rich_message_sets_fields_none(self): # Arrange - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) ctx = _make_ctx(stream_chunks=["x"], rich_message=None) # Act @@ -78,10 +108,10 @@ class TestLoadBuildStageProcess: assert ctx.outbound_payload.attachments == () @pytest.mark.asyncio - async def test_empty_stream_chunks(self): - # Arrange - stage = LoadBuildStage() - ctx = _make_ctx(stream_chunks=[]) + async def test_empty_stream_chunks_without_agent_run_id(self): + # Arrange — agent_run_id 为空时不触发内容加载,stream_chunks 保持空 + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) + ctx = _make_ctx(agent_run_id="", stream_chunks=[]) # Act ok = await stage.process(ctx) @@ -92,10 +122,10 @@ class TestLoadBuildStageProcess: @pytest.mark.asyncio async def test_rich_message_without_media_urls(self): - # Arrange:rich_message 无 image_url / video_url,附件仅含原始 - stage = LoadBuildStage() + # Arrange + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) rich = RichMessage(text="text-only") - ctx = _make_ctx(rich_message=rich) + ctx = _make_ctx(stream_chunks=["x"], rich_message=rich) # Act ok = await stage.process(ctx) @@ -119,32 +149,38 @@ class TestLoadBuildStageContract: """阶段契约属性(StageContract)。""" def test_reads_agent_run_and_payload_inputs(self): - stage = LoadBuildStage() - assert stage.reads == ("agent_run_id", "stream_chunks", "rich_message", "attachments") + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) + assert stage.reads == ( + "agent_run_id", + "stream_chunks", + "rich_message", + "attachments", + "delivery_mode", + ) - def test_writes_outbound_payload(self): - stage = LoadBuildStage() - assert stage.writes == ("outbound_payload",) + def test_writes_outbound_payload_and_stream_chunks(self): + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) + assert stage.writes == ("outbound_payload", "stream_chunks") def test_is_idempotent(self): - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) assert stage.idempotent is True def test_is_thread_safe(self): - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) assert stage.thread_safe is True def test_failure_strategy_is_terminate(self): # failure=TERMINATE:装配失败时终止管道(INV-8 agent_run_id 非空约束) - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) assert stage.failure is FailureStrategy.TERMINATE def test_compensate_is_none(self): - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) assert stage.compensate is None def test_condition_is_none(self): - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) assert stage.condition is None @@ -158,9 +194,9 @@ class TestLoadBuildStageEdgeCases: @pytest.mark.asyncio async def test_image_only_attachment_appended(self): # Arrange — rich_message 仅含 image_url,应转换为 image 附件并入 - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) rich = RichMessage(text="img-only", image_url="http://img") - ctx = _make_ctx(rich_message=rich) + ctx = _make_ctx(stream_chunks=["x"], rich_message=rich) # Act await stage.process(ctx) @@ -174,9 +210,9 @@ class TestLoadBuildStageEdgeCases: @pytest.mark.asyncio async def test_video_only_attachment_appended(self): # Arrange — rich_message 仅含 video_url,应转换为 video 附件并入 - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) rich = RichMessage(text="vid-only", video_url="http://vid") - ctx = _make_ctx(rich_message=rich) + ctx = _make_ctx(stream_chunks=["x"], rich_message=rich) # Act await stage.process(ctx) @@ -190,13 +226,13 @@ class TestLoadBuildStageEdgeCases: @pytest.mark.asyncio async def test_image_before_video_in_attachment_order(self): # Arrange — image_url 与 video_url 同时存在时,image 在前 video 在后 - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) rich = RichMessage( text="both", image_url="http://img", video_url="http://vid", ) - ctx = _make_ctx(rich_message=rich) + ctx = _make_ctx(stream_chunks=["x"], rich_message=rich) # Act await stage.process(ctx) @@ -208,7 +244,7 @@ class TestLoadBuildStageEdgeCases: @pytest.mark.asyncio async def test_stream_chunks_list_converted_to_tuple(self): # Arrange — list 输入应转为 tuple 以保证 OutboundPayload 不可变语义 - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) ctx = _make_ctx(stream_chunks=["a", "b", "c"]) # Act @@ -221,8 +257,9 @@ class TestLoadBuildStageEdgeCases: @pytest.mark.asyncio async def test_attachments_tuple_type_preserved(self): # Arrange — 原始 attachments 为 tuple,输出仍为 tuple - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) ctx = _make_ctx( + stream_chunks=["x"], attachments=(Attachment(type="file", url="http://file"),), ) @@ -235,9 +272,10 @@ class TestLoadBuildStageEdgeCases: @pytest.mark.asyncio async def test_original_attachments_prepend_media_attachments(self): # Arrange — 原始附件在前,image/video 在后 - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) rich = RichMessage(text="x", image_url="http://img", video_url="http://vid") ctx = _make_ctx( + stream_chunks=["x"], rich_message=rich, attachments=(Attachment(type="file", url="http://file"),), ) @@ -252,9 +290,9 @@ class TestLoadBuildStageEdgeCases: @pytest.mark.asyncio async def test_rich_message_fields_wraps_rich_message(self): # Arrange — rich_message 应被 RichMessageFields 包装 - stage = LoadBuildStage() + stage = LoadBuildStage(agent_run_port=_make_agent_run_port()) rich = RichMessage(text="wrap") - ctx = _make_ctx(rich_message=rich) + ctx = _make_ctx(stream_chunks=["x"], rich_message=rich) # Act await stage.process(ctx) @@ -262,3 +300,102 @@ class TestLoadBuildStageEdgeCases: # Assert assert ctx.outbound_payload.rich_message_fields is not None assert ctx.outbound_payload.rich_message_fields.rich_message is rich + + +# ─── 持久化模式 AgentRun 内容加载 ──────────────────────────────────────────── + + +@pytest.mark.unit +class TestLoadBuildStagePersistentContentLoad: + """持久化模式下 AgentRun 内容加载。""" + + @pytest.mark.asyncio + async def test_persistent_mode_loads_agent_run_output(self): + # Arrange — 持久化模式 + agent_run_id 非空 + stream_chunks 为空 → 触发加载 + port = _make_agent_run_port( + stream_events=["evt-1", "evt-2"], + final_output=Some("full response text"), + ) + stage = LoadBuildStage(agent_run_port=port) + ctx = _make_ctx(delivery_mode="persistent", agent_run_id="run-001", stream_chunks=[]) + + # Act + ok = await stage.process(ctx) + + # Assert + assert ok is True + # streamAgentRun 被调用(阻塞消费至完成) + port.streamAgentRun.assert_called_once() + # getAgentRunFinalOutput 被调用(获取完整文本) + port.getAgentRunFinalOutput.assert_called_once() + # 完整文本填充到 stream_chunks + assert ctx.stream_chunks == ["full response text"] + assert ctx.outbound_payload.stream_chunks == ("full response text",) + + @pytest.mark.asyncio + async def test_persistent_mode_final_output_nothing_keeps_empty(self): + # Arrange — getAgentRunFinalOutput 返回 Nothing(AgentRun 无文本输出) + port = _make_agent_run_port( + stream_events=["evt-1"], + final_output=Nothing(), + ) + stage = LoadBuildStage(agent_run_port=port) + ctx = _make_ctx(delivery_mode="persistent", agent_run_id="run-001", stream_chunks=[]) + + # Act + ok = await stage.process(ctx) + + # Assert — stream_chunks 保持空,不抛异常 + assert ok is True + assert ctx.stream_chunks == [] + assert ctx.outbound_payload.stream_chunks == () + + @pytest.mark.asyncio + async def test_streaming_mode_skips_agent_run_load(self): + # Arrange — 流式模式不触发内容加载(由 stream_chunk_stage 负责) + port = _make_agent_run_port() + stage = LoadBuildStage(agent_run_port=port) + ctx = _make_ctx(delivery_mode="streaming", agent_run_id="run-001", stream_chunks=[]) + + # Act + ok = await stage.process(ctx) + + # Assert + assert ok is True + port.streamAgentRun.assert_not_called() + port.getAgentRunFinalOutput.assert_not_called() + + @pytest.mark.asyncio + async def test_pre_filled_stream_chunks_skips_agent_run_load(self): + # Arrange — stream_chunks 已预填充(如静默命令响应路径)不触发加载 + port = _make_agent_run_port() + stage = LoadBuildStage(agent_run_port=port) + ctx = _make_ctx( + delivery_mode="persistent", + agent_run_id="run-001", + stream_chunks=["pre-filled"], + ) + + # Act + ok = await stage.process(ctx) + + # Assert + assert ok is True + port.streamAgentRun.assert_not_called() + port.getAgentRunFinalOutput.assert_not_called() + assert ctx.outbound_payload.stream_chunks == ("pre-filled",) + + @pytest.mark.asyncio + async def test_empty_agent_run_id_skips_agent_run_load(self): + # Arrange — agent_run_id 为空(如管理员消息)不触发加载 + port = _make_agent_run_port() + stage = LoadBuildStage(agent_run_port=port) + ctx = _make_ctx(delivery_mode="persistent", agent_run_id="", stream_chunks=["manual"]) + + # Act + ok = await stage.process(ctx) + + # Assert + assert ok is True + port.streamAgentRun.assert_not_called() + port.getAgentRunFinalOutput.assert_not_called() diff --git a/backend/test/unit/channels/application/pipeline/outbound/test_prefix_stage.py b/backend/test/unit/channels/application/pipeline/outbound/test_prefix_stage.py index 7db656ff..d3494e92 100644 --- a/backend/test/unit/channels/application/pipeline/outbound/test_prefix_stage.py +++ b/backend/test/unit/channels/application/pipeline/outbound/test_prefix_stage.py @@ -123,11 +123,11 @@ class TestPrefixStageContract: stage = PrefixStage() assert stage.thread_safe is True - def test_failure_strategy_is_skip(self): - # failure=SKIP:trusted_message 缺失时返回 False,管道跳过本阶段但不 - # 终止,下游阶段继续执行(消息仍可投递,仅无前缀) + def test_failure_strategy_is_terminate(self): + # failure=TERMINATE:非流式路径下 content 必须非空(fail-closed), + # ValidationError 终止管道而非吞掉空消息继续投递 stage = PrefixStage() - assert stage.failure is FailureStrategy.SKIP + assert stage.failure is FailureStrategy.TERMINATE def test_compensate_is_none(self): stage = PrefixStage() diff --git a/backend/test/unit/channels/application/pipeline/outbound/test_stream_chunk_stage.py b/backend/test/unit/channels/application/pipeline/outbound/test_stream_chunk_stage.py index feca51e0..19a2c865 100644 --- a/backend/test/unit/channels/application/pipeline/outbound/test_stream_chunk_stage.py +++ b/backend/test/unit/channels/application/pipeline/outbound/test_stream_chunk_stage.py @@ -4,18 +4,26 @@ - ``process(context)``:适配器不存在降级、正常流式投递、TTL 超时停止、 投递异常降级、预填充分块遍历 - ``_iterChunkContents()``:agent_run_id 非空走 streamAgentRun、为空遍历 stream_chunks +- ``_extractTextFromStreamEvent()``:event_type 过滤、两层 payload 穿透、 + stream_event.content/response 回退 +- 降级场景:TTL/异常降级后 formatted_message 重新格式化 """ from __future__ import annotations +import itertools from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channels.application.context.outbound_context import OutboundContext from yuxi.channels.application.pipeline.outbound.stream_chunk_stage import ( StreamChunkStage, + _extractTextFromStreamEvent, ) from yuxi.channels.contract.dtos.channel import ChannelType +from yuxi.channels.contract.dtos.common import MessageFormat +from yuxi.channels.contract.dtos.outbound import FormattedMessage +from yuxi.channels.contract.dtos.stream_event import StreamEvent from yuxi.channels.contract.dtos.streaming import StreamingConfig from yuxi.channels.contract.errors.base import Error from yuxi.channels.contract.errors.domain import ChannelDegradedError @@ -77,6 +85,40 @@ def _make_async_iter(items: list): return _Iter() +def _make_stream_event(content: str, *, event_type: str = "messages") -> StreamEvent: + """构造 messages 类型的 StreamEvent,payload 含单 item。 + + item 结构与 run_worker 写入 Redis Stream 的一致: + {"stream_event": {"content": ...}, "response": ...} + """ + return StreamEvent( + event_type=event_type, + payload={ + "event": event_type, + "payload": { + "items": [ + { + "stream_event": {"content": content}, + "response": content, + "status": "loading", + } + ] + }, + }, + seq="1-0", + ) + + +def _make_outbound_adapter(*, formatted_content: str = "formatted") -> AsyncMock: + """构造出站适配器桩,formatOutbound 返回指定内容。""" + adapter = AsyncMock() + adapter.formatOutbound.return_value = FormattedMessage( + content=formatted_content, + format=MessageFormat.TEXT, + ) + return adapter + + @pytest.mark.unit class TestStreamChunkStageProcess: @pytest.mark.asyncio @@ -86,6 +128,7 @@ class TestStreamChunkStageProcess: streaming_adapter_registry={}, config_port=_make_config_port(), agent_run_port=AsyncMock(), + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx() @@ -101,16 +144,21 @@ class TestStreamChunkStageProcess: @pytest.mark.asyncio async def test_normal_streaming_delivery_with_agent_run(self): - # Arrange + # Arrange:streamAgentRun 返回真实 StreamEvent 对象 adapter = AsyncMock() adapter.sendChunk.return_value = MagicMock(success=True) agent_run_port = AsyncMock() - # streamAgentRun 在源码中为同步调用(未 await),返回异步迭代器 - agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["chunk-1", "chunk-2"])) + agent_run_port.streamAgentRun = MagicMock( + return_value=_make_async_iter([ + _make_stream_event("chunk-1"), + _make_stream_event("chunk-2"), + ]) + ) stage = StreamChunkStage( streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(min_interval_ms=0, ttl_ms=60000), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx(agent_run_id="run-001") @@ -134,6 +182,7 @@ class TestStreamChunkStageProcess: streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(), agent_run_port=AsyncMock(), + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx(agent_run_id="", stream_chunks=["a", "b", "c"]) @@ -152,12 +201,19 @@ class TestStreamChunkStageProcess: adapter = AsyncMock() adapter.sendChunk.return_value = MagicMock(success=True) agent_run_port = AsyncMock() - agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["chunk-1", "chunk-2", "chunk-3"])) + agent_run_port.streamAgentRun = MagicMock( + return_value=_make_async_iter([ + _make_stream_event("chunk-1"), + _make_stream_event("chunk-2"), + _make_stream_event("chunk-3"), + ]) + ) logger = AsyncMock() stage = StreamChunkStage( streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(ttl_ms=0), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=logger, ) ctx = _make_ctx() @@ -165,8 +221,9 @@ class TestStreamChunkStageProcess: # monkeypatch time.monotonic 返回递增值,避免 Windows 时钟精度导致 # elapsed=0.0 > ttl_s=0.0 为 False 的 flaky 失败。 # 序列说明:started_at=0.0,首次循环 elapsed=0.0(不超时,发送 chunk-1), - # 第二次循环 elapsed=0.002 > 0(超时,break),results 含 1 个分块。 - time_values = iter([0.0, 0.0, 0.002, 0.003]) + # 第二次循环 elapsed=0.001 > 0(超时,break)。 + # 使用 chain+count 提供无限序列,避免 teardown 时迭代器耗尽。 + time_values = itertools.chain([0.0, 0.0], itertools.count(0.001, 0.001)) monkeypatch.setattr( "yuxi.channels.application.pipeline.outbound.stream_chunk_stage.time.monotonic", lambda: next(time_values), @@ -185,11 +242,14 @@ class TestStreamChunkStageProcess: adapter = AsyncMock() adapter.sendChunk.side_effect = Error("send failed") agent_run_port = AsyncMock() - agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["x"])) + agent_run_port.streamAgentRun = MagicMock( + return_value=_make_async_iter([_make_stream_event("x")]) + ) stage = StreamChunkStage( streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx() @@ -203,19 +263,22 @@ class TestStreamChunkStageProcess: assert ctx.stream_aborted_at_chunk == 0 @pytest.mark.asyncio - async def test_send_chunk_raises_asyncio_timeout_error_degrades(self): - # Arrange:原生 TimeoutError 不属于 Error 子类, - # 修复前会被 except Error 漏掉导致消息丢失(C2-O3) + async def test_send_chunk_raises_timeout_error_degrades(self): + # Arrange:原生 TimeoutError 不属于 Error 子类(C2-O3) adapter = AsyncMock() adapter.sendChunk.side_effect = TimeoutError() agent_run_port = AsyncMock() agent_run_port.streamAgentRun = MagicMock( - return_value=_make_async_iter(["chunk-a", "chunk-b"]) + return_value=_make_async_iter([ + _make_stream_event("chunk-a"), + _make_stream_event("chunk-b"), + ]) ) stage = StreamChunkStage( streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx() @@ -237,12 +300,16 @@ class TestStreamChunkStageProcess: adapter.sendChunk.side_effect = ConnectionError("network down") agent_run_port = AsyncMock() agent_run_port.streamAgentRun = MagicMock( - return_value=_make_async_iter(["chunk-a", "chunk-b"]) + return_value=_make_async_iter([ + _make_stream_event("chunk-a"), + _make_stream_event("chunk-b"), + ]) ) stage = StreamChunkStage( streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx() @@ -258,18 +325,23 @@ class TestStreamChunkStageProcess: assert exc_info.value.__cause__ is not None @pytest.mark.asyncio - async def test_send_chunk_raises_native_error_after_partial_send_sets_position(self): + async def test_send_chunk_raises_error_after_partial_send_sets_position(self): # Arrange:第二个分块投递失败,stream_aborted_at_chunk 应记录已投递数量 adapter = AsyncMock() adapter.sendChunk.side_effect = [MagicMock(success=True), ConnectionError("down")] agent_run_port = AsyncMock() agent_run_port.streamAgentRun = MagicMock( - return_value=_make_async_iter(["chunk-a", "chunk-b", "chunk-c"]) + return_value=_make_async_iter([ + _make_stream_event("chunk-a"), + _make_stream_event("chunk-b"), + _make_stream_event("chunk-c"), + ]) ) stage = StreamChunkStage( streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx() @@ -290,6 +362,7 @@ class TestStreamChunkStageProcess: streaming_adapter_registry={ChannelType("feishu"): adapter}, config_port=_make_config_port(), agent_run_port=AsyncMock(), + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx(agent_run_id="", stream_chunks=[]) @@ -300,9 +373,6 @@ class TestStreamChunkStageProcess: # Assert assert ok is True assert ctx.chunk_results == [] - # 空分块流不构造 StreamingCompleted(total_chunks=0 违反 DTO 不变量, - # 且 truncation_check 依赖 streaming_completed is not None 作为激活条件, - # 空流无需截断检测) assert ctx.streaming_completed is None @@ -310,13 +380,19 @@ class TestStreamChunkStageProcess: class TestIterChunkContents: @pytest.mark.asyncio async def test_agent_run_id_non_empty_uses_stream_agent_run(self): - # Arrange + # Arrange:streamAgentRun 返回 StreamEvent 对象 agent_run_port = AsyncMock() - agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["a", "b"])) + agent_run_port.streamAgentRun = MagicMock( + return_value=_make_async_iter([ + _make_stream_event("a"), + _make_stream_event("b"), + ]) + ) stage = StreamChunkStage( streaming_adapter_registry={}, config_port=_make_config_port(), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx(agent_run_id="run-001", stream_chunks=[]) @@ -326,7 +402,7 @@ class TestIterChunkContents: async for content in stage._iterChunkContents(ctx): contents.append(content) - # Assert:分块内容追加到 stream_chunks + # Assert:从 StreamEvent 提取的文本追加到 stream_chunks assert contents == ["a", "b"] assert ctx.stream_chunks == ["a", "b"] agent_run_port.streamAgentRun.assert_called_once() @@ -339,6 +415,7 @@ class TestIterChunkContents: streaming_adapter_registry={}, config_port=_make_config_port(), agent_run_port=agent_run_port, + outbound_adapter_registry={}, logger=AsyncMock(), ) ctx = _make_ctx(agent_run_id="", stream_chunks=["x", "y"]) @@ -351,3 +428,253 @@ class TestIterChunkContents: # Assert:不调用 streamAgentRun,直接遍历预填充列表 assert contents == ["x", "y"] agent_run_port.streamAgentRun.assert_not_awaited() + + +@pytest.mark.unit +class TestExtractTextFromStreamEvent: + """StreamEvent 内容提取逻辑。""" + + def test_messages_event_extracts_stream_event_content(self): + # Arrange:stream_event.content 存在时优先取 + event = _make_stream_event("hello") + + # Act + contents = list(_extractTextFromStreamEvent(event)) + + # Assert + assert contents == ["hello"] + + def test_non_messages_event_skipped(self): + # Arrange:metadata/custom/error/end 等事件不提取内容 + event = StreamEvent( + event_type="metadata", + payload={"payload": {"items": [{"response": "should-be-skipped"}]}}, + seq="1-0", + ) + + # Act + contents = list(_extractTextFromStreamEvent(event)) + + # Assert + assert contents == [] + + def test_falls_back_to_response_when_no_stream_event_content(self): + # Arrange:stream_event.content 缺失时回退到 response + event = StreamEvent( + event_type="messages", + payload={ + "payload": { + "items": [{"response": "fallback-text"}] + } + }, + seq="1-0", + ) + + # Act + contents = list(_extractTextFromStreamEvent(event)) + + # Assert + assert contents == ["fallback-text"] + + def test_multiple_items_extracted_in_order(self): + # Arrange:一个 messages 事件含多个 item + event = StreamEvent( + event_type="messages", + payload={ + "payload": { + "items": [ + {"stream_event": {"content": "part-1"}}, + {"stream_event": {"content": "part-2"}}, + {"response": "part-3"}, + ] + } + }, + seq="1-0", + ) + + # Act + contents = list(_extractTextFromStreamEvent(event)) + + # Assert + assert contents == ["part-1", "part-2", "part-3"] + + def test_empty_content_skipped(self): + # Arrange:空字符串 content 不产出 + event = StreamEvent( + event_type="messages", + payload={ + "payload": { + "items": [ + {"stream_event": {"content": ""}}, + {"response": ""}, + ] + } + }, + seq="1-0", + ) + + # Act + contents = list(_extractTextFromStreamEvent(event)) + + # Assert + assert contents == [] + + def test_non_dict_item_skipped(self): + # Arrange:非 dict item 跳过 + event = StreamEvent( + event_type="messages", + payload={"payload": {"items": ["not-a-dict", {"response": "valid"}]}}, + seq="1-0", + ) + + # Act + contents = list(_extractTextFromStreamEvent(event)) + + # Assert + assert contents == ["valid"] + + def test_empty_payload_returns_nothing(self): + # Arrange:payload 为空 dict + event = StreamEvent(event_type="messages", payload={}, seq="1-0") + + # Act + contents = list(_extractTextFromStreamEvent(event)) + + # Assert + assert contents == [] + + +@pytest.mark.unit +class TestReformatAfterDegrade: + """降级后 formatted_message 重新格式化。""" + + @pytest.mark.asyncio + async def test_ttl_degrade_reformats_formatted_message(self, monkeypatch): + # Arrange:TTL 超时降级后,stream_chunks 已填充,formatted_message 应更新 + adapter = AsyncMock() + adapter.sendChunk.return_value = MagicMock(success=True) + outbound_adapter = _make_outbound_adapter(formatted_content="reformatted") + agent_run_port = AsyncMock() + agent_run_port.streamAgentRun = MagicMock( + return_value=_make_async_iter([ + _make_stream_event("chunk-1"), + _make_stream_event("chunk-2"), + ]) + ) + stage = StreamChunkStage( + streaming_adapter_registry={ChannelType("feishu"): adapter}, + config_port=_make_config_port(ttl_ms=0), + agent_run_port=agent_run_port, + outbound_adapter_registry={ChannelType("feishu"): outbound_adapter}, + logger=AsyncMock(), + ) + ctx = _make_ctx() + ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT) + + # TTL=0 使首块后超时:started_at=0.0,首块 elapsed=0.0(不超时), + # 第二块 elapsed=0.001 > 0(超时,break)。 + time_values = itertools.chain([0.0, 0.0], itertools.count(0.001, 0.001)) + monkeypatch.setattr( + "yuxi.channels.application.pipeline.outbound.stream_chunk_stage.time.monotonic", + lambda: next(time_values), + ) + + # Act + await stage.process(ctx) + + # Assert:降级后 formatted_message 已更新,content 非空 + assert ctx.delivery_mode == "persistent" + assert ctx.formatted_message.content == "reformatted" + outbound_adapter.formatOutbound.assert_awaited_once() + + @pytest.mark.asyncio + async def test_exception_degrade_reformats_formatted_message(self): + # Arrange:投递异常降级后,stream_chunks 已填充,formatted_message 应更新 + adapter = AsyncMock() + adapter.sendChunk.side_effect = ConnectionError("down") + outbound_adapter = _make_outbound_adapter(formatted_content="reformatted") + agent_run_port = AsyncMock() + agent_run_port.streamAgentRun = MagicMock( + return_value=_make_async_iter([_make_stream_event("chunk-1")]) + ) + stage = StreamChunkStage( + streaming_adapter_registry={ChannelType("feishu"): adapter}, + config_port=_make_config_port(), + agent_run_port=agent_run_port, + outbound_adapter_registry={ChannelType("feishu"): outbound_adapter}, + logger=AsyncMock(), + ) + ctx = _make_ctx() + ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT) + + # Act + with pytest.raises(ChannelDegradedError): + await stage.process(ctx) + + # Assert:降级后(抛出前)formatted_message 已更新 + assert ctx.delivery_mode == "persistent" + assert ctx.formatted_message.content == "reformatted" + outbound_adapter.formatOutbound.assert_awaited_once() + + @pytest.mark.asyncio + async def test_adapter_not_registered_degrade_reformats(self): + # Arrange:适配器未注册降级时,stream_chunks 已预填充 + outbound_adapter = _make_outbound_adapter(formatted_content="reformatted") + stage = StreamChunkStage( + streaming_adapter_registry={}, + config_port=_make_config_port(), + agent_run_port=AsyncMock(), + outbound_adapter_registry={ChannelType("feishu"): outbound_adapter}, + logger=AsyncMock(), + ) + ctx = _make_ctx(stream_chunks=["pre-filled"], agent_run_id="") + ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT) + + # Act + await stage.process(ctx) + + # Assert + assert ctx.delivery_mode == "persistent" + assert ctx.formatted_message.content == "reformatted" + + @pytest.mark.asyncio + async def test_degrade_with_empty_chunks_skips_reformat(self): + # Arrange:降级时 stream_chunks 为空,不调用 formatOutbound + outbound_adapter = _make_outbound_adapter() + stage = StreamChunkStage( + streaming_adapter_registry={}, + config_port=_make_config_port(), + agent_run_port=AsyncMock(), + outbound_adapter_registry={ChannelType("feishu"): outbound_adapter}, + logger=AsyncMock(), + ) + ctx = _make_ctx(stream_chunks=[], agent_run_id="") + ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT) + + # Act + await stage.process(ctx) + + # Assert:空 chunks 不重新格式化 + outbound_adapter.formatOutbound.assert_not_awaited() + assert ctx.formatted_message.content == "" + + @pytest.mark.asyncio + async def test_degrade_without_outbound_adapter_falls_back_to_text(self): + # Arrange:出站适配器未注册时,直接构造 TEXT 格式 + stage = StreamChunkStage( + streaming_adapter_registry={}, + config_port=_make_config_port(), + agent_run_port=AsyncMock(), + outbound_adapter_registry={}, + logger=AsyncMock(), + ) + ctx = _make_ctx(stream_chunks=["fallback-text"], agent_run_id="") + ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT) + + # Act + await stage.process(ctx) + + # Assert:无出站适配器时直接用原始 content 构造 TEXT + assert ctx.delivery_mode == "persistent" + assert ctx.formatted_message.content == "fallback-text" + assert ctx.formatted_message.format == MessageFormat.TEXT diff --git a/backend/test/unit/channels/application/usecase/test_identity_merge_service.py b/backend/test/unit/channels/application/usecase/test_identity_merge_service.py index a893d5b1..89893bdd 100644 --- a/backend/test/unit/channels/application/usecase/test_identity_merge_service.py +++ b/backend/test/unit/channels/application/usecase/test_identity_merge_service.py @@ -410,6 +410,8 @@ class TestApproveMerge: user_identity_repo.getUserIdentityByIdentityId = AsyncMock( side_effect=lambda iid: {"canonical-1": canonical, "sibling-1": sibling}.get(iid) ) + # updateUserIdentity 返回输入 DTO(模拟 DB 更新后返回,version 与聚合根一致) + user_identity_repo.updateUserIdentity = AsyncMock(side_effect=lambda dto, **kw: dto) session_repo = AsyncMock() session_repo.findSessionsByFilter = AsyncMock(return_value=sessions) session_repo.updateChannelSession = AsyncMock(return_value=sessions[0]) @@ -431,7 +433,7 @@ class TestApproveMerge: assert result.canonical_identity_id == _CANONICAL_ID assert result.merged_identity_id == _SIBLING_ID assert result.canonical_version == 6 - # sibling.version: 3 (DB) +1 (clearPendingReview) +1 (unbindUser, M-7) = 5 + # sibling.version: 3 (DB) → +1 (clearPendingReview) 持久化 → +1 (unbindUser) 持久化 = 5 assert result.sibling_version == 5 # 校验 canonical 持久化:expected_version 用合并前 DB 版本 5 @@ -445,10 +447,11 @@ class TestApproveMerge: bindings_arg = call_kwargs.args[1] assert _PEER_SIBLING in bindings_arg.get("feishu", []) - # 校验 sibling pending_review 清除 - user_identity_repo.updateUserIdentity.assert_awaited_once() - sibling_dto = user_identity_repo.updateUserIdentity.call_args.args[0] - assert sibling_dto.pending_review is False + # 校验 sibling 两次持久化:clearPendingReview + unbindUser 拆分持久化 + assert user_identity_repo.updateUserIdentity.await_count == 2 + # 最后一次调用是 unbindUser 后的 DTO:pending_review 已在第一次清除 + last_call_dto = user_identity_repo.updateUserIdentity.await_args.args[0] + assert last_call_dto.pending_review is False # 校验会话迁移:sibling 名下的会话迁移到 canonical # findSessionsByFilter 用 sibling_id 过滤,返回 2 条会话 diff --git a/backend/test/unit/channels/core/model/test_outbox_entry.py b/backend/test/unit/channels/core/model/test_outbox_entry.py index 497c020e..2788f88f 100644 --- a/backend/test/unit/channels/core/model/test_outbox_entry.py +++ b/backend/test/unit/channels/core/model/test_outbox_entry.py @@ -161,7 +161,7 @@ class TestOutboxEntryMarkSent: # Assert assert entry.status == OutboxStatus.SENT assert entry.channel_msg_id == "channel-msg-1" - assert entry.version == original_version + 1 + assert entry.version == original_version assert entry.funnel_node == "sent" assert entry.latency_ms is not None assert entry.latency_ms >= 0 @@ -223,7 +223,7 @@ class TestOutboxEntryMarkSuppressed: assert entry.status == OutboxStatus.SUPPRESSED assert entry.last_error == "stale reply" assert entry.funnel_node == "suppressed" - assert entry.version == original_version + 1 + assert entry.version == original_version @pytest.mark.parametrize( "status", @@ -301,7 +301,7 @@ class TestOutboxEntryMarkFailed: assert entry.next_retry_at is not None assert entry.last_error == "network error" assert entry.funnel_node == "failed" - assert entry.version == original_version + 1 + assert entry.version == original_version assert entry.last_retry_at is not None def test_mark_failed_does_not_auto_transition_to_dead(self): @@ -365,7 +365,7 @@ class TestOutboxEntryMarkDead: assert entry.next_retry_at is None assert entry.last_error == "retry exhausted" assert entry.funnel_node == "dead" - assert entry.version == original_version + 1 + assert entry.version == original_version @pytest.mark.parametrize( "status", @@ -455,7 +455,7 @@ class TestOutboxEntryMarkPartialFailure: assert entry.status == OutboxStatus.SENT_UNCONFIRMED assert entry.partial_failure is True assert entry.last_error == "partial failure on parts: [1, 3]" - assert entry.version == original_version + 1 + assert entry.version == original_version @pytest.mark.parametrize( "status", @@ -492,7 +492,7 @@ class TestOutboxEntryRequeueForRetry: assert entry.status == OutboxStatus.PENDING assert entry.next_retry_at is None assert entry.funnel_node == "enter" - assert entry.version == original_version + 1 + assert entry.version == original_version # retry_count 保持不变 assert entry.retry_count == 2 @@ -656,14 +656,14 @@ class TestOutboxEntryRevive: # Assert assert entry.funnel_node == "failed" - def test_revive_increments_version(self): + def test_revive_preserves_version(self): # Arrange entry = self._make_dead_entry() original_version = entry.version # Act entry.revive() # Assert - assert entry.version == original_version + 1 + assert entry.version == original_version def test_revive_from_pending_raises_rule_violation_error(self): # Arrange diff --git a/backend/test/unit/channels/core/model/test_outbox_entry_aggregate.py b/backend/test/unit/channels/core/model/test_outbox_entry_aggregate.py index aad944e5..5502db8e 100644 --- a/backend/test/unit/channels/core/model/test_outbox_entry_aggregate.py +++ b/backend/test/unit/channels/core/model/test_outbox_entry_aggregate.py @@ -223,7 +223,7 @@ class TestMarkPartialFailure: assert entry.partial_failure is True assert "2" in entry.last_error assert "3" in entry.last_error - assert entry.version == original_version + 1 + assert entry.version == original_version def test_partial_failure_does_not_change_funnel_node(self): """markPartialFailure 不修改 funnel_node(仍是 SENT_UNCONFIRMED 路径)。""" @@ -326,7 +326,7 @@ class TestRequeueForRetry: assert entry.status == OutboxStatus.PENDING assert entry.next_retry_at is None assert entry.funnel_node == "enter" - assert entry.version == original_version + 1 + assert entry.version == original_version def test_retry_count_unchanged_after_requeue(self): """requeueForRetry 不重置 retry_count(重试预算已由 markFailed 递增)。""" diff --git a/backend/test/unit/channels/infrastructure/test_channel_use_cases.py b/backend/test/unit/channels/infrastructure/test_channel_use_cases.py index 43a1de44..06d85ece 100644 --- a/backend/test/unit/channels/infrastructure/test_channel_use_cases.py +++ b/backend/test/unit/channels/infrastructure/test_channel_use_cases.py @@ -3,7 +3,7 @@ 覆盖: - ``_fill_adapter_registry``:纯函数逻辑(最后注册者覆盖 / duplicate_rule 校验) - 8 个 ``create_channel_*_handler_dependencies`` 工厂函数:验证返回字典 - 键集与 ``*_repo_factory`` 闭包产物类型,确保 worker 进程装配契约稳定 + 键集与 ``xxx_repo`` 端口实例类型,确保 worker 进程装配契约稳定 不依赖运行中的 Docker 服务,纯单元测试。 """ @@ -25,7 +25,6 @@ from yuxi.channels.adapters.structured_logger_adapter import ( from yuxi.channels.contract.dtos.channel import ChannelType from yuxi.channels.contract.dtos.outbox import OutboxConfig from yuxi.channels.contract.errors import RuleViolationError -from yuxi.channels.contract.ports.driven.logger_port import LoggerPort from yuxi.channels.infrastructure.channel_use_cases import ( _fill_adapter_registry, create_channel_audit_log_retention_handler_dependencies, @@ -156,7 +155,7 @@ class TestFillAdapterRegistry: @pytest.fixture def mock_session_factory() -> MagicMock: - """worker 进程级 session_factory 桩。""" + """worker 进程级 session_factory 桩(``Callable[[], AsyncSession]``)。""" return MagicMock(name="session_factory") @@ -167,9 +166,9 @@ def mock_arq_pool() -> MagicMock: @pytest.fixture -def mock_db() -> MagicMock: - """AsyncSession 桩,传入 *_repo_factory 构造适配器。""" - return MagicMock(name="db") +def mock_cache_port() -> MagicMock: + """worker 进程级 CachePort 桩,供 handler 分布式锁使用。""" + return MagicMock(name="cache_port") def _assert_logger_port(logger: object) -> None: @@ -186,105 +185,94 @@ def _assert_outbox_config(outbox_config: object) -> None: class TestSessionInactiveCleanupHandlerDependencies: """create_channel_session_inactive_cleanup_handler_dependencies 测试。""" - def test_returns_dict_with_required_keys(self, mock_session_factory): + def test_returns_dict_with_required_keys( + self, mock_session_factory, mock_cache_port + ): # Act deps = create_channel_session_inactive_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) # Assert - assert set(deps.keys()) == { - "session_factory", - "session_repo_factory", - "logger", - } - assert deps["session_factory"] is mock_session_factory + assert set(deps.keys()) == {"session_repo", "cache_port", "logger"} + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) - def test_session_repo_factory_creates_persistence_adapter( - self, mock_session_factory, mock_db + def test_session_repo_is_persistence_adapter( + self, mock_session_factory, mock_cache_port ): - # Arrange + # Act deps = create_channel_session_inactive_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) - # Act - repo = deps["session_repo_factory"](mock_db) - # Assert: fat adapter 实现 ChannelSessionRepositoryPort - assert isinstance(repo, ChannelPersistenceAdapter) + assert isinstance(deps["session_repo"], ChannelPersistenceAdapter) @pytest.mark.unit class TestOutboxTerminalCleanupHandlerDependencies: """create_channel_outbox_terminal_cleanup_handler_dependencies 测试。""" - def test_returns_dict_with_required_keys(self, mock_session_factory): + def test_returns_dict_with_required_keys( + self, mock_session_factory, mock_cache_port + ): # Act deps = create_channel_outbox_terminal_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) # Assert assert set(deps.keys()) == { - "session_factory", - "outbox_repo_factory", + "outbox_repo", "event_publisher", + "cache_port", "logger", } - assert deps["session_factory"] is mock_session_factory # event_publisher 为 None:worker 进程未注入事件发布端口(best-effort) assert deps["event_publisher"] is None + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) - def test_outbox_repo_factory_creates_persistence_adapter( - self, mock_session_factory, mock_db + def test_outbox_repo_is_persistence_adapter( + self, mock_session_factory, mock_cache_port ): - # Arrange + # Act deps = create_channel_outbox_terminal_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) - # Act - repo = deps["outbox_repo_factory"](mock_db) - # Assert - assert isinstance(repo, ChannelPersistenceAdapter) + assert isinstance(deps["outbox_repo"], ChannelPersistenceAdapter) @pytest.mark.unit class TestPairingTerminalCleanupHandlerDependencies: """create_channel_pairing_terminal_cleanup_handler_dependencies 测试。""" - def test_returns_dict_with_required_keys(self, mock_session_factory): + def test_returns_dict_with_required_keys( + self, mock_session_factory, mock_cache_port + ): # Act deps = create_channel_pairing_terminal_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) # Assert - assert set(deps.keys()) == { - "session_factory", - "pairing_repo_factory", - "logger", - } - assert deps["session_factory"] is mock_session_factory + assert set(deps.keys()) == {"pairing_repo", "cache_port", "logger"} + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) - def test_pairing_repo_factory_creates_persistence_adapter( - self, mock_session_factory, mock_db + def test_pairing_repo_is_persistence_adapter( + self, mock_session_factory, mock_cache_port ): - # Arrange + # Act deps = create_channel_pairing_terminal_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) - # Act - repo = deps["pairing_repo_factory"](mock_db) - # Assert - assert isinstance(repo, ChannelPersistenceAdapter) + assert isinstance(deps["pairing_repo"], ChannelPersistenceAdapter) @pytest.mark.unit @@ -292,176 +280,160 @@ class TestOutboxRecoveryHandlerDependencies: """create_channel_outbox_recovery_handler_dependencies 测试。""" def test_returns_dict_with_required_keys( - self, mock_session_factory, mock_arq_pool + self, mock_session_factory, mock_arq_pool, mock_cache_port ): # Act deps = create_channel_outbox_recovery_handler_dependencies( - mock_session_factory, arq_pool=mock_arq_pool + mock_session_factory, arq_pool=mock_arq_pool, cache_port=mock_cache_port ) # Assert assert set(deps.keys()) == { - "session_factory", - "outbox_repo_factory", - "audit_log_repo_factory", + "outbox_repo", + "audit_log_repo", "queue_port", "outbox_config", + "cache_port", "logger", } - assert deps["session_factory"] is mock_session_factory + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) _assert_outbox_config(deps["outbox_config"]) - def test_outbox_and_audit_repo_factories_create_persistence_adapter( - self, mock_session_factory, mock_arq_pool, mock_db + def test_outbox_and_audit_repos_are_persistence_adapters( + self, mock_session_factory, mock_arq_pool, mock_cache_port ): - """outbox_repo_factory 与 audit_log_repo_factory 都构造 fat adapter。""" - # Arrange + """outbox_repo 与 audit_log_repo 都为 fat adapter 实例。""" + # Act deps = create_channel_outbox_recovery_handler_dependencies( - mock_session_factory, arq_pool=mock_arq_pool + mock_session_factory, arq_pool=mock_arq_pool, cache_port=mock_cache_port ) - # Act - outbox_repo = deps["outbox_repo_factory"](mock_db) - audit_repo = deps["audit_log_repo_factory"](mock_db) - # Assert: 两者都是 ChannelPersistenceAdapter(fat adapter 实现多端口) - assert isinstance(outbox_repo, ChannelPersistenceAdapter) - assert isinstance(audit_repo, ChannelPersistenceAdapter) + assert isinstance(deps["outbox_repo"], ChannelPersistenceAdapter) + assert isinstance(deps["audit_log_repo"], ChannelPersistenceAdapter) @pytest.mark.unit class TestPairingExpirationHandlerDependencies: """create_channel_pairing_expiration_handler_dependencies 测试。""" - def test_returns_dict_with_required_keys(self, mock_session_factory): + def test_returns_dict_with_required_keys( + self, mock_session_factory, mock_cache_port + ): # Act deps = create_channel_pairing_expiration_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) # Assert - assert set(deps.keys()) == { - "session_factory", - "pairing_repo_factory", - "logger", - } + assert set(deps.keys()) == {"pairing_repo", "cache_port", "logger"} + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) - def test_pairing_repo_factory_creates_persistence_adapter( - self, mock_session_factory, mock_db + def test_pairing_repo_is_persistence_adapter( + self, mock_session_factory, mock_cache_port ): - # Arrange + # Act deps = create_channel_pairing_expiration_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) - # Act - repo = deps["pairing_repo_factory"](mock_db) - # Assert - assert isinstance(repo, ChannelPersistenceAdapter) + assert isinstance(deps["pairing_repo"], ChannelPersistenceAdapter) @pytest.mark.unit class TestAuditLogRetentionHandlerDependencies: """create_channel_audit_log_retention_handler_dependencies 测试。""" - def test_returns_dict_with_required_keys(self, mock_session_factory): + def test_returns_dict_with_required_keys( + self, mock_session_factory, mock_cache_port + ): # Act deps = create_channel_audit_log_retention_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) # Assert - assert set(deps.keys()) == { - "session_factory", - "audit_log_repo_factory", - "logger", - } + assert set(deps.keys()) == {"audit_log_repo", "cache_port", "logger"} + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) - def test_audit_log_repo_factory_creates_persistence_adapter( - self, mock_session_factory, mock_db + def test_audit_log_repo_is_persistence_adapter( + self, mock_session_factory, mock_cache_port ): - # Arrange + # Act deps = create_channel_audit_log_retention_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) - # Act - repo = deps["audit_log_repo_factory"](mock_db) - # Assert - assert isinstance(repo, ChannelPersistenceAdapter) + assert isinstance(deps["audit_log_repo"], ChannelPersistenceAdapter) @pytest.mark.unit class TestContentReviewRetentionHandlerDependencies: """create_channel_content_review_retention_handler_dependencies 测试。""" - def test_returns_dict_with_required_keys(self, mock_session_factory): + def test_returns_dict_with_required_keys( + self, mock_session_factory, mock_cache_port + ): # Act deps = create_channel_content_review_retention_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) # Assert assert set(deps.keys()) == { - "session_factory", - "content_review_repo_factory", + "content_review_repo", + "cache_port", "logger", } + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) - def test_content_review_repo_factory_creates_review_adapter( - self, mock_session_factory, mock_db + def test_content_review_repo_is_review_adapter( + self, mock_session_factory, mock_cache_port ): """与其它 7 个 handler 不同,content_review 使用独立的 Adapter。""" - # Arrange + # Act deps = create_channel_content_review_retention_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) - # Act - repo = deps["content_review_repo_factory"](mock_db) - # Assert - assert isinstance(repo, ContentReviewRepositoryAdapter) - assert not isinstance(repo, ChannelPersistenceAdapter) + assert isinstance(deps["content_review_repo"], ContentReviewRepositoryAdapter) + assert not isinstance(deps["content_review_repo"], ChannelPersistenceAdapter) @pytest.mark.unit class TestIdempotencyCleanupHandlerDependencies: """create_channel_idempotency_cleanup_handler_dependencies 测试。""" - def test_returns_dict_with_required_keys(self, mock_session_factory): + def test_returns_dict_with_required_keys( + self, mock_session_factory, mock_cache_port + ): # Act deps = create_channel_idempotency_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) # Assert - assert set(deps.keys()) == { - "session_factory", - "idempotency_repo_factory", - "logger", - } + assert set(deps.keys()) == {"idempotency_repo", "cache_port", "logger"} + assert deps["cache_port"] is mock_cache_port _assert_logger_port(deps["logger"]) - def test_idempotency_repo_factory_creates_persistence_adapter( - self, mock_session_factory, mock_db + def test_idempotency_repo_is_persistence_adapter( + self, mock_session_factory, mock_cache_port ): - # Arrange + # Act deps = create_channel_idempotency_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ) - # Act - repo = deps["idempotency_repo_factory"](mock_db) - # Assert - assert isinstance(repo, ChannelPersistenceAdapter) + assert isinstance(deps["idempotency_repo"], ChannelPersistenceAdapter) # --------------------------------------------------------------------------- @@ -474,34 +446,36 @@ class TestHandlerDependenciesConsistency: """所有 worker 进程 handler 依赖工厂的一致性约束。""" def test_all_factories_produce_logger_port_compliant_instance( - self, mock_session_factory, mock_arq_pool + self, mock_session_factory, mock_arq_pool, mock_cache_port ): """8 个工厂的 logger 字段都为 StructuredLoggerAdapter(满足 LoggerPort)。""" # Act all_deps = [ create_channel_session_inactive_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ), create_channel_outbox_terminal_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ), create_channel_pairing_terminal_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ), create_channel_outbox_recovery_handler_dependencies( - mock_session_factory, arq_pool=mock_arq_pool + mock_session_factory, + arq_pool=mock_arq_pool, + cache_port=mock_cache_port, ), create_channel_pairing_expiration_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ), create_channel_audit_log_retention_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ), create_channel_content_review_retention_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ), create_channel_idempotency_cleanup_handler_dependencies( - mock_session_factory + mock_session_factory, cache_port=mock_cache_port ), ] @@ -509,5 +483,5 @@ class TestHandlerDependenciesConsistency: for deps in all_deps: assert "logger" in deps _assert_logger_port(deps["logger"]) - # 每个工厂都共享同一 session_factory - assert deps["session_factory"] is mock_session_factory + # 每个工厂都共享同一 cache_port + assert deps["cache_port"] is mock_cache_port diff --git a/backend/test/unit/channels/infrastructure/test_factory.py b/backend/test/unit/channels/infrastructure/test_factory.py index a6c44b87..8b401dd2 100644 --- a/backend/test/unit/channels/infrastructure/test_factory.py +++ b/backend/test/unit/channels/infrastructure/test_factory.py @@ -332,9 +332,7 @@ class TestRegisterBuiltinEventSubscribers: assert event_type in event_bus._subscriptions assert len(event_bus._subscriptions[event_type]) == 1 - def test_registers_config_event_pair_for_whitelist_and_route_match( - self, fake_logger - ): + def test_registers_config_event_pair_for_whitelist_and_route_match(self, fake_logger): """ConfigChanged/ConfigRollback 各注册 WhitelistConfigHandler + RouteMatchCacheHandler 共 2 个订阅。""" # Arrange @@ -375,9 +373,7 @@ class TestRegisterBuiltinEventSubscribers: assert "PairingApproved" in event_bus._subscriptions assert len(event_bus._subscriptions["PairingApproved"]) == 1 - def test_registers_channel_degraded_and_recovered_subscribers( - self, fake_logger - ): + def test_registers_channel_degraded_and_recovered_subscribers(self, fake_logger): # Arrange event_bus = self._make_event_bus(fake_logger) @@ -396,9 +392,7 @@ class TestRegisterBuiltinEventSubscribers: assert len(event_bus._subscriptions["ChannelDegraded"]) == 1 assert len(event_bus._subscriptions["ChannelRecovered"]) == 1 - def test_registers_three_channel_event_broadcaster_subscribers( - self, fake_logger - ): + def test_registers_three_channel_event_broadcaster_subscribers(self, fake_logger): """ChannelSessionUpdated/ChannelMessageReceived/ChannelMessageSent → 广播器。""" # Arrange event_bus = self._make_event_bus(fake_logger) @@ -501,9 +495,7 @@ def _make_plugin_loading_core() -> _PluginLoadingCore: class TestRegisterDiSingletons: """``_register_di_singletons`` 注册全量 DI 单例到容器。""" - def test_registers_core_registries_and_shared_dependencies( - self, fake_logger - ): + def test_registers_core_registries_and_shared_dependencies(self, fake_logger): # Arrange di = DependencyInjectionContainer() core = _make_plugin_loading_core() @@ -513,7 +505,7 @@ class TestRegisterDiSingletons: di_container=di, core=core, persistence_port=MagicMock(spec=PersistencePort), - persistence_db=MagicMock(), + session_factory=MagicMock(), queue_port=MagicMock(spec=QueuePort), config_port=MagicMock(spec=ConfigPort), stage_slot_injector=MagicMock(spec=StageSlotInjector), @@ -539,9 +531,7 @@ class TestRegisterDiSingletons: assert di.resolve(SensitiveFieldRegistry) is core.sensitive_registry assert di.resolve(ConfigScopeRegistry) is core.config_scope_registry - def test_registers_command_and_identity_resolver_registries( - self, fake_logger - ): + def test_registers_command_and_identity_resolver_registries(self, fake_logger): """命令注册表与身份解析器注册表由本函数构造并注册。""" # Arrange di = DependencyInjectionContainer() @@ -552,7 +542,7 @@ class TestRegisterDiSingletons: di_container=di, core=core, persistence_port=MagicMock(spec=PersistencePort), - persistence_db=MagicMock(), + session_factory=MagicMock(), queue_port=MagicMock(spec=QueuePort), config_port=MagicMock(spec=ConfigPort), stage_slot_injector=MagicMock(spec=StageSlotInjector), @@ -568,9 +558,7 @@ class TestRegisterDiSingletons: identity_registry = di.resolve(IdentityResolverRegistry) assert isinstance(identity_registry, IdentityResolverRegistry) - def test_registers_app_level_ports_and_orchestration_components( - self, fake_logger - ): + def test_registers_app_level_ports_and_orchestration_components(self, fake_logger): # Arrange di = DependencyInjectionContainer() core = _make_plugin_loading_core() @@ -587,7 +575,7 @@ class TestRegisterDiSingletons: di_container=di, core=core, persistence_port=persistence_port, - persistence_db=MagicMock(), + session_factory=MagicMock(), queue_port=queue_port, config_port=config_port, stage_slot_injector=stage_slot_injector, @@ -613,14 +601,14 @@ class TestRegisterDiSingletons: # Arrange di = DependencyInjectionContainer() core = _make_plugin_loading_core() - persistence_db = MagicMock() + session_factory = MagicMock() # Act _register_di_singletons( di_container=di, core=core, persistence_port=MagicMock(spec=PersistencePort), - persistence_db=persistence_db, + session_factory=session_factory, queue_port=MagicMock(spec=QueuePort), config_port=MagicMock(spec=ConfigPort), stage_slot_injector=MagicMock(spec=StageSlotInjector), @@ -639,9 +627,7 @@ class TestRegisterDiSingletons: # ContentReviewRepositoryPort 与 ContentReviewRepositoryAdapter 绑定到同一实例 assert review_port is review_adapter - def test_registers_redis_arq_and_execution_port_singletons( - self, fake_logger - ): + def test_registers_redis_arq_and_execution_port_singletons(self, fake_logger): # Arrange di = DependencyInjectionContainer() core = _make_plugin_loading_core() @@ -651,7 +637,7 @@ class TestRegisterDiSingletons: di_container=di, core=core, persistence_port=MagicMock(spec=PersistencePort), - persistence_db=MagicMock(), + session_factory=MagicMock(), queue_port=MagicMock(spec=QueuePort), config_port=MagicMock(spec=ConfigPort), stage_slot_injector=MagicMock(spec=StageSlotInjector), @@ -665,8 +651,8 @@ class TestRegisterDiSingletons: # (由 core 直接注册) assert di.resolve(AgentRunExecutionPort) is core.execution_port_impl - def test_persistence_db_not_registered_as_singleton(self, fake_logger): - """persistence_db 是请求级资源,不应注册为 AsyncSession 单例。""" + def test_async_session_not_registered_as_singleton(self, fake_logger): + """AsyncSession 是请求级资源,不应注册为单例;session_factory 注册为单例供按需创建。""" # Arrange di = DependencyInjectionContainer() core = _make_plugin_loading_core() @@ -677,7 +663,7 @@ class TestRegisterDiSingletons: di_container=di, core=core, persistence_port=MagicMock(spec=PersistencePort), - persistence_db=MagicMock(), + session_factory=MagicMock(), queue_port=MagicMock(spec=QueuePort), config_port=MagicMock(spec=ConfigPort), stage_slot_injector=MagicMock(spec=StageSlotInjector), diff --git a/backend/test/unit/channels/infrastructure/test_host_shutdown.py b/backend/test/unit/channels/infrastructure/test_host_shutdown.py index f81f12e1..4c5643c7 100644 --- a/backend/test/unit/channels/infrastructure/test_host_shutdown.py +++ b/backend/test/unit/channels/infrastructure/test_host_shutdown.py @@ -229,9 +229,7 @@ class TestTeardownPipelines: # Arrange stage_slot_registry = MagicMock(spec=StageSlotRegistry) # 模拟 inbound 有 2 个槽位,outbound 1 个,control 0 个 - stage_slot_registry.findByPipeline = MagicMock( - side_effect=[["slot1", "slot2"], ["slot3"], []] - ) + stage_slot_registry.findByPipeline = MagicMock(side_effect=[["slot1", "slot2"], ["slot3"], []]) shutdown = _make_shutdown(fake_logger, stage_slot_registry=stage_slot_registry) # Act @@ -239,9 +237,7 @@ class TestTeardownPipelines: # Assert: info 日志携带槽位数量 info_calls = fake_logger.info.call_args_list - teardown_call = next( - c for c in info_calls if "管道已拆除" in c.args[0] - ) + teardown_call = next(c for c in info_calls if "管道已拆除" in c.args[0]) assert teardown_call.kwargs.get("inbound_slot_count") == 2 assert teardown_call.kwargs.get("outbound_slot_count") == 1 assert teardown_call.kwargs.get("control_slot_count") == 0 @@ -251,9 +247,7 @@ class TestTeardownPipelines: """clear 抛异常时记录告警并继续(不中断 shutdown)。""" # Arrange stage_slot_registry = MagicMock(spec=StageSlotRegistry) - stage_slot_registry.findByPipeline = MagicMock( - side_effect=RuntimeError("registry corrupted") - ) + stage_slot_registry.findByPipeline = MagicMock(side_effect=RuntimeError("registry corrupted")) shutdown = _make_shutdown(fake_logger, stage_slot_registry=stage_slot_registry) # Act: 不抛异常即证明异常被吞并 @@ -339,9 +333,7 @@ class TestStopPlugins: plugin_registry = MagicMock(spec=PluginRegistry) # STARTED 返回 [a, b],PAUSED 返回 [] plugin_registry.listPluginsByState = MagicMock( - side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b] - if state == LifecycleState.STARTED - else [] + side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b] if state == LifecycleState.STARTED else [] ) # resolver 返回 [a, b](a 先,b 后),反向后为 [b, a] @@ -349,9 +341,7 @@ class TestStopPlugins: resolver.resolve = MagicMock(return_value=[manifest_a, manifest_b]) plugin_lifecycle = MagicMock(spec=PluginLifecycleManager) - plugin_lifecycle.stop = AsyncMock( - return_value=LifecycleResult(state="stopped") - ) + plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped")) shutdown = _make_shutdown( fake_logger, @@ -382,9 +372,7 @@ class TestStopPlugins: plugin_registry = MagicMock(spec=PluginRegistry) plugin_registry.listPluginsByState = MagicMock( - side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b] - if state == LifecycleState.STARTED - else [] + side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b] if state == LifecycleState.STARTED else [] ) # resolver.resolve 抛异常 @@ -392,9 +380,7 @@ class TestStopPlugins: resolver.resolve = MagicMock(side_effect=RuntimeError("cycle detected")) plugin_lifecycle = MagicMock(spec=PluginLifecycleManager) - plugin_lifecycle.stop = AsyncMock( - return_value=LifecycleResult(state="stopped") - ) + plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped")) shutdown = _make_shutdown( fake_logger, @@ -424,18 +410,14 @@ class TestStopPlugins: plugin_registry = MagicMock(spec=PluginRegistry) # STARTED 返回 [],PAUSED 返回 [paused-plugin] plugin_registry.listPluginsByState = MagicMock( - side_effect=lambda state: [plugin_manifest_paused] - if state == LifecycleState.PAUSED - else [] + side_effect=lambda state: [plugin_manifest_paused] if state == LifecycleState.PAUSED else [] ) resolver = MagicMock(spec=PluginDependencyResolver) resolver.resolve = MagicMock(return_value=[manifest_paused]) plugin_lifecycle = MagicMock(spec=PluginLifecycleManager) - plugin_lifecycle.stop = AsyncMock( - return_value=LifecycleResult(state="stopped") - ) + plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped")) shutdown = _make_shutdown( fake_logger, @@ -460,9 +442,7 @@ class TestStopSinglePlugin: # Arrange manifest = _make_manifest("ok-plugin") plugin_lifecycle = MagicMock(spec=PluginLifecycleManager) - plugin_lifecycle.stop = AsyncMock( - return_value=LifecycleResult(state="stopped") - ) + plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped")) shutdown = _make_shutdown(fake_logger, plugin_lifecycle=plugin_lifecycle) # Act @@ -481,11 +461,7 @@ class TestStopSinglePlugin: # Arrange manifest = _make_manifest("bad-plugin") plugin_lifecycle = MagicMock(spec=PluginLifecycleManager) - plugin_lifecycle.stop = AsyncMock( - return_value=LifecycleResult( - state="failed", error="cleanup failed" - ) - ) + plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="failed", error="cleanup failed")) shutdown = _make_shutdown(fake_logger, plugin_lifecycle=plugin_lifecycle) # Act @@ -563,8 +539,9 @@ class TestCloseDrivenAdapters: # Act await shutdown._closeAllDrivenAdapters() - # Assert: 记录跳过日志,不调用 aclose - fake_logger.info.assert_any_call( + # Assert: 记录跳过日志(DEBUG 级别,因跳过是预期行为而非异常事件), + # 不调用 aclose + fake_logger.debug.assert_any_call( "应用级被驱动适配器未实现 CloseablePort,跳过", adapter_type=type(non_closeable).__name__, ) @@ -601,9 +578,7 @@ class TestCloseDrivenAdapters: plugin_adapters = MagicMock() plugin_adapters.close = AsyncMock() plugin_registry = MagicMock(spec=PluginRegistry) - plugin_registry.listPluginAdapters = MagicMock( - return_value=[("feishu", plugin_adapters)] - ) + plugin_registry.listPluginAdapters = MagicMock(return_value=[("feishu", plugin_adapters)]) shutdown = _make_shutdown( fake_logger, plugin_registry=plugin_registry, diff --git a/backend/test/unit/channels/infrastructure/test_scheduler.py b/backend/test/unit/channels/infrastructure/test_scheduler.py index ee2f2d4f..0f540177 100644 --- a/backend/test/unit/channels/infrastructure/test_scheduler.py +++ b/backend/test/unit/channels/infrastructure/test_scheduler.py @@ -86,7 +86,9 @@ def patched_scheduler(monkeypatch): """替换 ``scheduler`` 模块的全部外部依赖,返回 ``_PatchTracker``。 替换内容: - - ``pg_manager.get_async_session_context``:返回 sentinel mock + - ``pg_manager.AsyncSession``:返回 sentinel mock(scheduler 通过 + ``pg_manager.AsyncSession`` 获取 ``async_sessionmaker`` 实例作为 + ``session_factory`` 透传给各 handler 工厂) - ``get_arq_pool``:AsyncMock 返回 sentinel mock - 8 个 ``create_channel_*_handler_dependencies`` 工厂:返回空字典, 调用参数记录到 ``factory_calls`` @@ -94,16 +96,19 @@ def patched_scheduler(monkeypatch): """ tracker = _PatchTracker() - # 替换 pg_manager.get_async_session_context + # 替换 pg_manager.AsyncSession:scheduler 源码通过 + # ``session_factory = pg_manager.AsyncSession`` 获取 ``async_sessionmaker`` + # 实例(callable,调用返回 ``AsyncSession``),符合 ``Callable[[], AsyncSession]`` 协议 monkeypatch.setattr( scheduler_module.pg_manager, - "get_async_session_context", + "AsyncSession", tracker.session_factory, ) # 替换 get_arq_pool async def _fake_get_arq_pool(): return tracker.arq_pool + monkeypatch.setattr(scheduler_module, "get_arq_pool", _fake_get_arq_pool) # 替换 8 个 handler 类:实例化时返回带正确 name 的 mock @@ -112,18 +117,23 @@ def patched_scheduler(monkeypatch): def _make_side_effect(h_name: str): def _instantiate(**kwargs): return _make_handler_mock(h_name) + return _instantiate + class_mock = MagicMock(side_effect=_make_side_effect(handler_name)) tracker.handler_class_mocks[class_name] = class_mock monkeypatch.setattr(scheduler_module, class_name, class_mock) # 替换 8 个工厂函数:返回空字典,记录调用参数 for factory_name, handler_name in _FACTORY_TO_NAME.items(): + def _make_factory(h_name: str): def _factory(*args, **kwargs): tracker.factory_calls[h_name].append({"args": args, "kwargs": kwargs}) return {} + return _factory + monkeypatch.setattr(scheduler_module, factory_name, _make_factory(handler_name)) return tracker @@ -179,9 +189,7 @@ class TestRegisterSchedulerHandlers: # Assert: 每个工厂的首个位置参数为 session_factory for handler_name in EXPECTED_HANDLER_NAMES: factory_calls = patched_scheduler.factory_calls[handler_name] - assert len(factory_calls) == 1, ( - f"factory for {handler_name} called {len(factory_calls)} times" - ) + assert len(factory_calls) == 1, f"factory for {handler_name} called {len(factory_calls)} times" assert factory_calls[0]["args"][0] is patched_scheduler.session_factory @pytest.mark.asyncio @@ -240,9 +248,7 @@ class TestRegisterSchedulerHandlers: # Assert: 每个 handler 类的 mock 被调用一次 for class_name in _HANDLER_CLASS_TO_NAME: class_mock = patched_scheduler.handler_class_mocks[class_name] - assert class_mock.call_count == 1, ( - f"{class_name} instantiated {class_mock.call_count} times" - ) + assert class_mock.call_count == 1, f"{class_name} instantiated {class_mock.call_count} times" # 工厂返回空字典,所以实例化参数为空 kwargs call_kwargs = class_mock.call_args.kwargs assert call_kwargs == {} diff --git a/backend/test/unit/channels/plugins/wechat_woc/adapters/test_inbound_adapter.py b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_inbound_adapter.py index baa45832..a06fc3b3 100644 --- a/backend/test/unit/channels/plugins/wechat_woc/adapters/test_inbound_adapter.py +++ b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_inbound_adapter.py @@ -88,7 +88,7 @@ def _make_text_payload() -> dict: "type": 1, "content": "hello", "create_time": "1000", - "session_type": "single", + "session_type": "p2p", } @@ -102,7 +102,7 @@ def _make_image_payload() -> dict: "type": 3, "content": "", "create_time": "1000", - "session_type": "single", + "session_type": "p2p", } @@ -668,6 +668,38 @@ class TestNormalizeInboundMessageTypes: await adapter.normalizeInbound(raw) assert exc_info.value.field == "talker" + @pytest.mark.asyncio + async def test_p2p_session_type_is_accepted(self): + # 回归测试:bridge 返回的单聊 session_type 为 "p2p"(非 "single"), + # 必须通过校验,不应抛出 session_type_invalid。 + client = _make_client() + adapter = WeChatWocInboundAdapter(client, _make_logger()) + payload = _make_text_payload() + payload["session_type"] = "p2p" + raw = _make_raw_event(payload, account_id="acct-1") + + # Act + content = await adapter.normalizeInbound(raw) + + # Assert + assert content.metadata["session_type"] == "p2p" + + @pytest.mark.asyncio + async def test_invalid_session_type_raises_validation_error(self): + # Arrange - session_type 不在合法枚举内 + client = _make_client() + adapter = WeChatWocInboundAdapter(client, _make_logger()) + payload = _make_text_payload() + payload["session_type"] = "unknown" + raw = _make_raw_event(payload, account_id="acct-1") + + # Act / Assert + from yuxi.channels.contract.errors import ValidationError + + with pytest.raises(ValidationError) as exc_info: + await adapter.normalizeInbound(raw) + assert exc_info.value.field == "session_type" + # ------------------------------------------------------------------ # normalizeInbound: trace_id 通过 ContextVar 隐式传递(P1-10) @@ -706,7 +738,7 @@ class TestNormalizeInboundTraceId: msg_type=1, sender="sender-1", is_sender=0, - session_type="single", + session_type="p2p", text_preview="hello", attachment_count=0, trace_id="ctx-trace-id-123",