diff --git a/backend/test/integration/api/channels/conftest.py b/backend/test/integration/api/channels/conftest.py index 297eab64..36b153e7 100644 --- a/backend/test/integration/api/channels/conftest.py +++ b/backend/test/integration/api/channels/conftest.py @@ -28,8 +28,6 @@ NON_EXISTENT_GROUP_ID = "nonexistent_group_00000000" NON_EXISTENT_PLUGIN_ID = "nonexistent_plugin_00000000" NON_EXISTENT_CONFIG_KEY = "nonexistent_config_key_0000" NON_EXISTENT_LOG_ID = "nonexistent_log_00000000" -NON_EXISTENT_TASK_ID = "nonexistent_task_00000000" -NON_EXISTENT_REPORT_ID = "nonexistent_report_00000000" NON_EXISTENT_CHECK_ID = "nonexistent_check_00000000" NON_EXISTENT_REVIEW_ID = "nonexistent_review_00000000" diff --git a/backend/test/integration/api/channels/test_reports_router.py b/backend/test/integration/api/channels/test_reports_router.py deleted file mode 100644 index 5492e2fc..00000000 --- a/backend/test/integration/api/channels/test_reports_router.py +++ /dev/null @@ -1,135 +0,0 @@ -"""Integration tests for channels reports_router endpoints.""" - -from __future__ import annotations - -import httpx -import pytest - -from .conftest import BASE_URL, NON_EXISTENT_REPORT_ID, NON_EXISTENT_TASK_ID - -pytestmark = [pytest.mark.asyncio, pytest.mark.integration] - -REPORTS_URL = f"{BASE_URL}/reports" -ONEOFF_URL = f"{REPORTS_URL}/oneoff" - - -# ============================================================================= -# === Auth three-tier for GET /reports/oneoff === -# ============================================================================= - - -async def test_list_oneoff_reports_requires_auth(test_client: httpx.AsyncClient): - # Act - response = await test_client.get(ONEOFF_URL) - # Assert - assert response.status_code == 401, response.text - - -async def test_list_oneoff_reports_requires_admin(test_client: httpx.AsyncClient, standard_user): - # Act - response = await test_client.get(ONEOFF_URL, headers=standard_user["headers"]) - # Assert - assert response.status_code == 403, response.text - - -async def test_admin_can_list_oneoff_reports(test_client: httpx.AsyncClient, admin_headers): - # Act - response = await test_client.get(ONEOFF_URL, headers=admin_headers) - # Assert - assert response.status_code == 200, response.text - payload = response.json() - assert payload["success"] is True - assert "data" in payload - - -# ============================================================================= -# === POST /reports/oneoff (create oneoff report) === -# ============================================================================= - - -async def test_admin_can_create_oneoff_report(test_client: httpx.AsyncClient, admin_headers): - # Arrange - body = {"report_type": "message_stats"} - # Act - response = await test_client.post(ONEOFF_URL, json=body, headers=admin_headers) - # Assert - assert response.status_code == 200, response.text - payload = response.json() - assert payload["success"] is True - assert "data" in payload - assert "task_id" in payload["data"] - - -async def test_create_oneoff_report_rejects_invalid_report_type(test_client: httpx.AsyncClient, admin_headers): - # Arrange — "invalid_type" is not in the Literal enum - body = {"report_type": "invalid_type"} - # Act - response = await test_client.post(ONEOFF_URL, json=body, headers=admin_headers) - # Assert - assert response.status_code == 422, response.text - - -# ============================================================================= -# === GET /reports/oneoff pagination validation === -# ============================================================================= - - -async def test_list_oneoff_reports_rejects_limit_zero(test_client: httpx.AsyncClient, admin_headers): - # Act — limit=0 violates ge=1 - response = await test_client.get(ONEOFF_URL, params={"limit": 0}, headers=admin_headers) - # Assert - assert response.status_code == 422, response.text - - -async def test_list_oneoff_reports_rejects_limit_over_max(test_client: httpx.AsyncClient, admin_headers): - # Act — limit=201 violates le=200 - response = await test_client.get(ONEOFF_URL, params={"limit": 201}, headers=admin_headers) - # Assert - assert response.status_code == 422, response.text - - -# ============================================================================= -# === Non-existent resource 404 paths === -# ============================================================================= - - -async def test_download_oneoff_report_returns_404_or_409_for_nonexistent(test_client: httpx.AsyncClient, admin_headers): - # Act — non-existent task returns 404 (not found) or 409 (not ready) - response = await test_client.get( - f"{ONEOFF_URL}/{NON_EXISTENT_TASK_ID}/download", - headers=admin_headers, - ) - # Assert - assert response.status_code in (404, 409), response.text - - -async def test_retry_oneoff_report_returns_404_for_nonexistent(test_client: httpx.AsyncClient, admin_headers): - # Act - response = await test_client.post( - f"{ONEOFF_URL}/{NON_EXISTENT_TASK_ID}/retry", - headers=admin_headers, - ) - # Assert - assert response.status_code == 404, response.text - - -async def test_get_report_returns_404_for_nonexistent(test_client: httpx.AsyncClient, admin_headers): - # Act - response = await test_client.get(f"{REPORTS_URL}/{NON_EXISTENT_REPORT_ID}", headers=admin_headers) - # Assert - assert response.status_code == 404, response.text - - -# ============================================================================= -# === Static path /oneoff not captured as {report_id} === -# ============================================================================= - - -async def test_static_oneoff_path_not_captured_as_report_id(test_client: httpx.AsyncClient, admin_headers): - # Act — GET /reports/oneoff must hit the list endpoint (200), not - # GET /reports/{report_id} with report_id="oneoff" (which would 404) - response = await test_client.get(ONEOFF_URL, headers=admin_headers) - # Assert - assert response.status_code == 200, response.text - payload = response.json() - assert payload["success"] is True 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 59113862..e3e86ee3 100644 --- a/backend/test/unit/channels/adapters/test_channel_persistence_adapter.py +++ b/backend/test/unit/channels/adapters/test_channel_persistence_adapter.py @@ -18,11 +18,13 @@ from yuxi.channels.adapters.channel_persistence_adapter import ( ChannelPersistenceAdapter, ) from yuxi.channels.contract.dtos.channel import ( + AccountFilter, ChannelAccount, ChannelSession, ChannelType, ) from yuxi.channels.contract.dtos.outbox import OutboxConfig +from yuxi.channels.contract.dtos.pairing import PairingStatus from yuxi.channels.contract.dtos.persistence import ( SaveChannelAccountCmd, SaveChannelSessionCmd, @@ -79,6 +81,13 @@ def _make_account_orm(*, account_id: str = "acc-1", channel_type: str = "feishu" orm.transport_cursor = "" orm.last_rotated_at = None orm.is_deleted = 0 + orm.version = 1 + orm.onboarding_status = "online" + orm.service_user_uid = None + orm.credential_ref = None + orm.credential_version = 0 + orm.last_error = None + orm.plugin_status = "stopped" return orm @@ -187,6 +196,7 @@ class TestChannelPersistenceAdapterAccount: cmd = UpdateChannelAccountCmd( channel_type=ChannelType("feishu"), account_id="acc-1", + expected_version=1, display_name="new-name", ) @@ -206,6 +216,7 @@ class TestChannelPersistenceAdapterAccount: cmd = UpdateChannelAccountCmd( channel_type=ChannelType("feishu"), account_id="acc-1", + expected_version=1, display_name="new-name", ) @@ -272,7 +283,9 @@ class TestChannelPersistenceAdapterAccount: adapter = _build_adapter(db, repos=_make_repos()) # Act - result = await adapter.findAccountsByFilter({"channel_type": "feishu"}) + result = await adapter.findAccountsByFilter( + AccountFilter(channel_type=ChannelType("feishu")) + ) # Assert assert len(result) == 1 @@ -644,3 +657,101 @@ class TestChannelPersistenceAdapterClose: # Assert db.close.assert_awaited_once() + + +def _make_pairing_orm(*, version: int = 1, status: str = "pending") -> MagicMock: + """构造 ChannelPairing ORM 桩。""" + orm = MagicMock() + orm.pairing_id = "pair-1" + orm.account_id = 1 + orm.peer_id = "peer-1" + orm.peer_name = "Peer" + orm.status = status + orm.approver_id = None + orm.approved_at = None + orm.rejected_at = None + orm.revoked_at = None + orm.expired_at = None + orm.reason = None + orm.expires_at = None + orm.requested_at = datetime(2026, 1, 1) + orm.created_at = datetime(2026, 1, 1) + orm.updated_at = datetime(2026, 1, 1) + orm.created_by = None + orm.updated_by = None + orm.version = version + return orm + + +def _make_repos_with_pairing(*, pairing_orm=None, account_orm=None) -> MagicMock: + """构造含 pairing 仓储的 Repositories 桩。""" + repos = _make_repos() + repos.pairing = MagicMock() + repos.pairing.get_by_pairing_id = AsyncMock(return_value=pairing_orm) + repos.pairing.update = AsyncMock(return_value=pairing_orm) + repos.account.get_by_id = AsyncMock(return_value=account_orm) + return repos + + +@pytest.mark.unit +class TestChannelPersistenceAdapterPairing: + """updatePairingStatus 乐观锁与状态写入测试。""" + + @pytest.mark.asyncio + async def test_update_raises_conflict_when_version_mismatch(self): + # Arrange: ORM version=2, 调用方传 expected_version=1 + orm = _make_pairing_orm(version=2) + repos = _make_repos_with_pairing(pairing_orm=orm) + adapter = _build_adapter(_make_db(), repos) + + # Act & Assert + with pytest.raises(ConflictError): + await adapter.updatePairingStatus( + "pair-1", + PairingStatus.APPROVED, + expected_version=1, + approver_id="admin-1", + ) + repos.pairing.update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_update_raises_not_found_when_pairing_missing(self): + # Arrange + repos = _make_repos_with_pairing(pairing_orm=None) + adapter = _build_adapter(_make_db(), repos) + + # Act & Assert + with pytest.raises(NotFoundError): + await adapter.updatePairingStatus( + "pair-1", + PairingStatus.APPROVED, + expected_version=1, + ) + + @pytest.mark.asyncio + async def test_update_approve_succeeds_and_increments_version(self): + # Arrange + orm = _make_pairing_orm(version=1) + account_orm = _make_account_orm() + repos = _make_repos_with_pairing(pairing_orm=orm, account_orm=account_orm) + adapter = _build_adapter(_make_db(), repos) + + # Act + result = await adapter.updatePairingStatus( + "pair-1", + PairingStatus.APPROVED, + expected_version=1, + approver_id="admin-1", + reason="ok", + updated_by="admin-1", + ) + + # Assert + assert result.pairing_id == "pair-1" + assert orm.status == "approved" + assert orm.approver_id == "admin-1" + assert orm.approved_at is not None + assert orm.reason == "ok" + assert orm.updated_by == "admin-1" + assert orm.version == 2 + repos.pairing.update.assert_awaited_once() 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 001db386..d893124c 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 @@ -50,6 +50,7 @@ def _make_db() -> MagicMock: db.rollback = AsyncMock() db.scalar = AsyncMock() db.execute = AsyncMock() + db.refresh = AsyncMock() return db @@ -117,6 +118,8 @@ def _make_orm(*, review_id: str = "rev-1") -> MagicMock: orm.reviewer = "user-1" orm.source = "manual_preview" orm.trace_id = "trace-1" + orm.created_at = datetime(2026, 1, 1, 0, 0, 0) + orm.updated_at = datetime(2026, 1, 1, 0, 0, 0) orm.is_deleted = 0 return orm @@ -141,6 +144,20 @@ def _make_result_with_one(one_value: Any) -> MagicMock: 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( @@ -154,8 +171,9 @@ 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 = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) record = _make_record() @@ -164,14 +182,15 @@ class TestContentReviewSaveReviewResult: # Assert db.add.assert_called_once() - db.flush.assert_awaited_once() db.commit.assert_awaited_once() + db.refresh.assert_awaited_once() + db.flush.assert_not_awaited() @pytest.mark.asyncio async def test_save_review_result_skips_commit_when_tx_provided(self): # Arrange db = _make_db() - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) record = _make_record() tx = MagicMock() @@ -186,9 +205,10 @@ class TestContentReviewSaveReviewResult: @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.flush = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig"))) - logger = MagicMock() + db.commit = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig"))) + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) record = _make_record() @@ -200,9 +220,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.flush = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = MagicMock() + db.commit = AsyncMock(side_effect=SQLAlchemyError("db failure")) + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) record = _make_record() @@ -221,7 +242,7 @@ class TestContentReviewQueryHistory: 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 = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -239,7 +260,7 @@ class TestContentReviewQueryHistory: result_mock = MagicMock() result_mock.scalars.return_value.all.return_value = [] db.execute = AsyncMock(return_value=result_mock) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -253,7 +274,7 @@ class TestContentReviewQueryHistory: # Arrange db = _make_db() db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert @@ -263,15 +284,17 @@ class TestContentReviewQueryHistory: @pytest.mark.asyncio async def test_query_history_translates_invalid_enum_to_dependency(self): # Arrange - # ORM 持有非法 channel_type 值(迁移残留 / 数据损坏), + # ORM 持有非法 verdict 值(迁移残留 / 数据损坏), # ``_enum`` 应翻译 ValueError 为 DependencyError,禁止穿透核心层(INV-7) + # 注:``ChannelType`` 是 ``str`` 子类而非 Enum,``ChannelType("unknown")`` + # 不抛异常;真正受 ``_enum`` 保护的是 verdict / source / resource_type db = _make_db() orm = _make_orm(review_id="rev-1") - orm.channel_type = "unknown_channel" + orm.verdict = "unknown_verdict" result_mock = MagicMock() result_mock.scalars.return_value.all.return_value = [orm] db.execute = AsyncMock(return_value=result_mock) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert @@ -288,7 +311,7 @@ class TestContentReviewCountHistory: result_mock = MagicMock() result_mock.scalar.return_value = 42 db.execute = AsyncMock(return_value=result_mock) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -304,7 +327,7 @@ class TestContentReviewCountHistory: result_mock = MagicMock() result_mock.scalar.return_value = None db.execute = AsyncMock(return_value=result_mock) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -318,7 +341,7 @@ class TestContentReviewCountHistory: # Arrange db = _make_db() db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert @@ -331,9 +354,12 @@ 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() - db.scalar = AsyncMock(return_value=_make_orm(review_id="rev-1")) - logger = MagicMock() + orm = _make_orm(review_id="rev-1") + db.execute = AsyncMock(return_value=_make_scalar_result(orm)) + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -347,8 +373,8 @@ class TestContentReviewGetDetail: async def test_get_detail_returns_none_when_not_found(self): # Arrange db = _make_db() - db.scalar = AsyncMock(return_value=None) - logger = MagicMock() + db.execute = AsyncMock(return_value=_make_scalar_result(None)) + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -361,8 +387,8 @@ class TestContentReviewGetDetail: async def test_get_detail_translates_sqlalchemy_error(self): # Arrange db = _make_db() - db.scalar = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = MagicMock() + db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert @@ -377,8 +403,8 @@ class TestContentReviewGetDetail: db = _make_db() orm = _make_orm(review_id="rev-1") orm.verdict = "unknown_verdict" - db.scalar = AsyncMock(return_value=orm) - logger = MagicMock() + db.execute = AsyncMock(return_value=_make_scalar_result(orm)) + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert @@ -413,7 +439,7 @@ class TestContentReviewGetAnalytics: trend_result, ] ) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -444,7 +470,7 @@ class TestContentReviewGetAnalytics: trend_result, ] ) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -459,7 +485,7 @@ class TestContentReviewGetAnalytics: # Arrange db = _make_db() db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert @@ -481,6 +507,7 @@ class TestContentReviewGetStats: 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() @@ -492,7 +519,7 @@ class TestContentReviewGetStats: trend_result, ] ) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -503,7 +530,7 @@ class TestContentReviewGetStats: assert result.total_reviews == 20 assert result.pass_count == 12 assert result.block_count == 5 - assert result.avg_decision_seconds == 0.0 + assert result.avg_decision_seconds == 42.5 assert len(result.trend) == 1 assert isinstance(result.trend[0], ReviewTrendPoint) @@ -517,6 +544,7 @@ class TestContentReviewGetStats: 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() @@ -528,7 +556,7 @@ class TestContentReviewGetStats: trend_result, ] ) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -544,7 +572,7 @@ class TestContentReviewGetStats: # Arrange db = _make_db() db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert @@ -562,19 +590,21 @@ class TestContentReviewUpdateVerdict: result_mock = MagicMock() result_mock.scalar_one_or_none.return_value = orm db.execute = AsyncMock(return_value=result_mock) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # 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.flush.assert_awaited_once() db.commit.assert_awaited_once() + db.refresh.assert_awaited_once() + db.flush.assert_not_awaited() @pytest.mark.asyncio async def test_update_verdict_returns_none_when_not_found(self): @@ -583,7 +613,7 @@ class TestContentReviewUpdateVerdict: result_mock = MagicMock() result_mock.scalar_one_or_none.return_value = None db.execute = AsyncMock(return_value=result_mock) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act @@ -601,7 +631,7 @@ class TestContentReviewUpdateVerdict: result_mock = MagicMock() result_mock.scalar_one_or_none.return_value = orm db.execute = AsyncMock(return_value=result_mock) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) tx = MagicMock() @@ -617,7 +647,7 @@ class TestContentReviewUpdateVerdict: # Arrange db = _make_db() db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure")) - logger = MagicMock() + logger = AsyncMock() adapter = ContentReviewRepositoryAdapter(db, logger) # Act / Assert diff --git a/backend/test/unit/channels/adapters/test_conversation_adapter.py b/backend/test/unit/channels/adapters/test_conversation_adapter.py index 1f5c090e..fa502b97 100644 --- a/backend/test/unit/channels/adapters/test_conversation_adapter.py +++ b/backend/test/unit/channels/adapters/test_conversation_adapter.py @@ -73,12 +73,17 @@ def _make_operator() -> Operator: ) -def _make_save_cmd(*, conversation_id: str = "100") -> SaveMessageCmd: +def _make_save_cmd( + *, + conversation_id: str = "100", + role: str = "assistant", + content: str = "hello", +) -> SaveMessageCmd: """构造保存消息命令。""" return SaveMessageCmd( conversation_id=conversation_id, - role="assistant", - content="hello", + role=role, + content=content, channel_status="sent", channel_msg_id="cmid-1", ) @@ -233,6 +238,7 @@ def _build_adapter( message_repo = MagicMock() conversation_repo = MagicMock() session_repo = MagicMock() + logger = AsyncMock() with ( patch( "yuxi.channels.adapters.conversation_adapter.ChannelMessageRepository", @@ -247,7 +253,7 @@ def _build_adapter( return_value=session_repo, ), ): - adapter = ConversationAdapter(db, MagicMock()) + adapter = ConversationAdapter(db, logger) return adapter, message_repo, conversation_repo, session_repo @@ -316,6 +322,9 @@ class TestConversationSaveMessage: assert message.message_id == "200" message_repo.init_channel_fields.assert_awaited_once() assert message_repo.init_channel_fields.call_args.kwargs["commit"] is True + # init_channel_fields 的 commit 会过期 conversation 关系, + # 需 refresh 重新加载以避免 lazy-load 抛 MissingGreenlet + db.refresh.assert_awaited_once_with(message_orm, ["conversation"]) @pytest.mark.asyncio async def test_save_message_skips_commit_when_tx_provided(self): @@ -323,7 +332,8 @@ class TestConversationSaveMessage: db = _make_db() db.execute = AsyncMock(return_value=_make_scalar_result(_make_message_orm())) adapter, message_repo, _, _ = _build_adapter(db) - message_repo.init_channel_fields = AsyncMock(return_value=_make_message_orm()) + message_orm = _make_message_orm() + message_repo.init_channel_fields = AsyncMock(return_value=message_orm) tx = MagicMock() # Act @@ -331,17 +341,72 @@ class TestConversationSaveMessage: # Assert db.commit.assert_not_awaited() + # 即使 flush 也需 refresh conversation(统一防御 lazy-load) + db.refresh.assert_awaited_once_with(message_orm, ["conversation"]) @pytest.mark.asyncio - async def test_save_message_raises_not_found_when_no_latest_message(self): + async def test_save_message_raises_not_found_when_init_returns_none(self): + """入站回复路径:init_channel_fields 返回 None(消息被并发删除)时抛 NotFoundError。""" + db = _make_db() + found_orm = _make_message_orm(message_id=5) + db.execute = AsyncMock(return_value=_make_scalar_result(found_orm)) + adapter, message_repo, _, _ = _build_adapter(db) + message_repo.init_channel_fields = AsyncMock(return_value=None) + + with pytest.raises(NotFoundError) as exc_info: + await adapter.saveMessage(_make_save_cmd()) + + assert exc_info.value.details["id"] == "5" + # 不应调用 refresh(latest is None) + db.refresh.assert_not_awaited() + + @pytest.mark.asyncio + async def test_save_message_creates_new_when_no_latest_message(self): + """管理员直发路径:会话无消息时创建新 Message 记录。 + + 验证 ``saveMessage`` 在 ``_findLatestMessage`` 返回 None 时: + 1. 调用 ``db.add`` 添加新 MessageORM + 2. 调用 ``db.commit``(``commit=True``) + 3. 返回的 Message 携带 cmd 中的 ``content`` 与 ``channel_status`` + """ # Arrange db = _make_db() db.execute = AsyncMock(return_value=_make_scalar_result(None)) adapter, message_repo, _, _ = _build_adapter(db) + cmd = _make_save_cmd(role="admin", content="亲亲") - # Act / Assert - with pytest.raises(NotFoundError): - await adapter.saveMessage(_make_save_cmd()) + # Act + message = await adapter.saveMessage(cmd) + + # Assert + assert isinstance(message, Message) + # 新增 MessageORM 至 session + db.add.assert_called_once() + added_orm = db.add.call_args.args[0] + assert added_orm.role == "admin" + assert added_orm.content == "亲亲" + assert added_orm.channel_status == "sent" + # commit=True 自主提交 + db.commit.assert_awaited_once() + # 不应调用 init_channel_fields(新消息在构造时已设置渠道字段) + message_repo.init_channel_fields.assert_not_called() + + @pytest.mark.asyncio + async def test_save_message_creates_new_flushes_when_tx_provided(self): + """管理员直发路径且事务由应用层控制时仅 flush。""" + # Arrange + db = _make_db() + db.execute = AsyncMock(return_value=_make_scalar_result(None)) + adapter, _, _, _ = _build_adapter(db) + tx = MagicMock() + + # Act + await adapter.saveMessage(_make_save_cmd(), tx=tx) + + # Assert + db.add.assert_called_once() + db.flush.assert_awaited_once() + db.commit.assert_not_awaited() @pytest.mark.asyncio async def test_save_message_raises_not_found_when_invalid_conversation_id(self): diff --git a/backend/test/unit/channels/adapters/test_mappers.py b/backend/test/unit/channels/adapters/test_mappers.py index 1969e06c..d388094a 100644 --- a/backend/test/unit/channels/adapters/test_mappers.py +++ b/backend/test/unit/channels/adapters/test_mappers.py @@ -46,6 +46,13 @@ def _make_account_orm() -> MagicMock: orm.updated_at = datetime(2026, 1, 1) orm.transport_cursor = "cursor-1" orm.last_rotated_at = None + orm.onboarding_status = "online" + orm.service_user_uid = None + orm.credential_ref = None + orm.credential_version = 0 + orm.last_error = None + orm.plugin_status = "stopped" + orm.version = 1 return orm @@ -81,6 +88,7 @@ def _make_pairing_orm() -> MagicMock: orm.approved_at = None orm.rejected_at = None orm.revoked_at = None + orm.expired_at = None orm.reason = None orm.created_at = datetime(2026, 1, 1) orm.updated_at = datetime(2026, 1, 1) diff --git a/backend/test/unit/channels/adapters/test_redis_config_adapter.py b/backend/test/unit/channels/adapters/test_redis_config_adapter.py index ef79fd98..24deae16 100644 --- a/backend/test/unit/channels/adapters/test_redis_config_adapter.py +++ b/backend/test/unit/channels/adapters/test_redis_config_adapter.py @@ -629,11 +629,7 @@ class TestRedisConfigScopeValidation: # Act - 以 ACCOUNT 作用域读取,回退到 GLOBAL result = await adapter.get("admin_users", scope=ConfigScope.ACCOUNT, target="acc-1") - # Assert - ACCOUNT 读取时声明一致无告警;GLOBAL 回退时声明 ACCOUNT 不一致产生告警 + # Assert - ACCOUNT 读取时声明一致无告警;GLOBAL 回退路径跳过 F-02 校验 + # (回退是设计内降级行为,非作用域使用错误),无告警 assert result.value == {"v": "global"} - # 第一次 _build_key (ACCOUNT) 无告警,第二次 _build_key (GLOBAL) 产生告警 - logger.warn.assert_awaited() - call_kwargs = logger.warn.call_args - assert call_kwargs.kwargs["key"] == "admin_users" - assert call_kwargs.kwargs["declared_scope"] == "account" - assert call_kwargs.kwargs["actual_scope"] == "global" + logger.warn.assert_not_awaited() 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 799c22f8..43d7af47 100644 --- a/backend/test/unit/channels/adapters/test_route_binding_repository.py +++ b/backend/test/unit/channels/adapters/test_route_binding_repository.py @@ -1,16 +1,16 @@ """yuxi.channels.adapters.channel_persistence_adapter 路由绑定仓储契约测试。 覆盖 ``ChannelPersistenceAdapter`` 实现 ``RouteBindingRepositoryPort`` 的 -6 个方法(saveRouteBinding / updateRouteBinding / getRouteBinding / -listRouteBindings / deleteRouteBinding / listEnabledByAccount),patch +5 个方法(saveRouteBinding / updateRouteBinding / getRouteBinding / +listRouteBindings / deleteRouteBinding),patch ``create_repositories`` 返回 mock 聚合,不连接真实 DB。 测试契约点: - saveRouteBinding 正常创建返回 RouteBindingRule,IntegrityError → ConflictError -- updateRouteBinding 局部更新,binding_id 不存在 → NotFoundError +- updateRouteBinding 局部更新,binding_id 不存在 → NotFoundError, + 乐观锁 version 不匹配 → ConflictError - deleteRouteBinding 软删除,不存在 → NotFoundError - listRouteBindings 按过滤条件查询 -- listEnabledByAccount 仅返回启用规则 """ from __future__ import annotations @@ -57,7 +57,6 @@ def _make_repos() -> MagicMock: repos.route_binding.soft_delete = AsyncMock() repos.route_binding.list = AsyncMock(return_value=[]) repos.route_binding.count = AsyncMock(return_value=0) - repos.route_binding.list_enabled_by_account = AsyncMock(return_value=[]) return repos @@ -214,7 +213,7 @@ class TestSaveRouteBinding: class TestUpdateRouteBinding: @pytest.mark.asyncio async def test_update_route_binding_updates_agent_binding(self): - """覆盖点:仅更新 cmd 中非 None 字段(agent_binding)。""" + """覆盖点:仅更新 cmd 中非 None 字段(agent_binding),version 递增。""" # Arrange db = _make_db() repos = _make_repos() @@ -225,6 +224,7 @@ class TestUpdateRouteBinding: adapter = _build_adapter(db, repos) cmd = UpdateRouteBindingCmd( binding_id="bnd-001", + expected_version=1, agent_binding="agent-new", ) @@ -240,6 +240,31 @@ class TestUpdateRouteBinding: data = repos.route_binding.update.call_args.args[1] assert data["agent_binding"] == "agent-new" assert data["updated_by"] == "admin-001" + # 乐观锁:加载的 ORM 实例 version 已递增 + assert orm.version == 2 + + @pytest.mark.asyncio + async def test_update_route_binding_raises_conflict_on_version_mismatch(self): + """覆盖点:expected_version 与持久化 version 不匹配 → ConflictError。""" + # Arrange - 持久化 version 已为 2(被其它并发更新推进),客户端仍带 1 + db = _make_db() + repos = _make_repos() + orm = _make_binding_orm() + orm.version = 2 + repos.route_binding.get_by_binding_id = AsyncMock(return_value=orm) + repos.route_binding.update = AsyncMock() + adapter = _build_adapter(db, repos) + cmd = UpdateRouteBindingCmd( + binding_id="bnd-001", + expected_version=1, + agent_binding="agent-new", + ) + + # Act / Assert + with pytest.raises(ConflictError): + await adapter.updateRouteBinding(cmd, _make_operator()) + # 冲突时不得写库 + repos.route_binding.update.assert_not_awaited() @pytest.mark.asyncio async def test_update_route_binding_raises_not_found_when_missing(self): @@ -249,7 +274,11 @@ class TestUpdateRouteBinding: repos = _make_repos() repos.route_binding.get_by_binding_id = AsyncMock(return_value=None) adapter = _build_adapter(db, repos) - cmd = UpdateRouteBindingCmd(binding_id="bnd-missing", agent_binding="agent-new") + cmd = UpdateRouteBindingCmd( + binding_id="bnd-missing", + expected_version=1, + agent_binding="agent-new", + ) # Act / Assert with pytest.raises(NotFoundError): @@ -393,47 +422,6 @@ class TestCountRouteBindings: await adapter.countRouteBindings(filter=RouteBindingFilter()) -# ─── listEnabledByAccount ───────────────────────────────────────────────── - - -@pytest.mark.unit -class TestListEnabledByAccount: - @pytest.mark.asyncio - async def test_list_enabled_by_account_returns_only_enabled_rules(self): - """覆盖点:仅返回启用规则,仓储 list_enabled_by_account 已过滤 enabled=True。""" - # Arrange - db = _make_db() - repos = _make_repos() - orm1 = _make_binding_orm(binding_id="bnd-1", enabled=True) - orm2 = _make_binding_orm(binding_id="bnd-2", enabled=True) - repos.route_binding.list_enabled_by_account = AsyncMock(return_value=[orm1, orm2]) - adapter = _build_adapter(db, repos) - - # Act - result = await adapter.listEnabledByAccount(ChannelType("feishu"), "acc-1") - - # Assert - assert isinstance(result, tuple) - assert len(result) == 2 - assert all(r.enabled for r in result) - repos.route_binding.list_enabled_by_account.assert_awaited_once_with(ChannelType("feishu"), "acc-1") - - @pytest.mark.asyncio - async def test_list_enabled_by_account_empty_returns_empty_tuple(self): - """覆盖点:账户无启用规则返回空元组。""" - # Arrange - db = _make_db() - repos = _make_repos() - repos.route_binding.list_enabled_by_account = AsyncMock(return_value=[]) - adapter = _build_adapter(db, repos) - - # Act - result = await adapter.listEnabledByAccount(ChannelType("feishu"), "acc-empty") - - # Assert - assert result == () - - # ─── getRouteBinding ────────────────────────────────────────────────────── diff --git a/backend/test/unit/channels/adapters/test_sqlalchemy_transaction_adapter.py b/backend/test/unit/channels/adapters/test_sqlalchemy_transaction_adapter.py index 7d138a02..30ccfd1a 100644 --- a/backend/test/unit/channels/adapters/test_sqlalchemy_transaction_adapter.py +++ b/backend/test/unit/channels/adapters/test_sqlalchemy_transaction_adapter.py @@ -19,9 +19,15 @@ pytestmark = pytest.mark.unit def _make_session() -> MagicMock: - """构造 AsyncSession 桩。""" + """构造 AsyncSession 桩。 + + 默认 ``in_transaction`` 返回 False(无隐式事务),模拟干净 session。 + 需要模拟 autobegin 场景时,在测试中覆盖 ``session.in_transaction`` 返回值。 + """ session = MagicMock() session.begin = AsyncMock() + session.in_transaction = MagicMock(return_value=False) + session.commit = AsyncMock() return session @@ -48,6 +54,32 @@ class TestSqlAlchemyTransactionContextAenter: assert result is ctx session.begin.assert_awaited_once() assert ctx._txn is not None + # 无隐式事务时不提交 + session.commit.assert_not_awaited() + + @pytest.mark.asyncio + async def test_aenter_commits_autobegin_transaction_before_begin(self): + """模拟 SQLAlchemy 2.0 autobegin 语义:前置 SELECT 触发隐式事务, + ``__aenter__`` 应先提交隐式事务再开启显式事务,避免 + ``InvalidRequestError: A transaction is already begun``。 + """ + # Arrange + session = _make_session() + session.in_transaction = MagicMock(return_value=True) + txn = _make_txn() + session.begin = AsyncMock(return_value=txn) + ctx = SqlAlchemyTransactionContext(session) + + # Act + result = await ctx.__aenter__() + + # Assert + assert result is ctx + # 先提交隐式事务 + session.commit.assert_awaited_once() + # 再开启显式事务 + session.begin.assert_awaited_once() + assert ctx._txn is txn @pytest.mark.unit 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 new file mode 100644 index 00000000..88edd6f5 --- /dev/null +++ b/backend/test/unit/channels/application/extension/scheduler_handlers/test_channel_content_review_retention_handler.py @@ -0,0 +1,154 @@ +"""ChannelContentReviewRetentionHandler 单元测试。 + +覆盖 ``yuxi.channels.application.extension.scheduler_handlers.channel_content_review_retention_handler``: +- execute 正常路径:无记录 / 有记录删除 +- payload 覆盖 retention_days 与 batch_size 默认值 +- session / repository 异常转换为 TaskResult(success=False) +""" + +from __future__ import annotations + +from contextlib import asynccontextmanager +from datetime import datetime, timedelta +from unittest.mock import AsyncMock, MagicMock + +import pytest +from yuxi.channels.application.extension.scheduler_handlers.channel_content_review_retention_handler import ( + ChannelContentReviewRetentionHandler, +) +from yuxi.scheduler.core.contracts import TaskContext +from yuxi.utils.datetime_utils import utc_now_naive + +pytestmark = pytest.mark.unit + + +# ─── 测试辅助 ───────────────────────────────────────────────────────────── + + +def _make_ctx(**payload_overrides) -> TaskContext: + """构造 TaskContext,payload 可按需覆盖。""" + return TaskContext( + task_id="task-content-review-retention", + run_id="run-001", + handler_name="channel_content_review_retention", + payload=payload_overrides, + triggered_by="auto", + scheduled_at=datetime(2024, 1, 1, 0, 0, 0), + ) + + +@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): + """构造 ChannelContentReviewRetentionHandler 及其依赖桩。""" + db = MagicMock() + if content_review_repo is None: + content_review_repo = AsyncMock() + content_review_repo.deleteOldReviewRecords.return_value = 0 + if logger is None: + logger = AsyncMock() + + handler = ChannelContentReviewRetentionHandler( + session_factory=lambda: _fake_session_factory(db), + content_review_repo_factory=lambda _db: content_review_repo, + logger=logger, + ) + return handler, { + "db": db, + "content_review_repo": content_review_repo, + "logger": logger, + } + + +# ─── 类属性 ─────────────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestHandlerAttributes: + def test_name_is_channel_content_review_retention(self): + assert ChannelContentReviewRetentionHandler.name == "channel_content_review_retention" + + def test_description_is_non_empty(self): + assert ChannelContentReviewRetentionHandler.description + assert isinstance(ChannelContentReviewRetentionHandler.description, str) + + +# ─── execute 正常路径 ───────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestHandlerExecute: + @pytest.mark.asyncio + async def test_successful_empty_deletion_returns_zero(self, fake_logger): + handler, deps = _make_handler(logger=fake_logger) + + result = await handler.execute(_make_ctx()) + + assert result.success is True + assert result.output["deleted_count"] == 0 + assert result.output["retention_days"] == 180 + deps["content_review_repo"].deleteOldReviewRecords.assert_awaited_once() + + @pytest.mark.asyncio + async def test_deletes_old_review_records(self, fake_logger): + handler, deps = _make_handler(logger=fake_logger) + deps["content_review_repo"].deleteOldReviewRecords.return_value = 42 + + result = await handler.execute(_make_ctx()) + + assert result.success is True + assert result.output["deleted_count"] == 42 + assert result.output["retention_days"] == 180 + deps["logger"].info.assert_awaited_once() + + @pytest.mark.asyncio + async def test_payload_overrides_defaults(self, fake_logger): + handler, deps = _make_handler(logger=fake_logger) + + await handler.execute(_make_ctx(retention_days=30, batch_size=100)) + + call_args = deps["content_review_repo"].deleteOldReviewRecords.call_args + assert call_args[1]["limit"] == 100 + cutoff = call_args[0][0] + # retention_days=30 → cutoff ≈ now - 30 days + expected_earliest = utc_now_naive() - timedelta(days=29) + expected_latest = utc_now_naive() - timedelta(days=31) + assert expected_latest < cutoff < expected_earliest + + +# ─── execute 异常路径 ───────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestHandlerErrors: + @pytest.mark.asyncio + async def test_session_factory_exception_returns_failure(self, fake_logger): + @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(), + 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) + deps["content_review_repo"].deleteOldReviewRecords.side_effect = RuntimeError("db error") + + result = await handler.execute(_make_ctx()) + + assert result.success is False + assert "db error" 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 5aad30d8..e14bcfd8 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 @@ -5,7 +5,7 @@ - payload 覆盖 batch_size 默认值 - 单条记录异常隔离 - session / repository 异常转换为 TaskResult(success=False) -- _reconstruct_pairing_approval 使用记录自身 version +- fromRecord 使用记录自身 version """ from __future__ import annotations @@ -17,10 +17,10 @@ from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channels.application.extension.scheduler_handlers.channel_pairing_expiration_handler import ( ChannelPairingExpirationHandler, - _reconstruct_pairing_approval, ) from yuxi.channels.contract.dtos.channel import ChannelType from yuxi.channels.contract.dtos.pairing import PairingRecord, PairingStatus +from yuxi.channels.core.model.pairing_approval import PairingApproval from yuxi.scheduler.core.contracts import TaskContext from yuxi.utils.datetime_utils import utc_now_naive @@ -137,6 +137,8 @@ class TestHandlerExecute: deps["pairing_repo"].updatePairingStatus.assert_awaited_once_with( "pr-001", PairingStatus.EXPIRED, + expected_version=3, + updated_by="system", ) @pytest.mark.asyncio @@ -250,7 +252,7 @@ class TestReconstructPairingApproval: def test_uses_record_version(self): record = _make_pairing_record(pairing_id="pr-001", version=7) - approval = _reconstruct_pairing_approval(record) + approval = PairingApproval.fromRecord(record) assert approval.pairing_id == "pr-001" assert approval.status == PairingStatus.PENDING @@ -260,6 +262,6 @@ class TestReconstructPairingApproval: created = utc_now_naive() - timedelta(hours=1) record = _make_pairing_record(pairing_id="pr-001", created_at=created) - approval = _reconstruct_pairing_approval(record) + approval = PairingApproval.fromRecord(record) assert approval.requested_at == created diff --git a/backend/test/unit/channels/application/lifecycle/test_route_match_registry_loader.py b/backend/test/unit/channels/application/lifecycle/test_route_match_registry_loader.py index 34647358..ce235d76 100644 --- a/backend/test/unit/channels/application/lifecycle/test_route_match_registry_loader.py +++ b/backend/test/unit/channels/application/lifecycle/test_route_match_registry_loader.py @@ -47,6 +47,7 @@ def _make_rule( updated_by="admin-001", created_at=datetime(2026, 1, 1), updated_at=datetime(2026, 1, 1), + version=1, ) diff --git a/backend/test/unit/channels/application/pipeline/control_plane/handlers/test_route_binding_handler.py b/backend/test/unit/channels/application/pipeline/control_plane/handlers/test_route_binding_handler.py index 4cfe54cd..d378868e 100644 --- a/backend/test/unit/channels/application/pipeline/control_plane/handlers/test_route_binding_handler.py +++ b/backend/test/unit/channels/application/pipeline/control_plane/handlers/test_route_binding_handler.py @@ -80,6 +80,7 @@ def _make_rule( updated_by="admin-001", created_at=datetime(2026, 1, 1), updated_at=datetime(2026, 1, 1), + version=1, ) @@ -400,6 +401,7 @@ class TestRouteBindingUpdate: operation="route_binding/update", params={ "binding_id": "bnd-001", + "version": 1, "agent_binding": "agent-new", }, target_channel=ChannelType("feishu"), diff --git a/backend/test/unit/channels/application/pipeline/inbound/test_inbound_pipeline.py b/backend/test/unit/channels/application/pipeline/inbound/test_inbound_pipeline.py index b2ddf152..484d5620 100644 --- a/backend/test/unit/channels/application/pipeline/inbound/test_inbound_pipeline.py +++ b/backend/test/unit/channels/application/pipeline/inbound/test_inbound_pipeline.py @@ -1,8 +1,8 @@ """yuxi.channels.application.pipeline.inbound.inbound_pipeline 单元测试。 覆盖 ``InboundPipeline(Pipeline)``: -- ``__init__()``:13 个阶段按序装配,``name="inbound"`` -- ``create()``:工厂方法接收全部依赖,构造 13 个阶段并返回实例 +- ``__init__()``:14 个阶段按序装配,``name="inbound"`` +- ``create()``:工厂方法接收全部依赖,构造 14 个阶段并返回实例 """ from __future__ import annotations @@ -20,6 +20,9 @@ from yuxi.channels.application.pipeline.inbound.command_check_stage import ( from yuxi.channels.application.pipeline.inbound.identity_resolve_stage import ( IdentityResolveStage, ) +from yuxi.channels.application.pipeline.inbound.inbound_idempotency_stage import ( + InboundIdempotencyStage, +) from yuxi.channels.application.pipeline.inbound.inbound_pipeline import InboundPipeline from yuxi.channels.application.pipeline.inbound.media_fetch_stage import MediaFetchStage from yuxi.channels.application.pipeline.inbound.receive_stage import ReceiveStage @@ -67,7 +70,7 @@ def _make_stage(stage_id: str) -> MagicMock: class TestInboundPipelineInit: """``InboundPipeline.__init__`` 阶段装配。""" - def test_init_stores_thirteen_stages_in_order(self): + def test_init_stores_fourteen_stages_in_order(self): # Arrange stages = { "receive": _make_stage("receive"), @@ -81,6 +84,7 @@ class TestInboundPipelineInit: "command-check": _make_stage("command-check"), "session-resolve": _make_stage("session-resolve"), "route": _make_stage("route"), + "inbound-idempotency": _make_stage("inbound-idempotency"), "agent-run-enqueue": _make_stage("agent-run-enqueue"), "reply": _make_stage("reply"), } @@ -98,13 +102,14 @@ class TestInboundPipelineInit: command_check_stage=stages["command-check"], session_resolve_stage=stages["session-resolve"], route_stage=stages["route"], + inbound_idempotency_stage=stages["inbound-idempotency"], agent_run_enqueue_stage=stages["agent-run-enqueue"], reply_stage=stages["reply"], ) # Assert assert pipeline.name == "inbound" - assert len(pipeline.stages) == 13 + assert len(pipeline.stages) == 14 expected_order = [ "receive", "signature-verify", @@ -117,6 +122,7 @@ class TestInboundPipelineInit: "command-check", "session-resolve", "route", + "inbound-idempotency", "agent-run-enqueue", "reply", ] @@ -139,6 +145,7 @@ class TestInboundPipelineInit: "command-check", "session-resolve", "route", + "inbound-idempotency", "agent-run-enqueue", "reply", ] @@ -157,6 +164,7 @@ class TestInboundPipelineInit: command_check_stage=stages["command-check"], session_resolve_stage=stages["session-resolve"], route_stage=stages["route"], + inbound_idempotency_stage=stages["inbound-idempotency"], agent_run_enqueue_stage=stages["agent-run-enqueue"], reply_stage=stages["reply"], logger=logger, @@ -181,6 +189,7 @@ class TestInboundPipelineInit: "command-check", "session-resolve", "route", + "inbound-idempotency", "agent-run-enqueue", "reply", ] @@ -199,6 +208,7 @@ class TestInboundPipelineInit: command_check_stage=stages["command-check"], session_resolve_stage=stages["session-resolve"], route_stage=stages["route"], + inbound_idempotency_stage=stages["inbound-idempotency"], agent_run_enqueue_stage=stages["agent-run-enqueue"], reply_stage=stages["reply"], ) @@ -214,7 +224,7 @@ class TestInboundPipelineInit: class TestInboundPipelineCreate: """``InboundPipeline.create`` 工厂方法。""" - def test_create_returns_inbound_pipeline_with_thirteen_stages(self): + def test_create_returns_inbound_pipeline_with_fourteen_stages(self): # Arrange — 所有依赖使用 MagicMock 桩 kwargs = dict( status_adapter_registry={}, @@ -248,7 +258,7 @@ class TestInboundPipelineCreate: # Assert assert isinstance(pipeline, InboundPipeline) assert pipeline.name == "inbound" - assert len(pipeline.stages) == 13 + assert len(pipeline.stages) == 14 stage_ids = [s.id for s in pipeline.stages] assert stage_ids[0] == "receive" assert stage_ids[-1] == "reply" @@ -296,8 +306,9 @@ class TestInboundPipelineCreate: assert isinstance(pipeline.stages[8], CommandCheckStage) assert isinstance(pipeline.stages[9], SessionResolveStage) assert isinstance(pipeline.stages[10], RouteStage) - assert isinstance(pipeline.stages[11], AgentRunEnqueueStage) - assert isinstance(pipeline.stages[12], ReplyStage) + assert isinstance(pipeline.stages[11], InboundIdempotencyStage) + assert isinstance(pipeline.stages[12], AgentRunEnqueueStage) + assert isinstance(pipeline.stages[13], ReplyStage) def test_create_accepts_optional_registries_and_ports(self): # Arrange — 可选参数传 None / 提供值 @@ -355,6 +366,7 @@ _INBOUND_STAGE_IDS = [ "command-check", "session-resolve", "route", + "inbound-idempotency", "agent-run-enqueue", "reply", ] @@ -381,7 +393,7 @@ class TestInboundPipelineRun: @pytest.mark.asyncio async def test_run_executes_all_stages_in_order_when_all_succeed(self): - # Arrange - 13 个阶段全部成功 + # Arrange - 14 个阶段全部成功 stages = [_make_stage(sid) for sid in _INBOUND_STAGE_IDS] pipeline = _make_pipeline_with_stages(stages) ctx = MagicMock() @@ -496,5 +508,5 @@ class TestInboundPipelineRun: ctx.trace_id = "trace-008" # Act await pipeline.run(ctx) - # Assert - info 日志被调用 13 次(每个成功阶段一次) - assert logger.info.await_count == 13 + # Assert - info 日志被调用 14 次(每个成功阶段一次) + assert logger.info.await_count == 14 diff --git a/backend/test/unit/channels/application/pipeline/inbound/test_route_stage.py b/backend/test/unit/channels/application/pipeline/inbound/test_route_stage.py index 50551db9..5f9eac95 100644 --- a/backend/test/unit/channels/application/pipeline/inbound/test_route_stage.py +++ b/backend/test/unit/channels/application/pipeline/inbound/test_route_stage.py @@ -7,7 +7,6 @@ - ``process(context)``:所有者保护未启用 → 更新绑定并写入 route_binding - ``process(context)``:会话为 None 时视为所有者 - ``process(context)``:会话无所有者时记录告警日志 -- ``process(context)``:从持久化端口刷新会话 """ from __future__ import annotations @@ -99,11 +98,11 @@ def _make_session( def _make_stage( *, route_resolver: AsyncMock | None = None, - persistence_port: AsyncMock | None = None, + session_repository: AsyncMock | None = None, ) -> RouteStage: return RouteStage( route_resolver=route_resolver or AsyncMock(), - persistence_port=persistence_port or AsyncMock(), + session_repository=session_repository or AsyncMock(), ) @@ -145,9 +144,7 @@ class TestRouteStageProcess: # Arrange:resolve 返回 None → 临时会话排除,current_uid 取 service_account.uid resolver = AsyncMock() resolver.resolve.return_value = None - persistence_port = AsyncMock() - persistence_port.getChannelSession.return_value = None - stage = _make_stage(route_resolver=resolver, persistence_port=persistence_port) + stage = _make_stage(route_resolver=resolver) session = _make_session(unified_identity_id="uid-001") ctx = _make_ctx(channel_session=session) @@ -188,9 +185,7 @@ class TestRouteStageProcess: resolver = AsyncMock() resolver.resolve.return_value = core_binding session = _make_session(owner_peer_id="owner-001", is_owner=False) - persistence = AsyncMock() - persistence.getChannelSession.return_value = session - stage = _make_stage(route_resolver=resolver, persistence_port=persistence) + stage = _make_stage(route_resolver=resolver) ctx = _make_ctx(channel_session=session, peer_id="user-002") # Act @@ -208,9 +203,7 @@ class TestRouteStageProcess: resolver = AsyncMock() resolver.resolve.return_value = core_binding session = _make_session(owner_peer_id="user-001", is_owner=True) - persistence = AsyncMock() - persistence.getChannelSession.return_value = session - stage = _make_stage(route_resolver=resolver, persistence_port=persistence) + stage = _make_stage(route_resolver=resolver) ctx = _make_ctx(channel_session=session, peer_id="user-001") # Act @@ -229,9 +222,7 @@ class TestRouteStageProcess: resolver = AsyncMock() resolver.resolve.return_value = core_binding session = _make_session(is_owner=False) - persistence = AsyncMock() - persistence.getChannelSession.return_value = session - stage = _make_stage(route_resolver=resolver, persistence_port=persistence) + stage = _make_stage(route_resolver=resolver) ctx = _make_ctx(channel_session=session) # Act @@ -259,45 +250,6 @@ class TestRouteStageProcess: assert ctx.route_binding is not None core_binding.updateBinding.assert_called_once() - @pytest.mark.asyncio - async def test_session_refreshed_from_persistence(self): - # Arrange:context.channel_session 存在时从持久化刷新 - core_binding = _make_core_binding(is_primary_owner_protected=False) - resolver = AsyncMock() - resolver.resolve.return_value = core_binding - original_session = _make_session(session_id="sess-old") - refreshed_session = _make_session(session_id="sess-new", is_owner=True) - persistence = AsyncMock() - persistence.getChannelSession.return_value = refreshed_session - stage = _make_stage(route_resolver=resolver, persistence_port=persistence) - ctx = _make_ctx(channel_session=original_session) - - # Act - await stage.process(ctx) - - # Assert:resolve 使用刷新后的会话 - persistence.getChannelSession.assert_awaited_once_with("sess-old") - assert resolver.resolve.call_args.args[1] is refreshed_session - - @pytest.mark.asyncio - async def test_session_not_refreshed_when_persistence_returns_none(self): - # Arrange:持久化返回 None → 使用原会话 - core_binding = _make_core_binding(is_primary_owner_protected=False) - resolver = AsyncMock() - resolver.resolve.return_value = core_binding - original_session = _make_session(session_id="sess-old", is_owner=True) - persistence = AsyncMock() - persistence.getChannelSession.return_value = None - stage = _make_stage(route_resolver=resolver, persistence_port=persistence) - ctx = _make_ctx(channel_session=original_session) - - # Act - ok = await stage.process(ctx) - - # Assert:使用原会话,绑定成功写入 - assert ok is True - assert ctx.route_binding is not None - # ─── 阶段契约属性 ─────────────────────────────────────────────────────────── diff --git a/backend/test/unit/channels/application/pipeline/inbound/test_security_stage.py b/backend/test/unit/channels/application/pipeline/inbound/test_security_stage.py index e7582202..846fa7d9 100644 --- a/backend/test/unit/channels/application/pipeline/inbound/test_security_stage.py +++ b/backend/test/unit/channels/application/pipeline/inbound/test_security_stage.py @@ -268,6 +268,39 @@ class TestSecurityStageProcess: # Assert assert ctx.is_admin is True + @pytest.mark.asyncio + async def test_process_dm_policy_none_defaults_to_allow(self): + # Arrange — 未配置 dm_policy 时回退到 schema 默认值 allow + config_port = _make_config_port(dm_policy_value=None) + dm_guard = MagicMock() + dm_guard.decide = AsyncMock(return_value=DmDecision.ALLOW) + audit_repo = AsyncMock() + audit_repo.saveAuditLog = AsyncMock() + bot_guard = MagicMock() + bot_guard.checkInbound = AsyncMock() + mention_eval = MagicMock() + mention_eval.isUserInitiated.return_value = True + admin_eval = MagicMock() + admin_eval.evaluate.return_value = False + stage = _make_stage( + config_port=config_port, + audit_log_repository=audit_repo, + dm_security_guard=dm_guard, + bot_loop_budget_guard=bot_guard, + mention_evaluator=mention_eval, + admin_evaluator=admin_eval, + ) + ctx = _make_ctx(chat_type="p2p") + + # Act + result = await stage.process(ctx) + + # Assert + assert result is True + assert ctx.dm_decision == DmDecision.ALLOW + call_kwargs = dm_guard.decide.await_args.args + assert call_kwargs[3] == DmPolicy.ALLOW + @pytest.mark.asyncio async def test_process_dm_policy_as_string_is_converted(self): # Arrange — dm_policy 配置为字符串 @@ -385,7 +418,6 @@ class TestSecurityStageWriteDmDecisionAuditLog: audit_repo.saveAuditLog.assert_awaited_once() call_args = audit_repo.saveAuditLog.await_args cmd = call_args.args[0] - tx = call_args.kwargs.get("tx") or call_args.args[1] if len(call_args.args) > 1 else None assert isinstance(cmd, SaveAuditLogCmd) assert cmd.operator == "system" assert cmd.result == "success" diff --git a/backend/test/unit/channels/application/pipeline/outbound/test_trusted_inject_stage.py b/backend/test/unit/channels/application/pipeline/outbound/test_trusted_inject_stage.py index 62db8ffe..b6d1d3be 100644 --- a/backend/test/unit/channels/application/pipeline/outbound/test_trusted_inject_stage.py +++ b/backend/test/unit/channels/application/pipeline/outbound/test_trusted_inject_stage.py @@ -1,7 +1,8 @@ """yuxi.channels.application.pipeline.outbound.trusted_inject_stage 单元测试。 覆盖 ``TrustedInjectStage``: -- ``process(context)``:正常路径、formatted_message 缺失、validate 异常隔离 +- ``process(context)``:正常路径、formatted_message 缺失抛 InternalError、 + validate 异常隔离 """ from __future__ import annotations @@ -16,7 +17,7 @@ from yuxi.channels.application.pipeline.outbound.trusted_inject_stage import ( from yuxi.channels.contract.dtos.channel import ChannelType from yuxi.channels.contract.dtos.outbound import FormattedMessage, RichMessage from yuxi.channels.contract.dtos.trusted import TrustedMessageContext -from yuxi.channels.contract.errors import PermissionDeniedError +from yuxi.channels.contract.errors import InternalError, PermissionDeniedError pytestmark = pytest.mark.unit @@ -32,6 +33,7 @@ def _make_ctx(**overrides) -> OutboundContext: delivery_mode="persistent", trusted_sender_id="user-001", conversation_id="conv-001", + formatted_message=FormattedMessage(content="rendered"), ) defaults.update(overrides) return OutboundContext(**defaults) @@ -76,8 +78,9 @@ class TestTrustedInjectStageProcess: assert ctx.trusted_message.rich_message is rich @pytest.mark.asyncio - async def test_formatted_message_none_uses_empty_content(self): - # Arrange + async def test_formatted_message_none_raises_internal_error(self): + # Arrange:formatted_message 缺失(上游 format 降级或异常), + # fail-closed 抛 InternalError 暴露前置条件违反,而非掩盖为空消息投递 guard = _make_guard() ctx = _make_ctx(formatted_message=None) stage = TrustedInjectStage( @@ -85,14 +88,10 @@ class TestTrustedInjectStageProcess: conversation_port=AsyncMock(), ) - # Act - ok = await stage.process(ctx) - - # Assert - assert ok is True - assert ctx.trusted_message is not None - assert ctx.trusted_message.content == "" - assert ctx.trusted_message.rich_message is None + # Act / Assert + with pytest.raises(InternalError) as exc_info: + await stage.process(ctx) + assert "formatted_message is None" in str(exc_info.value) @pytest.mark.asyncio async def test_validate_raises_propagates(self): diff --git a/backend/test/unit/channels/application/transport/test_base_worker.py b/backend/test/unit/channels/application/transport/test_base_worker.py index 994b683f..7216e52b 100644 --- a/backend/test/unit/channels/application/transport/test_base_worker.py +++ b/backend/test/unit/channels/application/transport/test_base_worker.py @@ -53,7 +53,7 @@ def _make_worker( """构造测试用 worker,所有依赖注入桩。""" return _TestableWorker( message_deliverer=message_deliverer or AsyncMock(), - logger=logger or MagicMock(), + logger=logger or AsyncMock(), config_port=config_port or AsyncMock(), event_publisher=event_publisher or AsyncMock(), circuit_breaker=circuit_breaker or AsyncMock(), diff --git a/backend/test/unit/channels/application/transport/test_base_worker_auth_expired.py b/backend/test/unit/channels/application/transport/test_base_worker_auth_expired.py index 31e410f3..66529b8f 100644 --- a/backend/test/unit/channels/application/transport/test_base_worker_auth_expired.py +++ b/backend/test/unit/channels/application/transport/test_base_worker_auth_expired.py @@ -53,7 +53,7 @@ def _make_worker( circuit_breaker.isOpen.return_value = False return _AuthExpiredWorker( message_deliverer=AsyncMock(), - logger=logger or MagicMock(), + logger=logger or AsyncMock(), config_port=config_port or AsyncMock(), event_publisher=event_publisher or AsyncMock(), circuit_breaker=circuit_breaker, diff --git a/backend/test/unit/channels/application/transport/test_manager.py b/backend/test/unit/channels/application/transport/test_manager.py index ddbf7d27..e9c5c649 100644 --- a/backend/test/unit/channels/application/transport/test_manager.py +++ b/backend/test/unit/channels/application/transport/test_manager.py @@ -44,7 +44,7 @@ def _make_manager( config_port=config_port or AsyncMock(), event_bus=event_bus or MagicMock(), circuit_breaker=cb, - logger=logger or MagicMock(), + logger=logger or AsyncMock(), message_deliverer=message_deliverer or AsyncMock(), transport_config=transport_config, ) @@ -499,3 +499,165 @@ class TestTransportManagerEventHandlers: # Assert manager.reloadConfig.assert_not_awaited() await manager.stop(timeout=0.1) + + +@pytest.mark.unit +class TestTransportManagerPluginStatus: + """plugin_status 持久化测试(FR-32,TransportManager 为写入方)。 + + 覆盖传输任务 start/stop/error 后调用 ``updatePluginStatus`` 持久化 + 按账户的插件运行态(running / stopped / error)。 + """ + + @pytest.mark.asyncio + async def test_on_account_online_marks_plugin_status_running( + self, fake_logger, fake_config + ): + # Arrange 上线启动 Puller 任务后应写 plugin_status="running" + puller_adapter = MagicMock() + plugin_da = _make_plugin_da(puller_adapter=puller_adapter) + plugin_registry = MagicMock() + plugin_registry.listPluginAdapters.return_value = [ + (ChannelType("feishu"), plugin_da) + ] + persistence_port = AsyncMock() + manager = _make_manager( + plugin_registry=plugin_registry, + persistence_port=persistence_port, + logger=fake_logger, + config_port=fake_config, + ) + await manager.start() + event = MagicMock() + event.payload = {"channel_type": "feishu", "account_id": "acc1"} + + # Act + await manager._on_account_online(event) + + # Assert + persistence_port.updatePluginStatus.assert_awaited_once_with( + ChannelType("feishu"), "acc1", "running" + ) + await manager.stop(timeout=0.1) + + @pytest.mark.asyncio + async def test_on_account_offline_marks_plugin_status_stopped( + self, fake_logger, fake_config + ): + # Arrange 下线停止任务后应写 plugin_status="stopped" + puller_adapter = MagicMock() + plugin_da = _make_plugin_da(puller_adapter=puller_adapter) + plugin_registry = MagicMock() + plugin_registry.listPluginAdapters.return_value = [ + (ChannelType("feishu"), plugin_da) + ] + persistence_port = AsyncMock() + manager = _make_manager( + plugin_registry=plugin_registry, + persistence_port=persistence_port, + logger=fake_logger, + config_port=fake_config, + ) + await manager.start() + online_event = MagicMock() + online_event.payload = {"channel_type": "feishu", "account_id": "acc1"} + await manager._on_account_online(online_event) + + offline_event = MagicMock() + offline_event.payload = { + "channel_type": "feishu", + "account_id": "acc1", + "reason": "test", + } + + # Act + await manager._on_account_offline(offline_event) + + # Assert 上线写 running、下线写 stopped(共 2 次,最近一次为 stopped) + assert persistence_port.updatePluginStatus.await_count == 2 + persistence_port.updatePluginStatus.assert_awaited_with( + ChannelType("feishu"), "acc1", "stopped" + ) + await manager.stop(timeout=0.1) + + @pytest.mark.asyncio + async def test_on_transport_error_marks_plugin_status_error( + self, fake_logger, fake_config + ): + # Arrange 传输错误事件应写 plugin_status="error" + persistence_port = AsyncMock() + manager = _make_manager( + persistence_port=persistence_port, + logger=fake_logger, + config_port=fake_config, + ) + await manager.start() + event = MagicMock() + event.trace_id = "trace-1" + event.payload = { + "channel_type": "feishu", + "account_id": "acc1", + "error_category": "network", + "error_code": "CONN_TIMEOUT", + } + + # Act + await manager._on_transport_error(event) + + # Assert + persistence_port.updatePluginStatus.assert_awaited_once_with( + ChannelType("feishu"), "acc1", "error" + ) + await manager.stop(timeout=0.1) + + @pytest.mark.asyncio + async def test_on_transport_error_skips_when_account_missing( + self, fake_logger, fake_config + ): + # Arrange 事件缺少 account_id 时不写 plugin_status + persistence_port = AsyncMock() + manager = _make_manager( + persistence_port=persistence_port, + logger=fake_logger, + config_port=fake_config, + ) + await manager.start() + event = MagicMock() + event.trace_id = "trace-1" + event.payload = {"channel_type": "feishu"} + + # Act + await manager._on_transport_error(event) + + # Assert + persistence_port.updatePluginStatus.assert_not_awaited() + await manager.stop(timeout=0.1) + + @pytest.mark.asyncio + async def test_touch_plugin_status_swallows_persistence_errors( + self, fake_logger, fake_config + ): + # Arrange updatePluginStatus 抛异常时仅告警,不穿透到传输主流程 + persistence_port = AsyncMock() + persistence_port.updatePluginStatus.side_effect = RuntimeError("db down") + manager = _make_manager( + persistence_port=persistence_port, + logger=fake_logger, + config_port=fake_config, + ) + await manager.start() + event = MagicMock() + event.trace_id = "trace-1" + event.payload = { + "channel_type": "feishu", + "account_id": "acc1", + "error_category": "network", + "error_code": "CONN_TIMEOUT", + } + + # Act 不应抛异常 + await manager._on_transport_error(event) + + # Assert + persistence_port.updatePluginStatus.assert_awaited_once() + await manager.stop(timeout=0.1) diff --git a/backend/test/unit/channels/application/transport/test_puller_worker.py b/backend/test/unit/channels/application/transport/test_puller_worker.py index af5e9759..ec8e29ab 100644 --- a/backend/test/unit/channels/application/transport/test_puller_worker.py +++ b/backend/test/unit/channels/application/transport/test_puller_worker.py @@ -42,7 +42,7 @@ def _make_puller( """构造测试用 PullerWorker。""" return PullerWorker( message_deliverer=message_deliverer or AsyncMock(), - logger=logger or MagicMock(), + logger=logger or AsyncMock(), config_port=config_port or AsyncMock(), event_publisher=event_publisher or AsyncMock(), circuit_breaker=circuit_breaker or AsyncMock(), @@ -55,6 +55,7 @@ def _account_with_cursor(cursor: str = "") -> MagicMock: """构造带 transport_cursor 的 channel account 桩。""" account = MagicMock() account.transport_cursor = cursor + account.version = 1 return account diff --git a/backend/test/unit/channels/application/transport/test_stream_worker.py b/backend/test/unit/channels/application/transport/test_stream_worker.py index 37847b99..08e8553f 100644 --- a/backend/test/unit/channels/application/transport/test_stream_worker.py +++ b/backend/test/unit/channels/application/transport/test_stream_worker.py @@ -39,7 +39,7 @@ def _make_stream_worker( """构造测试用 StreamWorker,所有依赖注入桩。""" return StreamWorker( message_deliverer=message_deliverer or AsyncMock(), - logger=logger or MagicMock(), + logger=logger or AsyncMock(), config_port=config_port or AsyncMock(), event_publisher=event_publisher or AsyncMock(), circuit_breaker=circuit_breaker or AsyncMock(), diff --git a/backend/test/unit/channels/application/transport/test_stream_worker_backoff.py b/backend/test/unit/channels/application/transport/test_stream_worker_backoff.py index a043d81b..612c939d 100644 --- a/backend/test/unit/channels/application/transport/test_stream_worker_backoff.py +++ b/backend/test/unit/channels/application/transport/test_stream_worker_backoff.py @@ -37,7 +37,7 @@ def _make_stream_worker( """构造测试用 StreamWorker,所有依赖注入桩。""" return StreamWorker( message_deliverer=message_deliverer or AsyncMock(), - logger=logger or MagicMock(), + logger=logger or AsyncMock(), config_port=config_port or AsyncMock(), event_publisher=event_publisher or AsyncMock(), circuit_breaker=circuit_breaker or AsyncMock(), diff --git a/backend/test/unit/channels/application/usecase/test_content_review_query_service.py b/backend/test/unit/channels/application/usecase/test_content_review_query_service.py index 75436a56..9aa821ac 100644 --- a/backend/test/unit/channels/application/usecase/test_content_review_query_service.py +++ b/backend/test/unit/channels/application/usecase/test_content_review_query_service.py @@ -47,6 +47,8 @@ def _make_history_cmd(**overrides) -> ContentReviewHistoryQueryCmd: channel_type=ChannelType("feishu"), account_id="acc1", verdict=None, + resource_type=None, + trace_id=None, start_time=None, end_time=None, limit=20, diff --git a/backend/test/unit/channels/contract/dtos/test_dtos_supplement.py b/backend/test/unit/channels/contract/dtos/test_dtos_supplement.py index d136a680..55a27753 100644 --- a/backend/test/unit/channels/contract/dtos/test_dtos_supplement.py +++ b/backend/test/unit/channels/contract/dtos/test_dtos_supplement.py @@ -1319,6 +1319,7 @@ class TestPersistenceCmdDTOs: cmd = UpdateChannelAccountCmd( channel_type=ChannelType("feishu"), account_id="acc-1", + expected_version=1, display_name="新名称", ) # Assert @@ -1329,14 +1330,19 @@ class TestPersistenceCmdDTOs: # Arrange # Act / Assert with pytest.raises(ValidationError) as exc_info: - UpdateChannelAccountCmd(channel_type=ChannelType("feishu"), account_id="", display_name="x") + UpdateChannelAccountCmd( + channel_type=ChannelType("feishu"), + account_id="", + expected_version=1, + display_name="x", + ) assert exc_info.value.field == "account_id" def test_update_channel_account_cmd_no_updatable_field_raises(self): # Arrange # Act / Assert - 至少一个可更新字段 with pytest.raises(ValidationError) as exc_info: - UpdateChannelAccountCmd(channel_type=ChannelType("feishu"), account_id="acc-1") + UpdateChannelAccountCmd(channel_type=ChannelType("feishu"), account_id="acc-1", expected_version=1) assert exc_info.value.field == "update" def test_save_channel_session_cmd_construction(self): diff --git a/backend/test/unit/channels/core/model/test_channel_account.py b/backend/test/unit/channels/core/model/test_channel_account.py index ae31f283..fdeecd51 100644 --- a/backend/test/unit/channels/core/model/test_channel_account.py +++ b/backend/test/unit/channels/core/model/test_channel_account.py @@ -111,9 +111,10 @@ class TestChannelAccountFromDto: transport_cursor="cursor-1", last_rotated_at=None, onboarding_status=OnboardingStatus.ONLINE, + version=5, ) # Act - account = ChannelAccount.fromDto(dto, version=5) + account = ChannelAccount.fromDto(dto) # Assert assert account.status == AccountStatus.DEGRADED assert account.version == 5 @@ -153,7 +154,7 @@ class TestChannelAccountUpdateConfig: account.updateConfig({"app_id": "new"}) # Assert assert account.config == {"app_id": "new"} - assert account.version == original_version + 1 + assert account.version == original_version assert account.updated_at is not None def test_rotate_credentials_requires_online_onboarding(self): @@ -198,7 +199,7 @@ class TestChannelAccountDisableEnable: # Assert assert account.status == AccountStatus.DISABLED assert account.enabled is False - assert account.version == original_version + 1 + assert account.version == original_version def test_disable_is_idempotent(self): # Arrange @@ -220,7 +221,7 @@ class TestChannelAccountDisableEnable: # Assert assert account.status == AccountStatus.ACTIVE assert account.enabled is True - assert account.version == original_version + 1 + assert account.version == original_version def test_enable_is_idempotent(self): # Arrange @@ -253,7 +254,7 @@ class TestChannelAccountDegradeRecover: account.degrade() # Assert assert account.status == AccountStatus.DEGRADED - assert account.version == original_version + 1 + assert account.version == original_version def test_degrade_is_idempotent(self): # Arrange @@ -283,7 +284,7 @@ class TestChannelAccountDegradeRecover: # Assert assert account.status == AccountStatus.ACTIVE assert account.enabled is True - assert account.version == original_version + 1 + assert account.version == original_version def test_recover_is_idempotent(self): # Arrange @@ -470,7 +471,7 @@ class TestChannelAccountMarkConfigured: # Act account.markConfigured("ref-001") # Assert - assert account.version == original_version + 1 + assert account.version == original_version @pytest.mark.parametrize( "onboarding", @@ -521,7 +522,7 @@ class TestChannelAccountMarkVerified: # Act account.markVerified() # Assert - assert account.version == original_version + 1 + assert account.version == original_version @pytest.mark.parametrize( "onboarding", @@ -565,7 +566,7 @@ class TestChannelAccountMarkOnline: # Act account.markOnline() # Assert - assert account.version == original_version + 1 + assert account.version == original_version @pytest.mark.parametrize( "onboarding", @@ -632,7 +633,7 @@ class TestChannelAccountMarkOffline: # Act account.markOffline() # Assert - assert account.version == original_version + 1 + assert account.version == original_version @pytest.mark.parametrize( "onboarding", @@ -689,7 +690,7 @@ class TestChannelAccountMarkFailed: # Act account.markFailed("boom") # Assert - assert account.version == original_version + 1 + assert account.version == original_version def test_mark_failed_from_online_raises_conflict_error(self): # Arrange: ONLINE 状态应使用 degrade 处理运行态故障 @@ -749,7 +750,7 @@ class TestChannelAccountReenterOnboarding: # Act account.reenterOnboarding() # Assert - assert account.version == original_version + 1 + assert account.version == original_version @pytest.mark.parametrize( "onboarding", @@ -792,7 +793,7 @@ class TestChannelAccountRecoverFromFailure: # Act account.recoverFromFailure() # Assert - assert account.version == original_version + 1 + assert account.version == original_version @pytest.mark.parametrize( "onboarding", @@ -840,7 +841,7 @@ class TestChannelAccountDisableEnableOnboardingSync: assert account.status == AccountStatus.DISABLED assert account.enabled is False # markOffline 递增一次 version - assert account.version == original_version + 1 + assert account.version == original_version def test_enable_when_onboarding_offline_transitions_status_and_onboarding(self): """覆盖点:enable() 当 onboarding_status==OFFLINE 时将 status→ACTIVE 且 onboarding_status→ONLINE。""" @@ -856,7 +857,7 @@ class TestChannelAccountDisableEnableOnboardingSync: assert account.enabled is True assert account.onboarding_status == OnboardingStatus.ONLINE # status 变更与 onboarding 变更各递增一次 version - assert account.version == original_version + 2 + assert account.version == original_version def test_enable_when_onboarding_online_is_idempotent(self): """覆盖点:enable() 当 onboarding_status==ONLINE 时直接返回(幂等)。""" @@ -893,9 +894,11 @@ class TestChannelAccountFromDtoFieldFallback: updated_at=datetime.now(UTC), transport_cursor="", last_rotated_at=None, + version=3, + last_error=None, ) # Act - account = ChannelAccount.fromDto(legacy_dto, version=3) + account = ChannelAccount.fromDto(legacy_dto) # Assert: 缺少 onboarding_status → PENDING,缺少 credential_ref → None,缺少 credential_version → 0 assert account.onboarding_status == OnboardingStatus.PENDING assert account.credential_ref is None diff --git a/backend/test/unit/channels/core/model/test_channel_session.py b/backend/test/unit/channels/core/model/test_channel_session.py index ab5ded28..a85f285d 100644 --- a/backend/test/unit/channels/core/model/test_channel_session.py +++ b/backend/test/unit/channels/core/model/test_channel_session.py @@ -4,7 +4,7 @@ - create() 静态工厂:session_id 生成、peer_id 空值校验 - bindIdentity():首次绑定 / 冲突 / 已关闭会话 - canBeMerged() / markMerged():主所有者保护与软删除 -- markTemporary() / updateActiveAgent():状态修改与前置校验 +- markTemporary():状态修改与前置校验 - isPrimaryOwner() / isOwner() / isDeleted() / isClosed():查询方法 - close():关闭与已关闭/已软删除校验 - isCronSession() / isTemporarySession():格式校验 @@ -75,8 +75,7 @@ class TestChannelSessionCreate: assert session.version == 1 assert session.deleted_at is None assert session.closed_at is None - assert session.active_agent_id is None - assert session.last_message_at is None + assert session.last_message_at is not None def test_create_empty_peer_id_raises(self): # Arrange & Act & Assert: SaveChannelSessionCmd.__post_init__ 校验 @@ -175,7 +174,7 @@ class TestChannelSessionMerge: session.markMerged() -# ─── markTemporary / updateActiveAgent ──────────────────────────────────── +# ─── markTemporary ────────────────────────────────────────────────────── @pytest.mark.unit @@ -198,31 +197,6 @@ class TestChannelSessionStateUpdate: with pytest.raises(RuleViolationError): session.markTemporary() - def test_update_active_agent_sets_field(self): - # Arrange - session = _make_session() - original_version = session.version - # Act - session.updateActiveAgent("agent-1") - # Assert - assert session.active_agent_id == "agent-1" - assert session.version == original_version + 1 - - def test_update_active_agent_empty_raises(self): - # Arrange - session = _make_session() - # Act & Assert - with pytest.raises(ValidationError): - session.updateActiveAgent("") - - def test_update_active_agent_closed_session_raises(self): - # Arrange - session = _make_session() - session.close() - # Act & Assert - with pytest.raises(RuleViolationError): - session.updateActiveAgent("agent-1") - # ─── isPrimaryOwner / isOwner / isDeleted / isClosed 查询 ────────────────── diff --git a/backend/test/unit/channels/core/model/test_pairing_approval.py b/backend/test/unit/channels/core/model/test_pairing_approval.py index 365b3482..4623f0d0 100644 --- a/backend/test/unit/channels/core/model/test_pairing_approval.py +++ b/backend/test/unit/channels/core/model/test_pairing_approval.py @@ -44,6 +44,8 @@ def _make_approval( expires_at: Any = None, approver_id: str | None = None, approved_at: Any = None, + rejected_at: Any = None, + revoked_at: Any = None, reason: str | None = None, version: int = 1, ) -> PairingApproval: @@ -61,6 +63,8 @@ def _make_approval( status=status, approver_id=approver_id, approved_at=approved_at, + rejected_at=rejected_at, + revoked_at=revoked_at, expires_at=expires_at, requested_at=requested_at, reason=reason, @@ -176,7 +180,7 @@ class TestPairingApprovalReject: # Assert assert approval.status == PairingStatus.REJECTED assert approval.approver_id == "admin-1" - assert approval.approved_at is not None + assert approval.rejected_at is not None assert approval.reason == "rejected reason" assert approval.version == original_version + 1 @@ -220,6 +224,7 @@ class TestPairingApprovalRevoke: approval.revoke(_operator("admin-1"), "manual revoke") # Assert assert approval.status == PairingStatus.REVOKED + assert approval.revoked_at is not None assert approval.reason == "manual revoke" assert approval.version == original_version + 1 diff --git a/backend/test/unit/channels/core/registry/test_route_match_registry.py b/backend/test/unit/channels/core/registry/test_route_match_registry.py index 6ae1fddf..b132c6a6 100644 --- a/backend/test/unit/channels/core/registry/test_route_match_registry.py +++ b/backend/test/unit/channels/core/registry/test_route_match_registry.py @@ -653,6 +653,7 @@ def _make_rule( updated_by="admin-001", created_at=datetime(2026, 1, 1), updated_at=datetime(2026, 1, 1), + version=1, ) diff --git a/backend/test/unit/channels/core/service/test_bot_loop_budget_guard.py b/backend/test/unit/channels/core/service/test_bot_loop_budget_guard.py index 311afbbd..413d1579 100644 --- a/backend/test/unit/channels/core/service/test_bot_loop_budget_guard.py +++ b/backend/test/unit/channels/core/service/test_bot_loop_budget_guard.py @@ -1,315 +1,147 @@ -"""Bot 循环预算守卫领域服务单元测试。 - -覆盖 ``yuxi.channels.core.service.bot_loop_budget_guard.BotLoopBudgetGuard``: -- 构造:依赖存储与 logger 可选 -- consumeOutbound():预算消耗、耗尽抛错、cache-miss 降级、配置作用域 -- checkInbound():用户主动发起重置预算、回复 Bot 仅检查耗尽、cache-miss 降级 -- 缓存故障降级:DependencyError 捕获、非 DependencyError 不捕获 -""" +"""BotLoopBudgetGuard 单元测试。""" from __future__ import annotations -from unittest.mock import MagicMock - import pytest -from yuxi.channels.contract.dtos.config import ConfigScope -from yuxi.channels.contract.dtos.option import Some -from yuxi.channels.contract.errors import BotLoopBudgetExceededError, DependencyError +from unittest.mock import AsyncMock, MagicMock + +from yuxi.channels.contract.dtos.config import ConfigScope, ConfigValue +from yuxi.channels.contract.dtos.option import Nothing, Some +from yuxi.channels.contract.errors import BotLoopBudgetExceededError, ValidationError from yuxi.channels.core.service.bot_loop_budget_guard import BotLoopBudgetGuard -pytestmark = pytest.mark.unit + +def _make_config_port(value: object | None = None) -> MagicMock: + port = MagicMock() + port.get = AsyncMock(return_value=ConfigValue( + key="bot_loop_budget", + value=value, + version=1, + scope=ConfigScope.ACCOUNT, + )) + return port -def _make_budget_config( - *, - max_replies: int | None = None, - cooldown: int | None = None, -) -> MagicMock: - """构造配置值对象(含 .value 属性)。 - - 模拟 config_port.get 返回的 ConfigEntry,``.value`` 为预算配置 dict。 - 缺省字段测试由调用方控制:传 None 表示该字段缺失。 - """ - cfg = MagicMock() - value: dict[str, int] = {} - if max_replies is not None: - value["max_replies_per_hour"] = max_replies - if cooldown is not None: - value["cooldown_seconds"] = cooldown - cfg.value = value - return cfg +def _make_guard( + cache_port: MagicMock | None = None, + config_port: MagicMock | None = None, +) -> BotLoopBudgetGuard: + return BotLoopBudgetGuard( + cache_port=cache_port or MagicMock(), + config_port=config_port or _make_config_port(), + logger=AsyncMock(), + ) -# ─── 构造 ───────────────────────────────────────────────────────────────── +@pytest.mark.asyncio +async def test_check_inbound_uses_default_when_config_is_int_zero(): + """旧 schema 中 bot_loop_budget 为 int 0 时,应兼容且不限制。""" + cache_port = MagicMock() + cache_port.get = AsyncMock(return_value=Nothing()) + cache_port.set = AsyncMock() + guard = _make_guard( + cache_port=cache_port, + config_port=_make_config_port(value=0), + ) + + result = await guard.checkInbound("account_1", "session_1", is_user_initiated=False) + + assert result is None + cache_port.get.assert_not_called() + cache_port.set.assert_not_called() -@pytest.mark.unit -class TestBotLoopBudgetGuardConstruction: - def test_construct_stores_dependencies(self, fake_cache, fake_config, fake_logger): - # Arrange / Act - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Assert - assert guard.cache_port is fake_cache - assert guard.config_port is fake_config - assert guard._logger is fake_logger +@pytest.mark.asyncio +async def test_check_inbound_user_initiated_resets_budget(): + cache_port = MagicMock() + cache_port.set = AsyncMock() + guard = _make_guard( + cache_port=cache_port, + config_port=_make_config_port(value={"max_replies_per_hour": 10, "cooldown_seconds": 30}), + ) + + result = await guard.checkInbound("account_1", "session_1", is_user_initiated=True) + + assert result is None + cache_port.set.assert_awaited_once_with("bot_loop_budget:account_1:session_1", 10, 30) -# ─── consumeOutbound ────────────────────────────────────────────────────── +@pytest.mark.asyncio +async def test_check_inbound_reply_exceeds_budget_raises(): + cache_port = MagicMock() + cache_port.get = AsyncMock(return_value=Some(0)) + guard = _make_guard( + cache_port=cache_port, + config_port=_make_config_port(value={"max_replies_per_hour": 5, "cooldown_seconds": 60}), + ) + + with pytest.raises(BotLoopBudgetExceededError): + await guard.checkInbound("account_1", "session_1", is_user_initiated=False) -@pytest.mark.unit -class TestBotLoopBudgetGuardConsumeOutbound: - @pytest.mark.asyncio - async def test_consume_decrements_count(self, fake_cache, fake_config, fake_logger): - # Arrange: 当前计数 5,默认配置 max=20 cooldown=60 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = Some(5) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: 写入 4,TTL=60 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 4, 60) +@pytest.mark.asyncio +async def test_check_inbound_invalid_config_type_raises_validation_error(): + guard = _make_guard(config_port=_make_config_port(value="invalid")) - @pytest.mark.asyncio - async def test_consume_count_one_to_zero(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = Some(1) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: 写入 0,不抛错 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 0, 60) - - @pytest.mark.asyncio - async def test_consume_zero_raises(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = Some(0) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert - with pytest.raises(BotLoopBudgetExceededError): - await guard.consumeOutbound("acc-1", "sess-1") - fake_cache.set.assert_not_awaited() - - @pytest.mark.asyncio - async def test_consume_negative_raises(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = Some(-3) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert - with pytest.raises(BotLoopBudgetExceededError): - await guard.consumeOutbound("acc-1", "sess-1") - - @pytest.mark.asyncio - async def test_consume_cache_miss_uses_max_budget(self, fake_cache, fake_config, fake_logger): - # Arrange: cache-miss 视为满预算 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = None - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: 写入 max_budget - 1 = 19 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 19, 60) - - @pytest.mark.asyncio - async def test_consume_uses_custom_config(self, fake_cache, fake_config, fake_logger): - # Arrange: 自定义 max=5, cooldown=30 - fake_config.get.return_value = _make_budget_config(max_replies=5, cooldown=30) - fake_cache.get.return_value = None - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: 写入 4,TTL=30 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 4, 30) - - @pytest.mark.asyncio - async def test_consume_config_missing_max_uses_default(self, fake_cache, fake_config, fake_logger): - # Arrange: 配置只有 cooldown,缺 max_replies_per_hour - fake_config.get.return_value = _make_budget_config(cooldown=30) - fake_cache.get.return_value = None - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: max=20(默认),写入 19,TTL=30 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 19, 30) - - @pytest.mark.asyncio - async def test_consume_config_missing_cooldown_uses_default(self, fake_cache, fake_config, fake_logger): - # Arrange: 配置只有 max,缺 cooldown_seconds - fake_config.get.return_value = _make_budget_config(max_replies=5) - fake_cache.get.return_value = None - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: cooldown=60(默认),TTL=60 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 4, 60) - - @pytest.mark.asyncio - async def test_consume_uses_account_config_scope(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = None - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: 配置作用域为 ACCOUNT - fake_config.get.assert_awaited_once_with("bot_loop_budget", ConfigScope.ACCOUNT, "acc-1") - - @pytest.mark.asyncio - async def test_consume_budget_exceeded_carries_dict(self, fake_cache, fake_config, fake_logger): - # Arrange: 自定义 max=5, cooldown=30,计数=0 - fake_config.get.return_value = _make_budget_config(max_replies=5, cooldown=30) - fake_cache.get.return_value = Some(0) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert - with pytest.raises(BotLoopBudgetExceededError) as exc_info: - await guard.consumeOutbound("acc-1", "sess-1") - assert exc_info.value.budget == {"max": 5, "cooldown": 30} - - @pytest.mark.asyncio - async def test_consume_cache_get_dependency_error_degrades(self, fake_cache, fake_config, fake_logger): - # Arrange: cache.get 抛 DependencyError,降级为 cache-miss - fake_config.get.return_value = _make_budget_config() - fake_cache.get.side_effect = DependencyError("cache", RuntimeError("down")) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act: 降级为 miss,按满预算处理,允许消耗 - await guard.consumeOutbound("acc-1", "sess-1") - # Assert: 仍写入缓存(max-1=19),logger.warn 被调用 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 19, 60) - fake_logger.warn.assert_awaited() - - @pytest.mark.asyncio - async def test_consume_cache_get_non_dependency_error_propagates(self, fake_cache, fake_config, fake_logger): - # Arrange: cache.get 抛非 DependencyError,不捕获 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.side_effect = ValueError("unexpected") - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert: 异常向上传播 - with pytest.raises(ValueError): - await guard.consumeOutbound("acc-1", "sess-1") + with pytest.raises(ValidationError): + await guard.checkInbound("account_1", "session_1", is_user_initiated=False) -# ─── checkInbound ───────────────────────────────────────────────────────── +@pytest.mark.asyncio +async def test_consume_outbound_unlimited_when_config_zero(): + cache_port = MagicMock() + guard = _make_guard( + cache_port=cache_port, + config_port=_make_config_port(value=0), + ) + + result = await guard.consumeOutbound("account_1", "session_1") + + assert result is None + cache_port.get.assert_not_called() + cache_port.set.assert_not_called() -@pytest.mark.unit -class TestBotLoopBudgetGuardCheckInbound: - @pytest.mark.asyncio - async def test_user_initiated_resets_budget(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act: 用户主动发起 → 重置为 max_budget - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=True) - # Assert: 写入 max_budget=20,TTL=60 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 20, 60) +@pytest.mark.asyncio +async def test_consume_outbound_decrements_budget(): + cache_port = MagicMock() + cache_port.get = AsyncMock(return_value=Some(3)) + cache_port.set = AsyncMock() + guard = _make_guard( + cache_port=cache_port, + config_port=_make_config_port(value={"max_replies_per_hour": 5, "cooldown_seconds": 60}), + ) - @pytest.mark.asyncio - async def test_user_initiated_does_not_check_depletion(self, fake_cache, fake_config, fake_logger): - # Arrange: 即使 cache.get 返回 0,用户主动发起也不抛错 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = Some(0) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=True) - # Assert: 早返回,不调用 cache.get - fake_cache.get.assert_not_awaited() - fake_cache.set.assert_awaited_once() + await guard.consumeOutbound("account_1", "session_1") - @pytest.mark.asyncio - async def test_user_initiated_uses_custom_config(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config(max_replies=5, cooldown=30) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=True) - # Assert: 重置为 5,TTL=30 - fake_cache.set.assert_awaited_once_with("bot_loop_budget:acc-1:sess-1", 5, 30) + cache_port.set.assert_awaited_once_with("bot_loop_budget:account_1:session_1", 2, 60) - @pytest.mark.asyncio - async def test_reply_when_budget_positive_allows(self, fake_cache, fake_config, fake_logger): - # Arrange: 回复 Bot,current=3 > 0 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = 3 - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) - # Assert: 允许,不写缓存 - fake_cache.set.assert_not_awaited() - @pytest.mark.asyncio - async def test_reply_when_budget_zero_raises(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = Some(0) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert - with pytest.raises(BotLoopBudgetExceededError): - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) +@pytest.mark.asyncio +async def test_release_outbound_unlimited_when_config_zero(): + cache_port = MagicMock() + guard = _make_guard( + cache_port=cache_port, + config_port=_make_config_port(value=0), + ) - @pytest.mark.asyncio - async def test_reply_when_budget_negative_raises(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = Some(-1) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert - with pytest.raises(BotLoopBudgetExceededError): - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) + result = await guard.releaseOutbound("account_1", "session_1") - @pytest.mark.asyncio - async def test_reply_cache_miss_allows(self, fake_cache, fake_config, fake_logger): - # Arrange: cache-miss 视为未耗尽 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = None - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) - # Assert: 允许,不写缓存 - fake_cache.set.assert_not_awaited() + assert result is None + cache_port.get.assert_not_called() + cache_port.set.assert_not_called() - @pytest.mark.asyncio - async def test_reply_cache_get_dependency_error_degrades(self, fake_cache, fake_config, fake_logger): - # Arrange: cache.get 抛 DependencyError,降级为 miss,允许 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.side_effect = DependencyError("cache", RuntimeError("down")) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) - # Assert: 降级为 None,视为未耗尽,允许 - fake_cache.set.assert_not_awaited() - fake_logger.warn.assert_awaited() - @pytest.mark.asyncio - async def test_reply_cache_get_non_dependency_error_propagates(self, fake_cache, fake_config, fake_logger): - # Arrange: cache.get 抛非 DependencyError,不捕获 - fake_config.get.return_value = _make_budget_config() - fake_cache.get.side_effect = ValueError("unexpected") - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert: 异常向上传播 - with pytest.raises(ValueError): - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) +@pytest.mark.asyncio +async def test_release_outbound_increments_budget(): + cache_port = MagicMock() + cache_port.get = AsyncMock(return_value=Some(3)) + cache_port.set = AsyncMock() + guard = _make_guard( + cache_port=cache_port, + config_port=_make_config_port(value={"max_replies_per_hour": 5, "cooldown_seconds": 60}), + ) - @pytest.mark.asyncio - async def test_check_uses_account_config_scope(self, fake_cache, fake_config, fake_logger): - # Arrange - fake_config.get.return_value = _make_budget_config() - fake_cache.get.return_value = 3 - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) - # Assert: 配置作用域为 ACCOUNT - fake_config.get.assert_awaited_once_with("bot_loop_budget", ConfigScope.ACCOUNT, "acc-1") + await guard.releaseOutbound("account_1", "session_1") - @pytest.mark.asyncio - async def test_reply_budget_exceeded_carries_dict(self, fake_cache, fake_config, fake_logger): - # Arrange: 自定义 max=5, cooldown=30,current=0 - fake_config.get.return_value = _make_budget_config(max_replies=5, cooldown=30) - fake_cache.get.return_value = Some(0) - guard = BotLoopBudgetGuard(fake_cache, fake_config, fake_logger) - # Act & Assert - with pytest.raises(BotLoopBudgetExceededError) as exc_info: - await guard.checkInbound("acc-1", "sess-1", is_user_initiated=False) - assert exc_info.value.budget == {"max": 5, "cooldown": 30} + cache_port.set.assert_awaited_once_with("bot_loop_budget:account_1:session_1", 4, 60) diff --git a/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_outbound_adapter.py b/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_outbound_adapter.py index 53b6d32a..f368fcd4 100644 --- a/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_outbound_adapter.py +++ b/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_outbound_adapter.py @@ -13,8 +13,9 @@ from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channels.contract.dtos.common import Attachment, MessageContent, MessageFormat +from yuxi.channels.contract.dtos.config import ConfigScope from yuxi.channels.contract.dtos.outbound import FormattedMessage -from yuxi.channels.contract.errors import NotImplementedError, ValidationError +from yuxi.channels.contract.errors import NotImplementedError, NotFoundError, ValidationError from yuxi.channels.plugins.wechat_ilink.adapters.outbound_adapter import ( WeChatILinkOutboundAdapter, ) @@ -31,6 +32,13 @@ def _make_logger() -> MagicMock: return logger +def _make_config_port() -> AsyncMock: + """构造 ConfigPort 桩,``get`` 默认抛 NotFoundError 触发默认值回退。""" + config_port = AsyncMock() + config_port.get.side_effect = NotFoundError("config", "not_found") + return config_port + + @pytest.mark.unit class TestWeChatILinkOutboundAdapter: """WeChatILinkOutboundAdapter 单元测试。""" @@ -38,7 +46,7 @@ class TestWeChatILinkOutboundAdapter: def test_supportsOutbound_returns_true(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = adapter.supportsOutbound() # Assert @@ -47,7 +55,7 @@ class TestWeChatILinkOutboundAdapter: def test_supportsMultiPart_returns_true(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = adapter.supportsMultiPart() # Assert @@ -57,7 +65,7 @@ class TestWeChatILinkOutboundAdapter: async def test_formatOutbound_keeps_text_format(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) content = MessageContent(text="hello", format=MessageFormat.TEXT) # Act result = await adapter.formatOutbound(content) @@ -69,7 +77,7 @@ class TestWeChatILinkOutboundAdapter: async def test_formatOutbound_downgrades_markdown_to_text(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) content = MessageContent(text="# hi", format=MessageFormat.MARKDOWN) # Act result = await adapter.formatOutbound(content) @@ -80,7 +88,7 @@ class TestWeChatILinkOutboundAdapter: async def test_formatOutbound_downgrades_rich_to_text(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) content = MessageContent(text="rich", format=MessageFormat.RICH) # Act result = await adapter.formatOutbound(content) @@ -92,7 +100,7 @@ class TestWeChatILinkOutboundAdapter: # Arrange fake_ilink_client.get_context_token.return_value = "ctx-token" logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="hello", format=MessageFormat.TEXT) # Act result = await adapter.sendMessage("acct-1", "user-1", payload) @@ -105,7 +113,7 @@ class TestWeChatILinkOutboundAdapter: # Arrange fake_ilink_client.get_context_token.return_value = None logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="hello", format=MessageFormat.TEXT) # Act / Assert with pytest.raises(ValidationError): @@ -117,7 +125,7 @@ class TestWeChatILinkOutboundAdapter: fake_ilink_client.get_context_token.return_value = "ctx-token" fake_ilink_client.send_message.return_value = {} logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="hello", format=MessageFormat.TEXT) # Act result = await adapter.sendMessage("acct-1", "user-1", payload) @@ -130,7 +138,7 @@ class TestWeChatILinkOutboundAdapter: # Arrange fake_ilink_client.get_context_token.return_value = "ctx-token" logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="short", format=MessageFormat.TEXT) # Act receipts = await adapter.sendMessageMultiPart("acct-1", "user-1", payload) @@ -144,7 +152,7 @@ class TestWeChatILinkOutboundAdapter: # Arrange fake_ilink_client.get_context_token.return_value = "ctx-token" logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) # 5000 字符需切分为 2 片(4000 + 1000) long_text = "x" * 5000 payload = FormattedMessage(content=long_text, format=MessageFormat.TEXT) @@ -157,13 +165,12 @@ class TestWeChatILinkOutboundAdapter: @pytest.mark.asyncio async def test_sendMessageMultiPart_reads_max_length_from_config(self, fake_ilink_client): - # Arrange - 验证分片阈值从 ConfigPort 读取(非硬编码) + # Arrange - 验证分片阈值从 ConfigPort CHANNEL 作用域读取(非硬编码) fake_ilink_client.get_context_token.return_value = "ctx-token" - fake_ilink_client.get_channel_config.side_effect = lambda key, default=None, **kw: { - "max_message_length": 2000, - }.get(key, default) + config_port = AsyncMock() + config_port.get.return_value = MagicMock(value=2000) logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, config_port, logger) # 5000 字符按 2000 分片 → 3 片(2000 + 2000 + 1000) long_text = "x" * 5000 payload = FormattedMessage(content=long_text, format=MessageFormat.TEXT) @@ -171,13 +178,18 @@ class TestWeChatILinkOutboundAdapter: receipts = await adapter.sendMessageMultiPart("acct-1", "user-1", payload) # Assert assert len(receipts) == 3 + config_port.get.assert_awaited_once_with( + "max_message_length", + scope=ConfigScope.CHANNEL, + target="wechat_ilink", + ) @pytest.mark.asyncio async def test_sendMessageMultiPart_image_attachment_uploads_and_sends(self, fake_ilink_client): # Arrange fake_ilink_client.get_context_token.return_value = "ctx-token" logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) attachment = Attachment( type="image", url="wechat_ilink://media?x=1", @@ -202,7 +214,7 @@ class TestWeChatILinkOutboundAdapter: # Arrange fake_ilink_client.get_context_token.return_value = "ctx-token" logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) attachment = Attachment(type="image", url="x", content=None) payload = FormattedMessage( content="text", @@ -220,7 +232,7 @@ class TestWeChatILinkOutboundAdapter: # Arrange fake_ilink_client.get_context_token.return_value = "ctx-token" logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) attachment = Attachment( type="file", url="x", @@ -246,7 +258,7 @@ class TestWeChatILinkOutboundAdapter: # Arrange fake_ilink_client.get_context_token.return_value = None logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="hi", format=MessageFormat.TEXT) # Act / Assert with pytest.raises(ValidationError): @@ -256,7 +268,7 @@ class TestWeChatILinkOutboundAdapter: async def test_updateMessage_raises_ValidationError(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="hi", format=MessageFormat.TEXT) # Act / Assert with pytest.raises(ValidationError): @@ -265,7 +277,7 @@ class TestWeChatILinkOutboundAdapter: def test_supportsThreadReply_returns_false(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = adapter.supportsThreadReply() # Assert @@ -276,7 +288,7 @@ class TestWeChatILinkOutboundAdapter: async def test_sendThreadMessage_raises_NotImplementedError(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="hi", format=MessageFormat.TEXT) # Act / Assert with pytest.raises(NotImplementedError): @@ -285,7 +297,7 @@ class TestWeChatILinkOutboundAdapter: def test_supportsBatchSend_returns_false(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = adapter.supportsBatchSend() # Assert @@ -296,7 +308,7 @@ class TestWeChatILinkOutboundAdapter: async def test_batchSendMessages_raises_NotImplementedError(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkOutboundAdapter(fake_ilink_client, logger) + adapter = WeChatILinkOutboundAdapter(fake_ilink_client, _make_config_port(), logger) payload = FormattedMessage(content="hi", format=MessageFormat.TEXT) # Act / Assert with pytest.raises(NotImplementedError): diff --git a/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_puller_adapter.py b/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_puller_adapter.py index b72fa30d..f9431d0b 100644 --- a/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_puller_adapter.py +++ b/backend/test/unit/channels/plugins/wechat_ilink/adapters/test_puller_adapter.py @@ -10,6 +10,8 @@ from __future__ import annotations from unittest.mock import AsyncMock, MagicMock import pytest +from yuxi.channels.contract.dtos.config import ConfigScope +from yuxi.channels.contract.errors import NotFoundError from yuxi.channels.contract.errors.transport import TransportError from yuxi.channels.plugins.wechat_ilink.adapters.puller_adapter import ( WeChatILinkPullerAdapter, @@ -27,6 +29,13 @@ def _make_logger() -> MagicMock: return logger +def _make_config_port() -> AsyncMock: + """构造 ConfigPort 桩,``get`` 默认抛 NotFoundError 触发默认值回退。""" + config_port = AsyncMock() + config_port.get.side_effect = NotFoundError("config", "not_found") + return config_port + + @pytest.mark.unit class TestWeChatILinkPullerAdapter: """WeChatILinkPullerAdapter 单元测试。""" @@ -39,7 +48,7 @@ class TestWeChatILinkPullerAdapter: "get_updates_buf": "next-cursor", } logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.poll("acct-1", "cursor-prev") # Assert @@ -52,7 +61,7 @@ class TestWeChatILinkPullerAdapter: async def test_poll_uses_empty_string_when_cursor_none(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.poll("acct-1", None) # Assert @@ -64,7 +73,7 @@ class TestWeChatILinkPullerAdapter: # Arrange fake_ilink_client.get_updates.side_effect = RuntimeError("api down") logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.poll("acct-1", "cursor") # Assert @@ -78,7 +87,7 @@ class TestWeChatILinkPullerAdapter: # Arrange fake_ilink_client.get_updates.return_value = {"get_updates_buf": "c1"} logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.poll("acct-1", "c0") # Assert @@ -90,7 +99,7 @@ class TestWeChatILinkPullerAdapter: # Arrange fake_ilink_client.get_updates.return_value = {"msgs": []} logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.poll("acct-1", "old-cursor") # Assert @@ -100,7 +109,7 @@ class TestWeChatILinkPullerAdapter: async def test_getPollingConfig_returns_expected_config(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act config = await adapter.getPollingConfig() # Assert @@ -110,23 +119,27 @@ class TestWeChatILinkPullerAdapter: @pytest.mark.asyncio async def test_getPollingConfig_reads_long_poll_timeout_from_config(self, fake_ilink_client): - # Arrange - 验证 long_poll_timeout_ms 从 ConfigPort 读取 - fake_ilink_client.get_channel_config.side_effect = lambda key, default=None, **kw: { - "long_poll_timeout_ms": 50000, - }.get(key, default) + # Arrange - 验证 long_poll_timeout_ms 从 ConfigPort CHANNEL 作用域读取 + config_port = AsyncMock() + config_port.get.return_value = MagicMock(value=50000) logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, config_port, logger) # Act config = await adapter.getPollingConfig() # Assert assert config.long_poll_timeout_ms == 50000 + config_port.get.assert_awaited_once_with( + "long_poll_timeout_ms", + scope=ConfigScope.CHANNEL, + target="wechat_ilink", + ) @pytest.mark.asyncio async def test_poll_logs_trace_id_on_error(self, fake_ilink_client): # Arrange - 验证 poll 失败时日志携带 trace_id fake_ilink_client.get_updates.side_effect = RuntimeError("api down") logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.poll("acct-1", "cursor") # Assert - warn 日志被调用且携带 trace_id 关键字参数 @@ -140,7 +153,7 @@ class TestWeChatILinkPullerAdapter: async def test_acknowledge_is_noop(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.acknowledge("acct-1", "cursor") # Assert @@ -150,7 +163,7 @@ class TestWeChatILinkPullerAdapter: async def test_onTransportReset_is_noop(self, fake_ilink_client): # Arrange logger = _make_logger() - adapter = WeChatILinkPullerAdapter(fake_ilink_client, logger) + adapter = WeChatILinkPullerAdapter(fake_ilink_client, _make_config_port(), logger) # Act result = await adapter.onTransportReset("acct-1") # Assert 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 c7ab07b2..4ef39068 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 @@ -28,6 +28,8 @@ def _make_logger() -> MagicMock: logger.info = AsyncMock() logger.warn = AsyncMock() logger.error = AsyncMock() + logger.debug = AsyncMock() + logger.exception = AsyncMock() return logger @@ -148,7 +150,7 @@ class TestWeChatWocInboundAdapter: @pytest.mark.asyncio async def test_normalizeInbound_text_message_does_not_require_account_id(self): - # Arrange — 文本消息无附件,account_id 不影响规范化结果 + # Arrange — 验证 account_id 缺失时文本消息仍能规范化 client = _make_client() adapter = WeChatWocInboundAdapter(client, _make_logger()) raw = _make_raw_event(_make_text_payload(), account_id=None) @@ -159,3 +161,26 @@ class TestWeChatWocInboundAdapter: # Assert assert content.text == "hello" assert content.attachments == () + + @pytest.mark.asyncio + async def test_normalizeInbound_falls_back_sender_to_talker_when_empty(self): + # Arrange — bridge 对 P2P 消息或自己发送的消息不填充 sender + client = _make_client() + logger = _make_logger() + adapter = WeChatWocInboundAdapter(client, logger) + payload = _make_text_payload() + payload["sender"] = "" + raw = _make_raw_event(payload, account_id="acct-1") + + # Act + content = await adapter.normalizeInbound(raw) + + # Assert — sender 回退到 talker,不抛异常 + assert content.metadata["sender"] == payload["talker"] + assert content.metadata["is_sender"] == payload["is_sender"] + logger.debug.assert_any_call( + "woc inbound normalize: sender fallback to talker", + account_id="acct-1", + msg_id="m1", + talker=payload["talker"], + ) diff --git a/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter.py b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter.py index 620a9783..4478002a 100644 --- a/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter.py +++ b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter.py @@ -25,6 +25,8 @@ def _make_logger() -> MagicMock: logger.info = AsyncMock() logger.warn = AsyncMock() logger.error = AsyncMock() + logger.debug = AsyncMock() + logger.exception = AsyncMock() return logger diff --git a/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter_cursor.py b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter_cursor.py index 889cf5a6..a6c8da74 100644 --- a/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter_cursor.py +++ b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_puller_adapter_cursor.py @@ -21,10 +21,13 @@ pytestmark = pytest.mark.unit def _make_logger() -> MagicMock: + """构造 LoggerPort 桩。""" logger = MagicMock() logger.info = AsyncMock() logger.warn = AsyncMock() logger.error = AsyncMock() + logger.debug = AsyncMock() + logger.exception = AsyncMock() return logger diff --git a/backend/test/unit/channels/plugins/wechat_woc/adapters/test_stream_connector_adapter.py b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_stream_connector_adapter.py new file mode 100644 index 00000000..54c80897 --- /dev/null +++ b/backend/test/unit/channels/plugins/wechat_woc/adapters/test_stream_connector_adapter.py @@ -0,0 +1,562 @@ +"""yuxi.channels.plugins.wechat_woc.adapters.stream_connector_adapter 单元测试。 + +覆盖 ``WeChatWocStreamConnectorAdapter`` 的 ``connect`` / ``ping`` / +``getStreamConfig`` 方法,以及 ``receive`` 闭包对 SSE 四类事件 +(sync / messages / status / heartbeat)与断线场景的处理。 + +通过 ``FakeSseStreamHandle`` 模拟 SSE 事件流,不发起真实网络请求。 +通过 ``FakeConfigPort`` 模拟 ConfigPort CHANNEL 作用域配置读取。 +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest +from yuxi.channels.contract.errors import NotFoundError +from yuxi.channels.contract.errors.transport import TransportError +from yuxi.channels.plugins.wechat_woc.adapters.stream_connector_adapter import ( + WeChatWocStreamConnectorAdapter, +) + +pytestmark = pytest.mark.unit + + +class FakeSseStreamHandle: + """模拟 ``SseStreamHandle``,支持配置事件序列与异常注入。 + + - ``events`` 为事件列表时,``events()`` 依次 yield 这些事件后结束 + - ``error`` 非 None 时,``events()`` 首次迭代即抛出该异常 + - ``aclose`` 调用后 ``closed=True``,供测试断言 + - 接受可选 logger 参数(与真实 SseStreamHandle 构造签名一致) + """ + + def __init__( + self, + events: list[dict[str, Any]] | None = None, + error: Exception | None = None, + trace_id: str = "test-trace-id", + logger: Any = None, + ) -> None: + self._events = events + self._error = error + self.trace_id = trace_id + self._logger = logger + self.closed = False + + async def events(self): + if self._error is not None: + raise self._error + if self._events is not None: + for event in self._events: + yield event + + async def aclose(self) -> None: + self.closed = True + + +class FakeConfigPort: + """模拟 ConfigPort,支持按 key 配置 CHANNEL 作用域返回值。 + + - ``values`` 字典中的 key 命中时返回 ``ConfigValue(value=...)`` + - 未命中时抛 ``NotFoundError``(模拟配置未设置场景) + """ + + def __init__(self, values: dict[str, Any] | None = None) -> None: + self._values = values or {} + + async def get(self, key: str, scope: Any = None, target: Any = None) -> Any: + if key in self._values: + # 返回带 value 属性的简单对象(模拟 ConfigValue) + cv = MagicMock() + cv.value = self._values[key] + return cv + raise NotFoundError(resource="config", id=f"{target}:{key}") + + +def _make_logger() -> MagicMock: + """构造 LoggerPort 桩。""" + logger = MagicMock() + logger.info = AsyncMock() + logger.warn = AsyncMock() + logger.error = AsyncMock() + logger.debug = AsyncMock() + logger.exception = AsyncMock() + return logger + + +def _make_client(handle: FakeSseStreamHandle | None = None) -> AsyncMock: + """构造 WocBridgeClient 桩。 + + ``stream_messages`` 返回传入的 ``handle``;``get_messages_since`` + 默认返回空消息列表(无补全消息),可由测试覆写。 + """ + client = AsyncMock() + if handle is not None: + client.stream_messages.return_value = handle + client.get_messages_since.return_value = {"messages": [], "next_cursor": None} + return client + + +def _make_adapter( + handle: FakeSseStreamHandle | None = None, + config_values: dict[str, Any] | None = None, +) -> tuple[WeChatWocStreamConnectorAdapter, AsyncMock, FakeConfigPort, MagicMock]: + """构造 adapter 与依赖桩,返回 (adapter, client, config_port, logger)。""" + client = _make_client(handle) + config_port = FakeConfigPort(config_values) + logger = _make_logger() + adapter = WeChatWocStreamConnectorAdapter(client, config_port, logger) + return adapter, client, config_port, logger + + +@pytest.mark.unit +class TestWeChatWocStreamConnectorAdapter: + """WeChatWocStreamConnectorAdapter 单元测试。""" + + # ------------------------------------------------------------------ + # getStreamConfig / ping + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_getStreamConfig_returns_sse_config_from_config_port(self): + """getStreamConfig 从 ConfigPort 读取配置。""" + adapter, _, _, _ = _make_adapter( + config_values={ + "sse_heartbeat_interval_ms": 25000, + "sse_connect_timeout_ms": 8000, + }, + ) + + config = await adapter.getStreamConfig() + + assert config.heartbeat_interval_ms == 25000 + assert config.connect_timeout_ms == 8000 + + @pytest.mark.asyncio + async def test_getStreamConfig_falls_back_to_defaults_when_not_set(self): + """ConfigPort 未命中时回退到 _constants 默认值。""" + adapter, _, _, _ = _make_adapter(config_values=None) + + config = await adapter.getStreamConfig() + + assert config.heartbeat_interval_ms == 30000 + assert config.connect_timeout_ms == 10000 + + @pytest.mark.asyncio + async def test_ping_returns_none(self): + adapter, _, _, _ = _make_adapter() + + result = await adapter.ping(MagicMock()) + + assert result is None + + # ------------------------------------------------------------------ + # connect → StreamConnection 基本结构 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_connect_returns_stream_connection_with_callbacks(self): + handle = FakeSseStreamHandle(events=[], logger=_make_logger()) + adapter, client, _, _ = _make_adapter(handle) + + conn = await adapter.connect("acct-1", token=None) + + assert callable(conn.receive) + assert callable(conn.send) + assert callable(conn.close) + assert conn.extra == handle.trace_id + client.stream_messages.assert_awaited_once_with("acct-1", read_timeout=None) + + @pytest.mark.asyncio + async def test_send_is_noop(self): + handle = FakeSseStreamHandle(events=[], logger=_make_logger()) + adapter, _, _, _ = _make_adapter(handle) + + conn = await adapter.connect("acct-1", token=None) + + result = await conn.send({"type": "text", "content": "hello"}) + assert result is None + + @pytest.mark.asyncio + async def test_close_calls_handle_aclose(self): + handle = FakeSseStreamHandle(events=[], logger=_make_logger()) + adapter, _, _, _ = _make_adapter(handle) + + conn = await adapter.connect("acct-1", token=None) + await conn.close() + + assert handle.closed is True + + # ------------------------------------------------------------------ + # receive: messages 事件 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_receive_messages_event_returns_one_msg_at_a_time(self): + """messages 事件拆分数组,逐条返回。""" + handle = FakeSseStreamHandle( + events=[ + { + "event": "messages", + "data": { + "messages": [ + {"msg_id": "m1", "content": "hello"}, + {"msg_id": "m2", "content": "world"}, + ], + "next_cursor": "2000", + }, + }, + ], + logger=_make_logger(), + ) + adapter, _, _, _ = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + msg1 = await conn.receive() + msg2 = await conn.receive() + + assert msg1 == {"msg_id": "m1", "content": "hello"} + assert msg2 == {"msg_id": "m2", "content": "world"} + + # ------------------------------------------------------------------ + # receive: sync 事件 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_receive_sync_event_triggers_backfill(self): + """sync 事件触发 /api/messages/since 补全,补全消息逐条返回。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "sync", "data": {"cursor": "1000"}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m_new"}]}, + }, + ], + logger=_make_logger(), + ) + adapter, client, _, _ = _make_adapter(handle) + # 补全返回 2 条消息 + client.get_messages_since.return_value = { + "messages": [ + {"msg_id": "m_backfill_1"}, + {"msg_id": "m_backfill_2"}, + ], + "next_cursor": None, + } + conn = await adapter.connect("acct-1", token=None) + + msg1 = await conn.receive() + msg2 = await conn.receive() + msg3 = await conn.receive() + + # 前两条来自补全 + assert msg1 == {"msg_id": "m_backfill_1"} + assert msg2 == {"msg_id": "m_backfill_2"} + # 第三条来自后续 messages 事件 + assert msg3 == {"msg_id": "m_new"} + # 验证补全调用了正确的 cursor + client.get_messages_since.assert_awaited_once_with("acct-1", "1000", 50) + + @pytest.mark.asyncio + async def test_receive_sync_only_processed_once(self): + """第二个 sync 事件被忽略,不重复触发补全。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "sync", "data": {"cursor": "1000"}}, + {"event": "sync", "data": {"cursor": "2000"}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m1"}]}, + }, + ], + logger=_make_logger(), + ) + adapter, client, _, _ = _make_adapter(handle) + client.get_messages_since.return_value = {"messages": [], "next_cursor": None} + conn = await adapter.connect("acct-1", token=None) + + msg = await conn.receive() + + assert msg == {"msg_id": "m1"} + # 补全只调用了一次(第一个 sync 触发,第二个 sync 忽略) + client.get_messages_since.assert_awaited_once() + + @pytest.mark.asyncio + async def test_receive_sync_cursor_zero_triggers_backfill(self): + """cursor=0 也应触发补全(边界判断修复 L-3)。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "sync", "data": {"cursor": 0}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m1"}]}, + }, + ], + logger=_make_logger(), + ) + adapter, client, _, _ = _make_adapter(handle) + client.get_messages_since.return_value = {"messages": [], "next_cursor": None} + conn = await adapter.connect("acct-1", token=None) + + msg = await conn.receive() + + assert msg == {"msg_id": "m1"} + # cursor=0 也触发了补全 + client.get_messages_since.assert_awaited_once_with("acct-1", "0", 50) + + # ------------------------------------------------------------------ + # receive: heartbeat 事件 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_receive_heartbeat_event_skipped(self): + """heartbeat 事件被跳过,继续读下一条事件。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "heartbeat", "data": {}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m1"}]}, + }, + ], + logger=_make_logger(), + ) + adapter, _, _, _ = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + msg = await conn.receive() + + assert msg == {"msg_id": "m1"} + + # ------------------------------------------------------------------ + # receive: status 事件 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_receive_status_db_inaccessible_raises_transport_error(self): + """status 事件 db_accessible=false 抛 TransportError(transient)。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "status", "data": {"db_accessible": False, "db_error_code": "DB_ENCRYPTED"}}, + ], + logger=_make_logger(), + ) + adapter, _, _, _ = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + with pytest.raises(TransportError) as exc_info: + await conn.receive() + + assert exc_info.value.category == "transient" + assert "DB inaccessible" in exc_info.value.message + + @pytest.mark.asyncio + async def test_receive_status_db_accessible_continues(self): + """status 事件 db_accessible=true 时跳过,继续读下一条。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "status", "data": {"db_accessible": True}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m1"}]}, + }, + ], + logger=_make_logger(), + ) + adapter, _, _, _ = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + msg = await conn.receive() + + assert msg == {"msg_id": "m1"} + + # ------------------------------------------------------------------ + # receive: 断线场景 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_receive_sse_disconnect_raises_transport_error(self): + """SSE 连接正常关闭(迭代结束)抛 TransportError(transient)。""" + handle = FakeSseStreamHandle(events=[], logger=_make_logger()) + adapter, _, _, _ = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + with pytest.raises(TransportError) as exc_info: + await conn.receive() + + assert exc_info.value.category == "transient" + assert "closed" in exc_info.value.message + + @pytest.mark.asyncio + async def test_receive_httpx_error_raises_transport_error(self): + """SSE 读取过程 httpx 异常抛 TransportError(transient)。""" + handle = FakeSseStreamHandle( + error=httpx.ConnectError("connection reset"), + logger=_make_logger(), + ) + adapter, _, _, _ = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + with pytest.raises(TransportError) as exc_info: + await conn.receive() + + assert exc_info.value.category == "transient" + assert "SSE read error" in exc_info.value.message + + @pytest.mark.asyncio + async def test_receive_json_parse_error_raises_transport_error(self): + """SSE events() 抛 JSON 解析异常时,receive 翻译为 TransportError(transient)。 + + 覆盖 H-1 修复:JSON 解析失败不再静默吞,而是触发重连。 + """ + # 构造一个 events() 抛 ValueError 的 handle + handle = FakeSseStreamHandle( + error=ValueError("Expecting value: line 1 column 1 (char 0)"), + logger=_make_logger(), + ) + adapter, _, _, _ = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + with pytest.raises(TransportError) as exc_info: + await conn.receive() + + assert exc_info.value.category == "transient" + assert "SSE unexpected error" in exc_info.value.message + + # ------------------------------------------------------------------ + # receive: 复合场景 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_receive_multiple_event_types_in_sequence(self): + """复合事件序列:sync → heartbeat → messages → heartbeat → messages。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "sync", "data": {"cursor": "500"}}, + {"event": "heartbeat", "data": {}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m1"}]}, + }, + {"event": "heartbeat", "data": {}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m2"}, {"msg_id": "m3"}]}, + }, + ], + logger=_make_logger(), + ) + adapter, client, _, _ = _make_adapter(handle) + client.get_messages_since.return_value = {"messages": [], "next_cursor": None} + conn = await adapter.connect("acct-1", token=None) + + msg1 = await conn.receive() + msg2 = await conn.receive() + msg3 = await conn.receive() + + assert msg1 == {"msg_id": "m1"} + assert msg2 == {"msg_id": "m2"} + assert msg3 == {"msg_id": "m3"} + + @pytest.mark.asyncio + async def test_receive_unknown_event_type_skipped(self): + """未知事件类型被跳过并记录日志。""" + handle = FakeSseStreamHandle( + events=[ + {"event": "unknown_event", "data": {"foo": "bar"}}, + { + "event": "messages", + "data": {"messages": [{"msg_id": "m1"}]}, + }, + ], + logger=_make_logger(), + ) + adapter, _, _, logger = _make_adapter(handle) + conn = await adapter.connect("acct-1", token=None) + + msg = await conn.receive() + + assert msg == {"msg_id": "m1"} + logger.debug.assert_awaited() + + # ------------------------------------------------------------------ + # backfill: 多页补全 + next_cursor 类型一致性 + # ------------------------------------------------------------------ + + @pytest.mark.asyncio + async def test_backfill_paginates_until_no_more_messages(self): + """补全循环拉取直到 next_cursor 为 None 或返回空批次。""" + handle = FakeSseStreamHandle( + events=[{"event": "sync", "data": {"cursor": "1000"}}], + logger=_make_logger(), + ) + adapter, client, _, _ = _make_adapter(handle) + # 模拟分页:第一页有消息 + next_cursor,第二页有消息 + next_cursor=None + client.get_messages_since.side_effect = [ + { + "messages": [{"msg_id": "m1"}], + "next_cursor": "2000", + }, + { + "messages": [{"msg_id": "m2"}], + "next_cursor": None, + }, + ] + conn = await adapter.connect("acct-1", token=None) + + msg1 = await conn.receive() + msg2 = await conn.receive() + + assert msg1 == {"msg_id": "m1"} + assert msg2 == {"msg_id": "m2"} + assert client.get_messages_since.await_count == 2 + + @pytest.mark.asyncio + async def test_backfill_handles_int_next_cursor(self): + """bridge 返回 int 类型 next_cursor 时也能正确比较(M-3 修复)。""" + handle = FakeSseStreamHandle( + events=[{"event": "sync", "data": {"cursor": "1000"}}], + logger=_make_logger(), + ) + adapter, client, _, _ = _make_adapter(handle) + # 第一页返回 int 类型 next_cursor=2000,第二页无更多 + client.get_messages_since.side_effect = [ + { + "messages": [{"msg_id": "m1"}], + "next_cursor": 2000, # int 类型 + }, + { + "messages": [{"msg_id": "m2"}], + "next_cursor": None, + }, + ] + conn = await adapter.connect("acct-1", token=None) + + msg1 = await conn.receive() + msg2 = await conn.receive() + + assert msg1 == {"msg_id": "m1"} + assert msg2 == {"msg_id": "m2"} + assert client.get_messages_since.await_count == 2 + + @pytest.mark.asyncio + async def test_backfill_uses_config_port_max_batch_size(self): + """_backfill 的 limit 从 ConfigPort 读取 max_batch_size(L-2 修复)。""" + handle = FakeSseStreamHandle( + events=[{"event": "sync", "data": {"cursor": "1000"}}], + logger=_make_logger(), + ) + adapter, client, _, _ = _make_adapter(handle, config_values={"max_batch_size": 100}) + client.get_messages_since.return_value = {"messages": [], "next_cursor": None} + conn = await adapter.connect("acct-1", token=None) + + # 触发 receive 直到补全完成(返回空列表后抛 TransportError) + with pytest.raises(TransportError): + await conn.receive() + + # 验证 limit 参数从 ConfigPort 读取为 100 + client.get_messages_since.assert_awaited_once_with("acct-1", "1000", 100) diff --git a/backend/test/unit/channels/plugins/wechat_woc/test_entry.py b/backend/test/unit/channels/plugins/wechat_woc/test_entry.py index 13b6ff7b..b7b338f4 100644 --- a/backend/test/unit/channels/plugins/wechat_woc/test_entry.py +++ b/backend/test/unit/channels/plugins/wechat_woc/test_entry.py @@ -1,7 +1,7 @@ """yuxi.channels.plugins.wechat_woc.entry 单元测试。 覆盖 ``channel_entry(host, manifest)`` 入口函数(F-03 单真相源新签名): -- 适配器注册(12 个适配器) +- 适配器注册(13 个适配器) - 生命周期钩子注册 - ``PluginManifest`` 返回值(manifest 直接复用入参,adapters 为运行时清单) @@ -77,11 +77,11 @@ class TestChannelEntry: # Assert - manifest 直接复用入参,不重新构造 assert result.manifest is woc_manifest - def test_registers_all_12_adapters(self, fake_plugin_host, woc_manifest): + def test_registers_all_13_adapters(self, fake_plugin_host, woc_manifest): # Act channel_entry(fake_plugin_host, woc_manifest) # Assert - assert fake_plugin_host.registerAdapter.call_count == 12 + assert fake_plugin_host.registerAdapter.call_count == 13 adapter_types = [call.args[0] for call in fake_plugin_host.registerAdapter.call_args_list] assert tuple(adapter_types) == _ADAPTER_TYPES @@ -98,7 +98,7 @@ class TestChannelEntry: result = channel_entry(fake_plugin_host, woc_manifest) # Assert assert result.adapters == _ADAPTER_TYPES - assert len(result.adapters) == 12 + assert len(result.adapters) == 13 def test_gets_required_ports(self, fake_plugin_host, woc_manifest): """验证 entry 调用 getConfigPort/getLoggerPort/getCachePort/getPersistencePort。""" @@ -222,4 +222,4 @@ class TestNoHardcodedDeclarations: import yuxi.channels.plugins.wechat_woc.entry as entry_mod assert hasattr(entry_mod, "_ADAPTER_TYPES") - assert len(entry_mod._ADAPTER_TYPES) == 12 + assert len(entry_mod._ADAPTER_TYPES) == 13 diff --git a/backend/test/unit/channels/plugins/wechat_woc/test_woc_bridge_client.py b/backend/test/unit/channels/plugins/wechat_woc/test_woc_bridge_client.py index c20363e4..a4031b02 100644 --- a/backend/test/unit/channels/plugins/wechat_woc/test_woc_bridge_client.py +++ b/backend/test/unit/channels/plugins/wechat_woc/test_woc_bridge_client.py @@ -2,7 +2,7 @@ 覆盖: - User-Agent 携带插件版本(F-03 单真相源) -- HTTP 超时 / 网络错误 / 5xx 服务端错误重试(P2-3 超时重试) +- HTTP 重试策略(P2-3):仅 NetworkError 重试 1 次;超时 / 5xx / 4xx 不重试 - 4xx 客户端错误不重试 """ @@ -94,40 +94,26 @@ class TestUserAgent: @pytest.mark.unit class TestHttpRetry: - """_execute_http 重试策略(P2-3)。""" + """_execute_http 重试策略(P2-3)。 + + 重试策略:仅对 ``httpx.NetworkError`` 做 1 次快速重试(不等待); + 超时与 5xx 不重试,交给 outbox 退避重试机制处理,避免双重重试放大延迟。 + """ @pytest.mark.asyncio - async def test_timeout_retries_then_succeeds(self, httpx_mock): - # Arrange - 前两次超时,第三次成功 + async def test_timeout_does_not_retry_raises_operation_timeout(self, httpx_mock): + # Arrange - 超时不重试,1 次请求即抛 OperationTimeoutError httpx_mock.add_exception(httpx.TimeoutException("timeout")) - httpx_mock.add_exception(httpx.TimeoutException("timeout")) - httpx_mock.add_response(json={"success": True, "wechat_running": True}) client = _make_client() - with patch("asyncio.sleep", new=AsyncMock()): - status = await client.get_status("acct-1") + with pytest.raises(OperationTimeoutError): + await client.get_status("acct-1") - # Assert - assert status["wechat_running"] is True + # Assert - 仅发起 1 次请求 requests = httpx_mock.get_requests() - assert len(requests) == 3 + assert len(requests) == 1 assert requests[0].headers["User-Agent"] == "yuxi-channels-wechat-woc/1.0.0" - @pytest.mark.asyncio - async def test_timeout_retries_exhausted_raises_operation_timeout(self, httpx_mock): - # Arrange - 全部 3 次超时 - httpx_mock.add_exception(httpx.TimeoutException("timeout")) - httpx_mock.add_exception(httpx.TimeoutException("timeout")) - httpx_mock.add_exception(httpx.TimeoutException("timeout")) - - client = _make_client() - with patch("asyncio.sleep", new=AsyncMock()): - with pytest.raises(OperationTimeoutError): - await client.get_status("acct-1") - - # Assert - 共发起 3 次请求 - assert len(httpx_mock.get_requests()) == 3 - @pytest.mark.asyncio async def test_network_error_retries_then_succeeds(self, httpx_mock): # Arrange - 第一次网络错误,第二次成功 @@ -143,19 +129,16 @@ class TestHttpRetry: assert len(httpx_mock.get_requests()) == 2 @pytest.mark.asyncio - async def test_5xx_retries_then_raises(self, httpx_mock): - # Arrange - 3 次 503 - httpx_mock.add_response(status_code=503) - httpx_mock.add_response(status_code=503) + async def test_5xx_does_not_retry_raises_dependency_error(self, httpx_mock): + # Arrange - 5xx 不重试,1 次请求即抛 DependencyError httpx_mock.add_response(status_code=503) client = _make_client() - with patch("asyncio.sleep", new=AsyncMock()): - with pytest.raises(DependencyError): - await client.get_status("acct-1") + with pytest.raises(DependencyError): + await client.get_status("acct-1") - # Assert - assert len(httpx_mock.get_requests()) == 3 + # Assert - 仅发起 1 次请求 + assert len(httpx_mock.get_requests()) == 1 @pytest.mark.asyncio async def test_4xx_does_not_retry(self, httpx_mock): @@ -309,8 +292,7 @@ class TestCapabilityNegotiation: @pytest.mark.asyncio async def test_negotiate_bridge_failure_returns_conservative_fallback(self, httpx_mock): - # Arrange - bridge 不可达,_execute_http 内部重试 3 次 - httpx_mock.add_exception(httpx.ConnectError("connection refused")) + # Arrange - bridge 不可达,_execute_http 对 NetworkError 重试 1 次(共 2 次) httpx_mock.add_exception(httpx.ConnectError("connection refused")) httpx_mock.add_exception(httpx.ConnectError("connection refused")) @@ -332,8 +314,7 @@ class TestCapabilityNegotiation: _CAPABILITIES_FAILURE_CACHE_TTL_SECONDS, ) - # Arrange - bridge 不可达 - httpx_mock.add_exception(httpx.ConnectError("connection refused")) + # Arrange - bridge 不可达,NetworkError 重试 1 次(共 2 次) httpx_mock.add_exception(httpx.ConnectError("connection refused")) httpx_mock.add_exception(httpx.ConnectError("connection refused"))