test: 批量完善单元测试与集成测试用例

1. 为测试桩函数新增version、debug/exception日志等必要字段
2. 重构SQLAlchemy事务上下文测试,修复隐式事务提交逻辑
3. 修复微信WOC插件适配器测试用例,补充缺失的日志方法与测试场景
4. 调整路由阶段测试,移除过时的会话刷新逻辑
5. 新增内容审核 retention 定时任务测试用例
6. 完善配对、会话、账号模型测试用例
7. 删除过时的报表集成测试文件
8. 为持久化适配器添加配对状态更新的乐观锁测试
9. 修复配置验证测试的断言逻辑
This commit is contained in:
Kris 2026-07-08 03:58:05 +08:00
parent 6f730dca06
commit ff07cef881
38 changed files with 1574 additions and 743 deletions

View File

@ -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"

View File

@ -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

View File

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

View File

@ -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非 flushIntegrityError 从 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非 flushSQLAlchemyError 从 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

View File

@ -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"
# 不应调用 refreshlatest 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):

View File

@ -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)

View File

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

View File

@ -1,16 +1,16 @@
"""yuxi.channels.adapters.channel_persistence_adapter 路由绑定仓储契约测试。
覆盖 ``ChannelPersistenceAdapter`` 实现 ``RouteBindingRepositoryPort``
6 个方法saveRouteBinding / updateRouteBinding / getRouteBinding /
listRouteBindings / deleteRouteBinding / listEnabledByAccountpatch
5 个方法saveRouteBinding / updateRouteBinding / getRouteBinding /
listRouteBindings / deleteRouteBindingpatch
``create_repositories`` 返回 mock 聚合不连接真实 DB
测试契约点
- saveRouteBinding 正常创建返回 RouteBindingRuleIntegrityError 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_bindingversion 递增"""
# 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 ──────────────────────────────────────────────────────

View File

@ -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

View File

@ -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:
"""构造 TaskContextpayload 可按需覆盖。"""
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_factoryyield 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 "")

View File

@ -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

View File

@ -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,
)

View File

@ -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"),

View File

@ -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

View File

@ -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:
# Arrangeresolve 返回 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):
# Arrangecontext.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)
# Assertresolve 使用刷新后的会话
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
# ─── 阶段契约属性 ───────────────────────────────────────────────────────────

View File

@ -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"

View File

@ -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):
# Arrangeformatted_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):

View File

@ -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(),

View File

@ -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,

View File

@ -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-32TransportManager 为写入方)。
覆盖传输任务 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)

View File

@ -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

View File

@ -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(),

View File

@ -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(),

View File

@ -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,

View File

@ -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):

View File

@ -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

View File

@ -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 查询 ──────────────────

View File

@ -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

View File

@ -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,
)

View File

@ -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: 写入 4TTL=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: 写入 4TTL=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默认写入 19TTL=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=19logger.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=20TTL=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: 重置为 5TTL=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: 回复 Botcurrent=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=30current=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)

View File

@ -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):

View File

@ -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

View File

@ -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"],
)

View File

@ -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

View File

@ -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

View File

@ -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_sizeL-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)

View File

@ -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

View File

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