test: 批量完善单元测试与集成测试用例
1. 为测试桩函数新增version、debug/exception日志等必要字段 2. 重构SQLAlchemy事务上下文测试,修复隐式事务提交逻辑 3. 修复微信WOC插件适配器测试用例,补充缺失的日志方法与测试场景 4. 调整路由阶段测试,移除过时的会话刷新逻辑 5. 新增内容审核 retention 定时任务测试用例 6. 完善配对、会话、账号模型测试用例 7. 删除过时的报表集成测试文件 8. 为持久化适配器添加配对状态更新的乐观锁测试 9. 修复配置验证测试的断言逻辑
This commit is contained in:
parent
6f730dca06
commit
ff07cef881
@ -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"
|
||||
|
||||
|
||||
@ -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
|
||||
@ -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()
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 "")
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@ -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"),
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
# ─── 阶段契约属性 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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(),
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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(),
|
||||
|
||||
@ -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(),
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 查询 ──────────────────
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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"],
|
||||
)
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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)
|
||||
@ -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
|
||||
|
||||
@ -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"))
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user