refactor(test): 批量重构测试代码,简化调度器handler测试逻辑

1.  统一替换多个调度器handler测试用例,移除冗余的session_factory相关代码和fake session上下文
2.  调整ChannelPersistenceAdapter测试,使用模块级patcher管理并统一初始化方式
3.  修正OutboxEntry聚合根操作的版本号逻辑,移除不必要的版本递增
4.  新增UTC时区转换相关测试用例,完善datetime类型映射测试
5.  更新微信WOC适配器测试,适配新的session_type枚举值
6.  优化测试代码的可读性和一致性,统一测试辅助函数的实现方式
This commit is contained in:
Kris 2026-07-10 14:29:23 +08:00
parent bb1934023e
commit 713fcdcd8e
24 changed files with 1105 additions and 913 deletions

View File

@ -549,8 +549,9 @@ class TestAgentRunAdapterStream:
service_account_port = MagicMock()
adapter = AgentRunAdapter(execution_port, logger, service_account_port)
# Act
results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))]
# Act — streamAgentRun 通过 _session_scope 获取会话调用 getRunUid
with _patch_pg_session():
results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))]
# Assert
assert len(results) == 2
@ -577,7 +578,8 @@ class TestAgentRunAdapterStream:
adapter = AgentRunAdapter(execution_port, logger, service_account_port)
# Act
results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))]
with _patch_pg_session():
results = [ev async for ev in adapter.streamAgentRun(AgentRunId("run-1"))]
# Assert
assert results[0].payload == {}
@ -594,7 +596,7 @@ class TestAgentRunAdapterStream:
adapter = AgentRunAdapter(execution_port, logger, service_account_port)
# Act / Assert
with pytest.raises(DependencyError):
with _patch_pg_session(), pytest.raises(DependencyError):
async for _ in adapter.streamAgentRun(AgentRunId("run-1")):
pass
logger.error.assert_awaited_once()

View File

@ -125,13 +125,34 @@ def _make_logger() -> MagicMock:
return logger
_active_patchers: list = []
def _build_adapter(db: MagicMock, repos: MagicMock) -> ChannelPersistenceAdapter:
"""构造 adapterpatch ``create_repositories`` 返回桩 repos。"""
with patch(
"""构造 adapterpatch ``create_repositories`` 返回桩 repos。
patch 通过 ``start()`` 启动并在模块级 ``_active_patchers`` 列表中注册
autouse fixture ``_cleanup_patchers`` 在测试方法结束后统一 ``stop()``
这是因为改造后 ``create_repositories`` ``_session_scope(tx)`` 方法调用
时执行而非 ``__init__`` 构造时patch 需要跨越 ``_build_adapter``
返回后继续生效
"""
patcher = patch(
"yuxi.channels.adapters.channel_persistence_adapter.create_repositories",
return_value=repos,
):
return ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=_make_logger())
)
patcher.start()
_active_patchers.append(patcher)
return ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=_make_logger())
@pytest.fixture(autouse=True)
def _cleanup_patchers():
"""每个测试结束后停止所有活跃的 patcher。"""
yield
for p in _active_patchers:
p.stop()
_active_patchers.clear()
@pytest.mark.unit
@ -476,11 +497,10 @@ class TestChannelPersistenceAdapterSession:
"yuxi.channels.adapters.channel_persistence_adapter.create_repositories",
return_value=repos,
):
adapter = ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=logger)
cmd = UpdateChannelSessionCmd(session_id="sess-1", is_temporary=True)
# Act
await adapter.updateChannelSession(cmd)
adapter = ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=logger)
cmd = UpdateChannelSessionCmd(session_id="sess-1", is_temporary=True)
# Act
await adapter.updateChannelSession(cmd)
# Assert: expected_version=None 时应记录 WARN 暴露乐观锁缺口
logger.warn.assert_awaited_once()
@ -500,15 +520,14 @@ class TestChannelPersistenceAdapterSession:
"yuxi.channels.adapters.channel_persistence_adapter.create_repositories",
return_value=repos,
):
adapter = ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=logger)
cmd = UpdateChannelSessionCmd(
session_id="sess-1",
is_temporary=True,
expected_version=1,
)
# Act
await adapter.updateChannelSession(cmd)
adapter = ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=logger)
cmd = UpdateChannelSessionCmd(
session_id="sess-1",
is_temporary=True,
expected_version=1,
)
# Act
await adapter.updateChannelSession(cmd)
# Assert: 传入 expected_version 时不应触发 WARN
logger.warn.assert_not_awaited()
@ -835,21 +854,6 @@ class TestChannelPersistenceAdapterCleanup:
await adapter.cleanupInactiveSessions(["sess-1"])
@pytest.mark.unit
class TestChannelPersistenceAdapterClose:
@pytest.mark.asyncio
async def test_aclose_closes_shared_session(self):
# Arrange
db = _make_db()
adapter = _build_adapter(db, _make_repos())
# Act
await adapter.aclose()
# Assert
db.close.assert_awaited_once()
def _make_pairing_orm(*, version: int = 1, status: str = "pending") -> MagicMock:
"""构造 ChannelPairing ORM 桩。"""
orm = MagicMock()

View File

@ -2,15 +2,15 @@
覆盖 ``ContentReviewRepositoryAdapter`` ``saveReviewResult`` /
``queryReviewHistory`` / ``countReviewHistory`` / ``getReviewDetail`` /
``getReviewAnalytics`` / ``getReviewStats`` / ``updateReviewVerdict`` 方法
使用 ``MagicMock`` 模拟 ``AsyncSession``不连接真实 DB
``getReviewAnalytics`` / ``getReviewStats`` / ``updateReviewVerdict`` /
``deleteOldReviewRecords`` 方法patch
``ChannelContentReviewRecordRepository`` 返回 mock 仓储不连接真实 DB
"""
from __future__ import annotations
from datetime import datetime
from typing import Any
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, patch
from zoneinfo import ZoneInfo
import pytest
@ -44,18 +44,28 @@ pytestmark = pytest.mark.unit
def _make_db() -> MagicMock:
"""构造 AsyncSession 桩,常用方法均为 AsyncMock"""
"""构造 AsyncSession 桩,仅需 close/rollback 供 ``_session_scope`` 使用"""
db = MagicMock()
db.add = MagicMock()
db.flush = AsyncMock()
db.commit = AsyncMock()
db.rollback = AsyncMock()
db.scalar = AsyncMock()
db.execute = AsyncMock()
db.refresh = AsyncMock()
db.close = AsyncMock()
return db
def _make_repo() -> MagicMock:
"""构造 ``ChannelContentReviewRecordRepository`` 桩。"""
repo = MagicMock()
repo.create = AsyncMock()
repo.list = AsyncMock(return_value=[])
repo.count = AsyncMock(return_value=0)
repo.get_by_review_id = AsyncMock(return_value=None)
repo.update_verdict = AsyncMock(return_value=None)
repo.query_analytics = AsyncMock()
repo.query_stats = AsyncMock()
repo.delete_old_records = AsyncMock(return_value=0)
return repo
def _make_outcome() -> ContentModerationOutcome:
"""构造审核结论 DTO。"""
return ContentModerationOutcome(
@ -134,32 +144,6 @@ def _make_analytics_query() -> ContentReviewAnalyticsQuery:
)
def _make_result_with_one(one_value: Any) -> MagicMock:
"""构造 execute 返回值桩,使 ``.one()`` 返回指定聚合行。
analytics / stats 源码用 ``(await db.execute(stmt)).one()`` 访问聚合行
execute 返回的 MagicMock 默认 ``.one()`` 返回新 MagicMock无法读到
预置字段 helper 显式绑定 ``one`` 返回值
"""
result = MagicMock()
result.one.return_value = one_value
return result
def _make_scalar_result(scalar_value: Any) -> MagicMock:
"""构造 execute 返回值桩,使 ``.scalar_one_or_none()`` 返回指定值。
get_by_review_id / update_verdict 源码用
``(await db.execute(stmt)).scalar_one_or_none()`` 访问单行
execute 返回的 MagicMock 默认 ``.scalar_one_or_none()`` 返回新
MagicMock None导致 not-found 分支无法触发 helper 显式
绑定 ``scalar_one_or_none`` 返回值
"""
result = MagicMock()
result.scalar_one_or_none.return_value = scalar_value
return result
def _make_stats_query() -> ContentReviewStatsQuery:
"""构造审核统计查询条件。"""
return ContentReviewStatsQuery(
@ -168,50 +152,80 @@ def _make_stats_query() -> ContentReviewStatsQuery:
)
_active_patchers: list = []
def _build_adapter(
db: MagicMock,
repo: MagicMock,
logger: AsyncMock | None = None,
) -> ContentReviewRepositoryAdapter:
"""构造 adapterpatch ``ChannelContentReviewRecordRepository`` 返回桩 repo。
patch 通过 ``start()`` 启动并在模块级 ``_active_patchers`` 列表中注册
autouse fixture ``_cleanup_patchers`` 在测试方法结束后统一 ``stop()``
这是因为改造后 ``ChannelContentReviewRecordRepository(session)``
``_session_scope(tx)`` 方法调用时执行而非 ``__init__`` 构造时
patch 需要跨越 ``_build_adapter`` 返回后继续生效
"""
patcher = patch(
"yuxi.channels.adapters.content_review_repository_adapter.ChannelContentReviewRecordRepository",
return_value=repo,
)
patcher.start()
_active_patchers.append(patcher)
return ContentReviewRepositoryAdapter(lambda: db, logger or AsyncMock())
@pytest.fixture(autouse=True)
def _cleanup_patchers():
"""每个测试结束后停止所有活跃的 patcher。"""
yield
for p in _active_patchers:
p.stop()
_active_patchers.clear()
@pytest.mark.unit
class TestContentReviewSaveReviewResult:
@pytest.mark.asyncio
async def test_save_review_result_commits_when_no_tx(self):
# Arrange
# repo.create 在 commit=True 时调用 add + commit + refresh不调用 flush
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
adapter = _build_adapter(db, repo)
record = _make_record()
# Act
await adapter.saveReviewResult(record)
# Assert
db.add.assert_called_once()
db.commit.assert_awaited_once()
db.refresh.assert_awaited_once()
db.flush.assert_not_awaited()
# Assert: tx 为 None 时 commit=True自主提交
repo.create.assert_awaited_once()
assert repo.create.call_args.kwargs["commit"] is True
@pytest.mark.asyncio
async def test_save_review_result_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
adapter = _build_adapter(db, repo)
record = _make_record()
tx = MagicMock()
tx.get_session.return_value = db
# Act
await adapter.saveReviewResult(record, tx=tx)
# Assert
db.flush.assert_awaited_once()
db.commit.assert_not_awaited()
# Assert: tx 非空时 commit=False加入应用层事务
assert repo.create.call_args.kwargs["commit"] is False
@pytest.mark.asyncio
async def test_save_review_result_translates_integrity_error_to_conflict(self):
# Arrange
# repo.create 在 commit=True 时调用 commit非 flushIntegrityError 从 commit 抛出
db = _make_db()
db.commit = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.create = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter = _build_adapter(db, repo)
record = _make_record()
# Act / Assert
@ -222,11 +236,10 @@ class TestContentReviewSaveReviewResult:
@pytest.mark.asyncio
async def test_save_review_result_translates_sqlalchemy_error_to_dependency(self):
# Arrange
# repo.create 在 commit=True 时调用 commit非 flushSQLAlchemyError 从 commit 抛出
db = _make_db()
db.commit = AsyncMock(side_effect=SQLAlchemyError("db failure"))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.create = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
record = _make_record()
# Act / Assert
@ -241,11 +254,9 @@ class TestContentReviewQueryHistory:
async def test_query_history_returns_items_tuple(self):
# Arrange
db = _make_db()
result_mock = MagicMock()
result_mock.scalars.return_value.all.return_value = [_make_orm(review_id="rev-1"), _make_orm(review_id="rev-2")]
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.list = AsyncMock(return_value=[_make_orm(review_id="rev-1"), _make_orm(review_id="rev-2")])
adapter = _build_adapter(db, repo)
# Act
items = await adapter.queryReviewHistory(_make_filter(), limit=10, offset=0)
@ -259,11 +270,9 @@ class TestContentReviewQueryHistory:
async def test_query_history_returns_empty_tuple_when_no_records(self):
# Arrange
db = _make_db()
result_mock = MagicMock()
result_mock.scalars.return_value.all.return_value = []
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.list = AsyncMock(return_value=[])
adapter = _build_adapter(db, repo)
# Act
items = await adapter.queryReviewHistory(_make_filter(), limit=10, offset=0)
@ -275,9 +284,9 @@ class TestContentReviewQueryHistory:
async def test_query_history_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.list = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -285,19 +294,14 @@ class TestContentReviewQueryHistory:
@pytest.mark.asyncio
async def test_query_history_translates_invalid_enum_to_dependency(self):
# Arrange
# ORM 持有非法 verdict 值(迁移残留 / 数据损坏),
# Arrange: ORM 持有非法 verdict 值(迁移残留 / 数据损坏),
# ``_enum`` 应翻译 ValueError 为 DependencyError禁止穿透核心层INV-7
# 注:``ChannelType`` 是 ``str`` 子类而非 Enum``ChannelType("unknown")``
# 不抛异常;真正受 ``_enum`` 保护的是 verdict / source / resource_type
db = _make_db()
repo = _make_repo()
orm = _make_orm(review_id="rev-1")
orm.verdict = "unknown_verdict"
result_mock = MagicMock()
result_mock.scalars.return_value.all.return_value = [orm]
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo.list = AsyncMock(return_value=[orm])
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -310,11 +314,9 @@ class TestContentReviewCountHistory:
async def test_count_history_returns_int(self):
# Arrange
db = _make_db()
result_mock = MagicMock()
result_mock.scalar.return_value = 42
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.count = AsyncMock(return_value=42)
adapter = _build_adapter(db, repo)
# Act
count = await adapter.countReviewHistory(_make_filter())
@ -326,11 +328,9 @@ class TestContentReviewCountHistory:
async def test_count_history_returns_zero_when_no_records(self):
# Arrange
db = _make_db()
result_mock = MagicMock()
result_mock.scalar.return_value = None
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.count = AsyncMock(return_value=0)
adapter = _build_adapter(db, repo)
# Act
count = await adapter.countReviewHistory(_make_filter())
@ -342,9 +342,9 @@ class TestContentReviewCountHistory:
async def test_count_history_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.count = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -356,13 +356,10 @@ class TestContentReviewGetDetail:
@pytest.mark.asyncio
async def test_get_detail_returns_detail_when_found(self):
# Arrange
# repo.get_by_review_id 用 ``(await db.execute(stmt)).scalar_one_or_none()``
# 访问单行,故桩 db.execute 返回带 scalar_one_or_none 的结果对象
db = _make_db()
orm = _make_orm(review_id="rev-1")
db.execute = AsyncMock(return_value=_make_scalar_result(orm))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.get_by_review_id = AsyncMock(return_value=_make_orm(review_id="rev-1"))
adapter = _build_adapter(db, repo)
# Act
detail = await adapter.getReviewDetail("rev-1")
@ -375,9 +372,9 @@ class TestContentReviewGetDetail:
async def test_get_detail_returns_none_when_not_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(None))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.get_by_review_id = AsyncMock(return_value=None)
adapter = _build_adapter(db, repo)
# Act
detail = await adapter.getReviewDetail("missing")
@ -389,9 +386,9 @@ class TestContentReviewGetDetail:
async def test_get_detail_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.get_by_review_id = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -399,15 +396,14 @@ class TestContentReviewGetDetail:
@pytest.mark.asyncio
async def test_get_detail_translates_invalid_verdict_to_dependency(self):
# Arrange
# ORM 持有非法 verdict 值(数据损坏),``_enum`` 应翻译 ValueError
# Arrange: ORM 持有非法 verdict 值(数据损坏),``_enum`` 应翻译 ValueError
# 为 DependencyError禁止穿透核心层INV-7
db = _make_db()
repo = _make_repo()
orm = _make_orm(review_id="rev-1")
orm.verdict = "unknown_verdict"
db.execute = AsyncMock(return_value=_make_scalar_result(orm))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo.get_by_review_id = AsyncMock(return_value=orm)
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -418,31 +414,22 @@ class TestContentReviewGetDetail:
class TestContentReviewGetAnalytics:
@pytest.mark.asyncio
async def test_get_analytics_returns_aggregated_result(self):
# Arrange
# 源码用 ``(await db.execute(agg_stmt)).one()`` 访问聚合行字段,
# 故 execute 须返回带 ``one()`` 方法的桩,``one()`` 返回预置聚合行
# Arrange: repo.query_analytics 返回聚合 dict适配器映射为 DTO
db = _make_db()
agg_row = MagicMock()
agg_row.total_reviews = 10
agg_row.block_count = 4
category_result = MagicMock()
category_result.all.return_value = [
MagicMock(category="politics", count=3),
MagicMock(category="violence", count=1),
]
trend_result = MagicMock()
trend_result.all.return_value = [
(datetime(2026, 1, 1), 5, 2),
]
db.execute = AsyncMock(
side_effect=[
_make_result_with_one(agg_row),
category_result,
trend_result,
]
repo = _make_repo()
repo.query_analytics = AsyncMock(
return_value={
"total_reviews": 10,
"block_count": 4,
"block_rate": 0.4,
"by_category": [
{"category": "politics", "count": 3},
{"category": "violence", "count": 1},
],
"trend": [(datetime(2026, 1, 1), 5, 2)],
}
)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
adapter = _build_adapter(db, repo)
# Act
result = await adapter.getReviewAnalytics(_make_analytics_query())
@ -458,22 +445,17 @@ class TestContentReviewGetAnalytics:
async def test_get_analytics_block_rate_zero_when_no_reviews(self):
# Arrange
db = _make_db()
agg_row = MagicMock()
agg_row.total_reviews = 0
agg_row.block_count = 0
category_result = MagicMock()
category_result.all.return_value = []
trend_result = MagicMock()
trend_result.all.return_value = []
db.execute = AsyncMock(
side_effect=[
_make_result_with_one(agg_row),
category_result,
trend_result,
]
repo = _make_repo()
repo.query_analytics = AsyncMock(
return_value={
"total_reviews": 0,
"block_count": 0,
"block_rate": 0.0,
"by_category": [],
"trend": [],
}
)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
adapter = _build_adapter(db, repo)
# Act
result = await adapter.getReviewAnalytics(_make_analytics_query())
@ -486,9 +468,9 @@ class TestContentReviewGetAnalytics:
async def test_get_analytics_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.query_analytics = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -499,30 +481,24 @@ class TestContentReviewGetAnalytics:
class TestContentReviewGetStats:
@pytest.mark.asyncio
async def test_get_stats_returns_aggregated_result(self):
# Arrange
# 源码用 ``(await db.execute(agg_stmt)).one()`` 访问聚合行字段,
# 故 execute 须返回带 ``one()`` 方法的桩,``one()`` 返回预置聚合行
# Arrange: repo.query_stats 返回聚合 dict适配器映射为 DTO
db = _make_db()
agg_row = MagicMock()
agg_row.total_reviews = 20
agg_row.pass_count = 12
agg_row.review_count = 3
agg_row.block_count = 5
agg_row.manual_count = 2
agg_row.avg_decision_seconds = 42.5
category_result = MagicMock()
category_result.all.return_value = [MagicMock(category="spam", count=4)]
trend_result = MagicMock()
trend_result.all.return_value = [(datetime(2026, 1, 1), 12, 5)]
db.execute = AsyncMock(
side_effect=[
_make_result_with_one(agg_row),
category_result,
trend_result,
]
repo = _make_repo()
repo.query_stats = AsyncMock(
return_value={
"total_reviews": 20,
"pass_count": 12,
"review_count": 3,
"block_count": 5,
"pass_rate": 0.6,
"block_rate": 0.25,
"manual_intervention_rate": 0.1,
"avg_decision_seconds": 42.5,
"by_category": [{"category": "spam", "count": 4}],
"trend": [{"timestamp": datetime(2026, 1, 1), "pass_count": 12, "block_count": 5}],
}
)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
adapter = _build_adapter(db, repo)
# Act
result = await adapter.getReviewStats(_make_stats_query())
@ -540,26 +516,22 @@ class TestContentReviewGetStats:
async def test_get_stats_rates_zero_when_no_reviews(self):
# Arrange
db = _make_db()
agg_row = MagicMock()
agg_row.total_reviews = 0
agg_row.pass_count = 0
agg_row.review_count = 0
agg_row.block_count = 0
agg_row.manual_count = 0
agg_row.avg_decision_seconds = None
category_result = MagicMock()
category_result.all.return_value = []
trend_result = MagicMock()
trend_result.all.return_value = []
db.execute = AsyncMock(
side_effect=[
_make_result_with_one(agg_row),
category_result,
trend_result,
]
repo = _make_repo()
repo.query_stats = AsyncMock(
return_value={
"total_reviews": 0,
"pass_count": 0,
"review_count": 0,
"block_count": 0,
"pass_rate": 0.0,
"block_rate": 0.0,
"manual_intervention_rate": 0.0,
"avg_decision_seconds": None,
"by_category": [],
"trend": [],
}
)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
adapter = _build_adapter(db, repo)
# Act
result = await adapter.getReviewStats(_make_stats_query())
@ -573,9 +545,9 @@ class TestContentReviewGetStats:
async def test_get_stats_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.query_stats = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -588,69 +560,56 @@ class TestContentReviewUpdateVerdict:
async def test_update_verdict_returns_detail_when_found(self):
# Arrange
db = _make_db()
orm = _make_orm(review_id="rev-1")
result_mock = MagicMock()
result_mock.scalar_one_or_none.return_value = orm
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.update_verdict = AsyncMock(return_value=_make_orm(review_id="rev-1"))
adapter = _build_adapter(db, repo)
# Act
detail = await adapter.updateReviewVerdict("rev-1", "pass", "admin-1")
# Assert
# repo.update_verdict 在 commit=True 时调用 commit + refresh不调用 flush
assert detail is not None
assert detail.review_id == "rev-1"
assert orm.verdict == "pass"
assert orm.reviewer == "admin-1"
db.commit.assert_awaited_once()
db.refresh.assert_awaited_once()
db.flush.assert_not_awaited()
repo.update_verdict.assert_awaited_once()
assert repo.update_verdict.call_args.kwargs["commit"] is True
@pytest.mark.asyncio
async def test_update_verdict_returns_none_when_not_found(self):
# Arrange
db = _make_db()
result_mock = MagicMock()
result_mock.scalar_one_or_none.return_value = None
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.update_verdict = AsyncMock(return_value=None)
adapter = _build_adapter(db, repo)
# Act
detail = await adapter.updateReviewVerdict("missing", "pass", "admin-1")
# Assert
assert detail is None
db.rollback.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_verdict_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
orm = _make_orm(review_id="rev-1")
result_mock = MagicMock()
result_mock.scalar_one_or_none.return_value = orm
db.execute = AsyncMock(return_value=result_mock)
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.update_verdict = AsyncMock(return_value=_make_orm(review_id="rev-1"))
adapter = _build_adapter(db, repo)
tx = MagicMock()
tx.get_session.return_value = db
# Act
await adapter.updateReviewVerdict("rev-1", "pass", "admin-1", tx=tx)
# Assert
db.flush.assert_awaited_once()
db.commit.assert_not_awaited()
assert repo.update_verdict.call_args.kwargs["commit"] is False
@pytest.mark.asyncio
async def test_update_verdict_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
repo = _make_repo()
repo.update_verdict = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
# Act / Assert
with pytest.raises(DependencyError):
@ -666,66 +625,63 @@ class TestContentReviewRepositoryDeleteOldRecords:
与各类异常翻译IntegrityError / SQLAlchemyError / ConflictError /
DependencyError / 通用 Exception
使用 AsyncMock 替换 ``adapter._repo.delete_old_records``不连接真实
DB聚焦适配器的参数归一化委托与错误翻译职责
通过 ``_build_adapter`` patch ``ChannelContentReviewRecordRepository``
返回 mock 仓储不连接真实 DB聚焦适配器的参数归一化委托与
错误翻译职责
"""
@pytest.mark.asyncio
async def test_delete_old_records_returns_count_on_success(self):
# 正常路径repo.delete_old_records 返回删除数量,适配器原样透传
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(return_value=42)
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(return_value=42)
adapter = _build_adapter(db, repo)
result = await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0))
assert result == 42
delete_mock.assert_awaited_once()
repo.delete_old_records.assert_awaited_once()
db.rollback.assert_not_awaited()
@pytest.mark.asyncio
async def test_delete_old_records_uses_default_limit_1000(self):
# limit 未传时默认 1000应原样传递给 repo
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(return_value=0)
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(return_value=0)
adapter = _build_adapter(db, repo)
before = datetime(2026, 1, 1, 0, 0, 0)
await adapter.deleteOldReviewRecords(before)
delete_mock.assert_awaited_once_with(before, limit=1000, commit=True)
repo.delete_old_records.assert_awaited_once_with(before, limit=1000, commit=True)
@pytest.mark.asyncio
async def test_delete_old_records_passes_custom_limit(self):
# 自定义 limit 应透传给 repo
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(return_value=0)
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(return_value=0)
adapter = _build_adapter(db, repo)
before = datetime(2026, 1, 1, 0, 0, 0)
await adapter.deleteOldReviewRecords(before, limit=500)
delete_mock.assert_awaited_once_with(before, limit=500, commit=True)
repo.delete_old_records.assert_awaited_once_with(before, limit=500, commit=True)
@pytest.mark.asyncio
async def test_delete_old_records_falls_back_when_before_is_none(self):
# before 为 None 时_to_naive_utc(None) or None = None
# 兜底传递 None 给 repo适配器不因 None 崩溃)
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(return_value=0)
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(return_value=0)
adapter = _build_adapter(db, repo)
await adapter.deleteOldReviewRecords(None) # type: ignore[arg-type]
delete_mock.assert_awaited_once_with(None, limit=1000, commit=True)
repo.delete_old_records.assert_awaited_once_with(None, limit=1000, commit=True)
@pytest.mark.asyncio
async def test_delete_old_records_normalizes_aware_before_to_naive_utc(self):
@ -733,26 +689,22 @@ class TestContentReviewRepositoryDeleteOldRecords:
# _to_naive_utc 应归一化为 naive UTC 后再传递给 repo
# Shanghai 2026-01-01 08:00:00 (+08:00) → UTC 2026-01-01 00:00:00 (naive)
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(return_value=0)
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(return_value=0)
adapter = _build_adapter(db, repo)
aware = datetime(2026, 1, 1, 8, 0, 0, tzinfo=ZoneInfo("Asia/Shanghai"))
await adapter.deleteOldReviewRecords(aware)
delete_mock.assert_awaited_once_with(
datetime(2026, 1, 1, 0, 0, 0), limit=1000, commit=True
)
repo.delete_old_records.assert_awaited_once_with(datetime(2026, 1, 1, 0, 0, 0), limit=1000, commit=True)
@pytest.mark.asyncio
async def test_delete_old_records_translates_integrity_error_to_conflict(self):
# IntegrityError → ConflictError并回滚事务
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter = _build_adapter(db, repo)
with pytest.raises(ConflictError):
await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0))
@ -762,10 +714,9 @@ class TestContentReviewRepositoryDeleteOldRecords:
async def test_delete_old_records_translates_sqlalchemy_error_to_dependency(self):
# SQLAlchemyError → DependencyError并回滚事务
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter = _build_adapter(db, repo)
with pytest.raises(DependencyError):
await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0))
@ -775,10 +726,9 @@ class TestContentReviewRepositoryDeleteOldRecords:
async def test_delete_old_records_reraises_conflict_error(self):
# 契约层 ConflictError 原样重抛,不二次翻译,并回滚事务
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(side_effect=ConflictError("content_review_record"))
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(side_effect=ConflictError("content_review_record"))
adapter = _build_adapter(db, repo)
with pytest.raises(ConflictError):
await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0))
@ -788,14 +738,11 @@ class TestContentReviewRepositoryDeleteOldRecords:
async def test_delete_old_records_reraises_dependency_error(self):
# 契约层 DependencyError 原样重抛,不二次翻译,并回滚事务
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(
side_effect=DependencyError(
"content_review_repository", Error("upstream failure")
)
repo = _make_repo()
repo.delete_old_records = AsyncMock(
side_effect=DependencyError("content_review_repository", Error("upstream failure"))
)
adapter._repo.delete_old_records = delete_mock
adapter = _build_adapter(db, repo)
with pytest.raises(DependencyError):
await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0))
@ -806,9 +753,9 @@ class TestContentReviewRepositoryDeleteOldRecords:
# 通用 Exception → DependencyError 包装,记录日志并回滚事务
db = _make_db()
logger = AsyncMock()
adapter = ContentReviewRepositoryAdapter(db, logger)
delete_mock = AsyncMock(side_effect=RuntimeError("unexpected boom"))
adapter._repo.delete_old_records = delete_mock
repo = _make_repo()
repo.delete_old_records = AsyncMock(side_effect=RuntimeError("unexpected boom"))
adapter = _build_adapter(db, repo, logger=logger)
with pytest.raises(DependencyError):
await adapter.deleteOldReviewRecords(datetime(2026, 1, 1, 0, 0, 0))

View File

@ -9,7 +9,7 @@
from __future__ import annotations
from datetime import date, datetime
from datetime import UTC, date, datetime
from unittest.mock import MagicMock, patch
import pytest
@ -314,6 +314,22 @@ class TestOrmToAuditLog:
assert result.target_channel == "feishu"
@pytest.mark.unit
class TestToAwareUtc:
def test_converts_naive_datetime_to_aware_utc(self):
naive = datetime(2026, 1, 1, 12, 0, 0)
result = mappers._to_aware_utc(naive)
assert result == datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC)
assert result.tzinfo is UTC
def test_returns_aware_datetime_unchanged(self):
aware = datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC)
assert mappers._to_aware_utc(aware) is aware
def test_returns_none_unchanged(self):
assert mappers._to_aware_utc(None) is None
@pytest.mark.unit
class TestOrmToOutboxEntry:
def test_returns_outbox_entry_with_str_ids(self):
@ -332,6 +348,25 @@ class TestOrmToOutboxEntry:
assert result.status == OutboxStatus.PENDING
assert result.durability_policy == MessageDurabilityPolicy.REQUIRED
def test_datetime_fields_are_aware_utc(self):
"""DB 读出的 naive datetime 必须转为 aware UTC避免领域层时区混用。"""
orm = _make_outbox_orm()
orm.created_at = datetime(2026, 1, 1)
orm.updated_at = datetime(2026, 1, 2)
orm.expires_at = datetime(2026, 1, 3)
orm.sent_at = datetime(2026, 1, 4)
orm.last_retry_at = datetime(2026, 1, 5)
orm.next_retry_at = datetime(2026, 1, 6)
result = mappers.orm_to_outbox_entry(orm, _make_account_orm())
assert result.created_at.tzinfo is UTC
assert result.updated_at.tzinfo is UTC
assert result.expires_at.tzinfo is UTC
assert result.sent_at.tzinfo is UTC
assert result.last_retry_at.tzinfo is UTC
assert result.next_retry_at.tzinfo is UTC
def test_handles_none_channel_session_id(self):
# Arrange
orm = _make_outbox_orm()
@ -343,6 +378,21 @@ class TestOrmToOutboxEntry:
# Assert
assert result.channel_session_id is None
def test_handles_none_datetime_fields(self):
"""可空 datetime 字段为 None 时不应报错。"""
orm = _make_outbox_orm()
orm.expires_at = None
orm.sent_at = None
orm.last_retry_at = None
orm.next_retry_at = None
result = mappers.orm_to_outbox_entry(orm, _make_account_orm())
assert result.expires_at is None
assert result.sent_at is None
assert result.last_retry_at is None
assert result.next_retry_at is None
@pytest.mark.unit
class TestOrmToUserIdentity:

View File

@ -90,13 +90,34 @@ def _make_binding_orm(
return orm
_active_patchers: list = []
def _build_adapter(db: MagicMock, repos: MagicMock) -> ChannelPersistenceAdapter:
"""构造 adapterpatch ``create_repositories`` 返回桩 repos。"""
with patch(
"""构造 adapterpatch ``create_repositories`` 返回桩 repos。
patch 通过 ``start()`` 启动并在模块级 ``_active_patchers`` 列表中注册
autouse fixture ``_cleanup_patchers`` 在测试方法结束后统一 ``stop()``
这是因为改造后 ``create_repositories`` ``_session_scope(tx)`` 方法调用
时执行而非 ``__init__`` 构造时patch 需要跨越 ``_build_adapter``
返回后继续生效
"""
patcher = patch(
"yuxi.channels.adapters.channel_persistence_adapter.create_repositories",
return_value=repos,
):
return ChannelPersistenceAdapter(db, OutboxConfig.default(), logger=MagicMock())
)
patcher.start()
_active_patchers.append(patcher)
return ChannelPersistenceAdapter(lambda: db, OutboxConfig.default(), logger=MagicMock())
@pytest.fixture(autouse=True)
def _cleanup_patchers():
"""每个测试结束后停止所有活跃的 patcher。"""
yield
for p in _active_patchers:
p.stop()
_active_patchers.clear()
def _make_operator() -> Operator:

View File

@ -3,14 +3,13 @@
覆盖 ``yuxi.channels.application.extension.scheduler_handlers.channel_audit_log_retention_handler``
- execute 正常路径无记录 / 有记录删除
- payload 覆盖 retention_days batch_size 默认值
- session / repository 异常转换为 TaskResult(success=False)
- repository 异常转换为 TaskResult(success=False)
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime, timedelta
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
from yuxi.channels.application.extension.scheduler_handlers.channel_audit_log_retention_handler import (
@ -38,15 +37,8 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
def _make_handler(*, audit_log_repo=None, logger=None, cache_port=None):
"""构造 ChannelAuditLogRetentionHandler 及其依赖桩。"""
db = MagicMock()
if audit_log_repo is None:
audit_log_repo = AsyncMock()
audit_log_repo.deleteOldAuditLogs.return_value = 0
@ -58,13 +50,11 @@ def _make_handler(*, audit_log_repo=None, logger=None, cache_port=None):
cache_port.releaseAdvisoryLock.return_value = True
handler = ChannelAuditLogRetentionHandler(
session_factory=lambda: _fake_session_factory(db),
audit_log_repo_factory=lambda _db: audit_log_repo,
audit_log_repo=audit_log_repo,
cache_port=cache_port,
logger=logger,
)
return handler, {
"db": db,
"audit_log_repo": audit_log_repo,
"cache_port": cache_port,
"logger": logger,
@ -131,25 +121,6 @@ class TestHandlerExecute:
@pytest.mark.unit
class TestHandlerErrors:
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelAuditLogRetentionHandler(
session_factory=_broken_session_factory,
audit_log_repo_factory=lambda _db: AsyncMock(),
cache_port=fake_cache,
logger=fake_logger,
)
result = await handler.execute(_make_ctx())
assert result.success is False
assert "session creation failed" in (result.error or "")
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger):
handler, deps = _make_handler(logger=fake_logger)

View File

@ -3,14 +3,13 @@
覆盖 ``yuxi.channels.application.extension.scheduler_handlers.channel_content_review_retention_handler``
- execute 正常路径无记录 / 有记录删除
- payload 覆盖 retention_days batch_size 默认值
- session / repository 异常转换为 TaskResult(success=False)
- repository 异常转换为 TaskResult(success=False)
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime, timedelta
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
from yuxi.channels.application.extension.scheduler_handlers.channel_content_review_retention_handler import (
@ -38,15 +37,8 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
def _make_handler(*, content_review_repo=None, logger=None, cache_port=None):
"""构造 ChannelContentReviewRetentionHandler 及其依赖桩。"""
db = MagicMock()
if content_review_repo is None:
content_review_repo = AsyncMock()
content_review_repo.deleteOldReviewRecords.return_value = 0
@ -58,13 +50,11 @@ def _make_handler(*, content_review_repo=None, logger=None, cache_port=None):
cache_port.releaseAdvisoryLock.return_value = True
handler = ChannelContentReviewRetentionHandler(
session_factory=lambda: _fake_session_factory(db),
content_review_repo_factory=lambda _db: content_review_repo,
content_review_repo=content_review_repo,
cache_port=cache_port,
logger=logger,
)
return handler, {
"db": db,
"content_review_repo": content_review_repo,
"cache_port": cache_port,
"logger": logger,
@ -132,25 +122,6 @@ class TestHandlerExecute:
@pytest.mark.unit
class TestHandlerErrors:
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelContentReviewRetentionHandler(
session_factory=_broken_session_factory,
content_review_repo_factory=lambda _db: AsyncMock(),
cache_port=fake_cache,
logger=fake_logger,
)
result = await handler.execute(_make_ctx())
assert result.success is False
assert "session creation failed" in (result.error or "")
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger):
handler, deps = _make_handler(logger=fake_logger)

View File

@ -3,14 +3,13 @@
覆盖 ``yuxi.channels.application.extension.scheduler_handlers.channel_idempotency_cleanup_handler``
- execute 正常路径无记录 / 有记录删除
- deleteExpiredRecords ``utc_now_naive`` 为截止时间
- session / repository 异常转换为 TaskResult(success=False)
- repository 异常转换为 TaskResult(success=False)
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
from yuxi.channels.application.extension.scheduler_handlers.channel_idempotency_cleanup_handler import (
@ -38,15 +37,8 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
def _make_handler(*, idempotency_repo=None, logger=None, cache_port=None):
"""构造 ChannelIdempotencyCleanupHandler 及其依赖桩。"""
db = MagicMock()
if idempotency_repo is None:
idempotency_repo = AsyncMock()
idempotency_repo.deleteExpiredRecords.return_value = 0
@ -58,13 +50,11 @@ def _make_handler(*, idempotency_repo=None, logger=None, cache_port=None):
cache_port.releaseAdvisoryLock.return_value = True
handler = ChannelIdempotencyCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
idempotency_repo_factory=lambda _db: idempotency_repo,
idempotency_repo=idempotency_repo,
cache_port=cache_port,
logger=logger,
)
return handler, {
"db": db,
"idempotency_repo": idempotency_repo,
"cache_port": cache_port,
"logger": logger,
@ -134,26 +124,6 @@ class TestHandlerExecute:
@pytest.mark.unit
class TestHandlerErrors:
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelIdempotencyCleanupHandler(
session_factory=_broken_session_factory,
idempotency_repo_factory=lambda _db: AsyncMock(),
cache_port=fake_cache,
logger=fake_logger,
)
result = await handler.execute(_make_ctx())
assert result.success is False
assert "session creation failed" in (result.error or "")
fake_logger.exception.assert_awaited_once()
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger):
handler, deps = _make_handler(logger=fake_logger)

View File

@ -6,14 +6,13 @@
- 死信审计日志写入及写入失败不阻塞
- ARQ 入队失败仅记录日志
- 单条异常隔离
- session / repository 异常转换为 TaskResult(success=False)
- repository 异常转换为 TaskResult(success=False)
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime, timedelta
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
from yuxi.channels.application.extension.scheduler_handlers.channel_outbox_recovery_handler import (
@ -49,12 +48,6 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
def _make_outbox_config(
ttl_seconds: int = 86400,
max_retry: int = 5,
@ -110,7 +103,6 @@ def _make_handler(
cache_port=None,
):
"""构造 ChannelOutboxRecoveryHandler 及其依赖桩。"""
db = MagicMock()
if outbox_repo is None:
outbox_repo = AsyncMock()
outbox_repo.listPendingOutboxEntries.return_value = []
@ -130,16 +122,14 @@ def _make_handler(
cache_port.releaseAdvisoryLock.return_value = True
handler = ChannelOutboxRecoveryHandler(
session_factory=lambda: _fake_session_factory(db),
outbox_repo_factory=lambda _db: outbox_repo,
audit_log_repo_factory=lambda _db: audit_log_repo,
outbox_repo=outbox_repo,
audit_log_repo=audit_log_repo,
queue_port=queue_port,
outbox_config=outbox_config,
cache_port=cache_port,
logger=logger,
)
return handler, {
"db": db,
"outbox_repo": outbox_repo,
"audit_log_repo": audit_log_repo,
"queue_port": queue_port,
@ -377,28 +367,6 @@ class TestHandlerErrors:
assert result.success is True
assert result.output["processed_count"] == 1
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelOutboxRecoveryHandler(
session_factory=_broken_session_factory,
outbox_repo_factory=lambda _db: AsyncMock(),
audit_log_repo_factory=lambda _db: AsyncMock(),
queue_port=AsyncMock(),
outbox_config=_make_outbox_config(),
cache_port=fake_cache,
logger=fake_logger,
)
result = await handler.execute(_make_ctx())
assert result.success is False
assert "session creation failed" in (result.error or "")
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger):
handler, deps = _make_handler(logger=fake_logger)

View File

@ -6,14 +6,13 @@
- payload 覆盖默认保留小时数与批量大小
- event_publisher None 时跳过事件发布
- event_publisher 发布异常仅记录 WARN不影响主流程
- 异常隔离仓储异常 / session_factory 异常转换为 TaskResult(success=False)
- 异常隔离仓储异常转换为 TaskResult(success=False)
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
from yuxi.channels.application.extension.scheduler_handlers.channel_outbox_terminal_cleanup_handler import (
@ -40,12 +39,6 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
# ─── 类属性 ───────────────────────────────────────────────────────────────
@ -67,15 +60,13 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_successful_cleanup_returns_success(self, fake_logger, fake_cache):
# Arrange
db = MagicMock()
outbox_repo = AsyncMock()
outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1", "dead-2"])
outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=["sent-1"])
event_publisher = AsyncMock()
handler = ChannelOutboxTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
outbox_repo_factory=lambda _db: outbox_repo,
outbox_repo=outbox_repo,
event_publisher=event_publisher,
cache_port=fake_cache,
logger=fake_logger,
@ -92,14 +83,12 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_payload_overrides_defaults(self, fake_logger, fake_cache):
# Arrange - payload 指定 dead_retention_hours=24, sent_retention_hours=168, batch_size=50
db = MagicMock()
outbox_repo = AsyncMock()
outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=[])
outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=[])
handler = ChannelOutboxTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
outbox_repo_factory=lambda _db: outbox_repo,
outbox_repo=outbox_repo,
event_publisher=None,
cache_port=fake_cache,
logger=fake_logger,
@ -123,14 +112,12 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_event_publisher_none_skips_publish(self, fake_logger, fake_cache):
# Arrange - event_publisher 为 None
db = MagicMock()
outbox_repo = AsyncMock()
outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1"])
outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=[])
handler = ChannelOutboxTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
outbox_repo_factory=lambda _db: outbox_repo,
outbox_repo=outbox_repo,
event_publisher=None,
cache_port=fake_cache,
logger=fake_logger,
@ -146,7 +133,6 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_event_publish_failure_continues(self, fake_logger, fake_cache):
# Arrange - event_publisher 发布异常仅记录 WARN不影响主流程
db = MagicMock()
outbox_repo = AsyncMock()
outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1"])
outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=["sent-1"])
@ -154,8 +140,7 @@ class TestHandlerExecute:
event_publisher.publishOutboxEntryPurged = AsyncMock(side_effect=RuntimeError("publish failed"))
handler = ChannelOutboxTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
outbox_repo_factory=lambda _db: outbox_repo,
outbox_repo=outbox_repo,
event_publisher=event_publisher,
cache_port=fake_cache,
logger=fake_logger,
@ -173,15 +158,13 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_event_published_with_purged_entry_ids(self, fake_logger, fake_cache):
# Arrange - 验证发布的事件包含被清理条目的完整 ID 列表
db = MagicMock()
outbox_repo = AsyncMock()
outbox_repo.cleanupOldDeadEntries = AsyncMock(return_value=["dead-1"])
outbox_repo.cleanupOldSentEntries = AsyncMock(return_value=["sent-1"])
event_publisher = AsyncMock()
handler = ChannelOutboxTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
outbox_repo_factory=lambda _db: outbox_repo,
outbox_repo=outbox_repo,
event_publisher=event_publisher,
cache_port=fake_cache,
logger=fake_logger,
@ -206,13 +189,11 @@ class TestHandlerErrors:
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger, fake_cache):
# Arrange - 仓储抛异常
db = MagicMock()
outbox_repo = AsyncMock()
outbox_repo.cleanupOldDeadEntries = AsyncMock(side_effect=RuntimeError("db error"))
handler = ChannelOutboxTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
outbox_repo_factory=lambda _db: outbox_repo,
outbox_repo=outbox_repo,
event_publisher=None,
cache_port=fake_cache,
logger=fake_logger,
@ -224,27 +205,3 @@ class TestHandlerErrors:
# Assert
assert result.success is False
assert "db error" in (result.error or "")
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
# Arrange - session_factory 自身抛异常
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelOutboxTerminalCleanupHandler(
session_factory=_broken_session_factory,
outbox_repo_factory=lambda _db: AsyncMock(),
event_publisher=None,
cache_port=fake_cache,
logger=fake_logger,
)
# Act
result = await handler.execute(_make_ctx())
# Assert
assert result.success is False
assert "session creation failed" in (result.error or "")

View File

@ -4,15 +4,14 @@
- execute 正常路径空扫描 / 过期标记 / 未过期跳过 / 混合记录
- payload 覆盖 batch_size 默认值
- 单条记录异常隔离
- session / repository 异常转换为 TaskResult(success=False)
- repository 异常转换为 TaskResult(success=False)
- fromRecord 使用记录自身 version
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime, timedelta
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
from yuxi.channels.application.extension.scheduler_handlers.channel_pairing_expiration_handler import (
@ -43,12 +42,6 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
def _make_pairing_record(
*,
pairing_id: str = "pr-001",
@ -73,7 +66,6 @@ def _make_pairing_record(
def _make_handler(*, pairing_repo=None, logger=None, cache_port=None):
"""构造 ChannelPairingExpirationHandler 及其依赖桩。"""
db = MagicMock()
if pairing_repo is None:
pairing_repo = AsyncMock()
pairing_repo.listExpiredPendingPairings.return_value = ()
@ -85,13 +77,11 @@ def _make_handler(*, pairing_repo=None, logger=None, cache_port=None):
cache_port.releaseAdvisoryLock.return_value = True
handler = ChannelPairingExpirationHandler(
session_factory=lambda: _fake_session_factory(db),
pairing_repo_factory=lambda _db: pairing_repo,
pairing_repo=pairing_repo,
cache_port=cache_port,
logger=logger,
)
return handler, {
"db": db,
"pairing_repo": pairing_repo,
"cache_port": cache_port,
"logger": logger,
@ -222,25 +212,6 @@ class TestHandlerErrors:
assert result.output["scanned_count"] == 2
deps["logger"].exception.assert_awaited_once()
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelPairingExpirationHandler(
session_factory=_broken_session_factory,
pairing_repo_factory=lambda _db: AsyncMock(),
cache_port=fake_cache,
logger=fake_logger,
)
result = await handler.execute(_make_ctx())
assert result.success is False
assert "session creation failed" in (result.error or "")
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger):
handler, deps = _make_handler(logger=fake_logger)

View File

@ -4,14 +4,13 @@
模块的 ``ChannelPairingTerminalCleanupHandler``
- execute 正常路径调用 cleanupOldTerminalPairings 并返回成功结果
- payload 覆盖默认保留天数与批量大小
- 异常隔离仓储异常 / session_factory 异常转换为 TaskResult(success=False)
- 异常隔离仓储异常转换为 TaskResult(success=False)
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
from yuxi.channels.application.extension.scheduler_handlers.channel_pairing_terminal_cleanup_handler import (
@ -37,12 +36,6 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
# ─── 类属性 ───────────────────────────────────────────────────────────────
@ -64,13 +57,11 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_successful_cleanup_returns_success(self, fake_logger, fake_cache):
# Arrange
db = MagicMock()
pairing_repo = AsyncMock()
pairing_repo.cleanupOldTerminalPairings = AsyncMock(return_value=42)
handler = ChannelPairingTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
pairing_repo_factory=lambda _db: pairing_repo,
pairing_repo=pairing_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -86,13 +77,11 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_payload_overrides_defaults(self, fake_logger, fake_cache):
# Arrange - payload 指定 retention_days=7, batch_size=10
db = MagicMock()
pairing_repo = AsyncMock()
pairing_repo.cleanupOldTerminalPairings = AsyncMock(return_value=5)
handler = ChannelPairingTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
pairing_repo_factory=lambda _db: pairing_repo,
pairing_repo=pairing_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -109,13 +98,11 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_default_retention_days_used_when_payload_missing(self, fake_logger, fake_cache):
# Arrange - payload 不含 retention_days 与 batch_size
db = MagicMock()
pairing_repo = AsyncMock()
pairing_repo.cleanupOldTerminalPairings = AsyncMock(return_value=0)
handler = ChannelPairingTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
pairing_repo_factory=lambda _db: pairing_repo,
pairing_repo=pairing_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -137,13 +124,11 @@ class TestHandlerErrors:
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger, fake_cache):
# Arrange - 仓储抛异常
db = MagicMock()
pairing_repo = AsyncMock()
pairing_repo.cleanupOldTerminalPairings = AsyncMock(side_effect=RuntimeError("db error"))
handler = ChannelPairingTerminalCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
pairing_repo_factory=lambda _db: pairing_repo,
pairing_repo=pairing_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -154,26 +139,3 @@ class TestHandlerErrors:
# Assert - 异常被捕获并转换为 TaskResult(success=False)
assert result.success is False
assert "db error" in (result.error or "")
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
# Arrange - session_factory 自身抛异常
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelPairingTerminalCleanupHandler(
session_factory=_broken_session_factory,
pairing_repo_factory=lambda _db: AsyncMock(),
cache_port=fake_cache,
logger=fake_logger,
)
# Act
result = await handler.execute(_make_ctx())
# Assert
assert result.success is False
assert "session creation failed" in (result.error or "")

View File

@ -7,12 +7,11 @@
- payload 覆盖默认阈值与批量大小
- realtime_metrics 非空时调用 releaseSession
- realtime_metrics None 时跳过释放
- 异常隔离仓储异常 / session_factory 异常转换为 TaskResult(success=False)
- 异常隔离仓储异常转换为 TaskResult(success=False)
"""
from __future__ import annotations
from contextlib import asynccontextmanager
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
@ -40,12 +39,6 @@ def _make_ctx(**payload_overrides) -> TaskContext:
)
@asynccontextmanager
async def _fake_session_factory(db_mock):
"""构造假的 session_factoryyield db_mock。"""
yield db_mock
def _make_session(session_id: str = "sess-001"):
"""构造带 session_id 属性的假会话对象。"""
session = MagicMock()
@ -74,14 +67,12 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_successful_cleanup_returns_success(self, fake_logger, fake_cache):
# Arrange
db = MagicMock()
session_repo = AsyncMock()
session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[_make_session("s1"), _make_session("s2")])
session_repo.cleanupInactiveSessions = AsyncMock(return_value=2)
handler = ChannelSessionInactiveCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
session_repo_factory=lambda _db: session_repo,
session_repo=session_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -98,13 +89,11 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_no_inactive_sessions_returns_zero(self, fake_logger, fake_cache):
# Arrange - 无非活跃会话时直接返回
db = MagicMock()
session_repo = AsyncMock()
session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[])
handler = ChannelSessionInactiveCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
session_repo_factory=lambda _db: session_repo,
session_repo=session_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -120,14 +109,12 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_payload_overrides_defaults(self, fake_logger, fake_cache):
# Arrange - payload 指定 inactive_threshold_minutes=30, batch_size=50
db = MagicMock()
session_repo = AsyncMock()
session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[])
session_repo.cleanupInactiveSessions = AsyncMock(return_value=0)
handler = ChannelSessionInactiveCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
session_repo_factory=lambda _db: session_repo,
session_repo=session_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -143,15 +130,13 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_realtime_metrics_release_called(self, fake_logger, fake_cache):
# Arrange - realtime_metrics 非空时对每个 session_id 调用 releaseSession
db = MagicMock()
session_repo = AsyncMock()
session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[_make_session("s1"), _make_session("s2")])
session_repo.cleanupInactiveSessions = AsyncMock(return_value=2)
realtime_metrics = AsyncMock()
handler = ChannelSessionInactiveCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
session_repo_factory=lambda _db: session_repo,
session_repo=session_repo,
cache_port=fake_cache,
logger=fake_logger,
realtime_metrics=realtime_metrics,
@ -169,14 +154,12 @@ class TestHandlerExecute:
@pytest.mark.asyncio
async def test_realtime_metrics_none_skips_release(self, fake_logger, fake_cache):
# Arrange - realtime_metrics 为 None 时不调用 releaseSession
db = MagicMock()
session_repo = AsyncMock()
session_repo.listInactiveTemporarySessions = AsyncMock(return_value=[_make_session("s1")])
session_repo.cleanupInactiveSessions = AsyncMock(return_value=1)
handler = ChannelSessionInactiveCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
session_repo_factory=lambda _db: session_repo,
session_repo=session_repo,
cache_port=fake_cache,
logger=fake_logger,
realtime_metrics=None,
@ -198,13 +181,11 @@ class TestHandlerErrors:
@pytest.mark.asyncio
async def test_repository_exception_returns_failure(self, fake_logger, fake_cache):
# Arrange - 仓储抛异常
db = MagicMock()
session_repo = AsyncMock()
session_repo.listInactiveTemporarySessions = AsyncMock(side_effect=RuntimeError("db error"))
handler = ChannelSessionInactiveCleanupHandler(
session_factory=lambda: _fake_session_factory(db),
session_repo_factory=lambda _db: session_repo,
session_repo=session_repo,
cache_port=fake_cache,
logger=fake_logger,
)
@ -215,26 +196,3 @@ class TestHandlerErrors:
# Assert
assert result.success is False
assert "db error" in (result.error or "")
@pytest.mark.asyncio
async def test_session_factory_exception_returns_failure(self, fake_logger, fake_cache):
# Arrange - session_factory 自身抛异常
@asynccontextmanager
async def _broken_session_factory():
raise RuntimeError("session creation failed")
yield # pragma: no cover
handler = ChannelSessionInactiveCleanupHandler(
session_factory=_broken_session_factory,
session_repo_factory=lambda _db: AsyncMock(),
cache_port=fake_cache,
logger=fake_logger,
)
# Act
result = await handler.execute(_make_ctx())
# Assert
assert result.success is False
assert "session creation failed" in (result.error or "")

View File

@ -3,15 +3,21 @@
覆盖 ``LoadBuildStage``
- ``process(context)``正常路径含富消息/附件无富消息空分块
image_url/video_url 转换为附件
- 持久化模式下 AgentRun 内容加载streamAgentRun 阻塞消费 + getAgentRunFinalOutput
"""
from __future__ import annotations
from collections.abc import AsyncIterator
from unittest.mock import AsyncMock, MagicMock
import pytest
from yuxi.channels.application.context.outbound_context import OutboundContext
from yuxi.channels.application.pipeline.outbound.load_build_stage import LoadBuildStage
from yuxi.channels.contract.dtos.agent_run import AgentRunId
from yuxi.channels.contract.dtos.channel import ChannelType
from yuxi.channels.contract.dtos.common import Attachment
from yuxi.channels.contract.dtos.option import Nothing, Some
from yuxi.channels.contract.dtos.outbound import RichMessage
from yuxi.channels.contract.plugin.extension_point import FailureStrategy
@ -33,12 +39,36 @@ def _make_ctx(**overrides) -> OutboundContext:
return OutboundContext(**defaults)
def _make_agent_run_port(
*,
stream_events: list | None = None,
final_output: Some | Nothing = None,
) -> MagicMock:
"""构造 mock AgentRunPort。
stream_eventsstreamAgentRun 产出的事件列表按序 yield 后结束迭代
final_outputgetAgentRunFinalOutput 返回的 Option
"""
port = MagicMock()
async def _stream(_run_id: AgentRunId) -> AsyncIterator:
for event in stream_events or []:
yield event
port.streamAgentRun = MagicMock(side_effect=_stream)
port.getAgentRunFinalOutput = AsyncMock(return_value=final_output or Nothing())
return port
# ─── process 正常路径 ────────────────────────────────────────────────────────
@pytest.mark.unit
class TestLoadBuildStageProcess:
@pytest.mark.asyncio
async def test_normal_path_builds_payload_with_rich_message(self):
# Arrange
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
rich = RichMessage(text="hello", image_url="http://img", video_url="http://vid")
ctx = _make_ctx(
stream_chunks=["chunk-a", "chunk-b"],
@ -65,7 +95,7 @@ class TestLoadBuildStageProcess:
@pytest.mark.asyncio
async def test_no_rich_message_sets_fields_none(self):
# Arrange
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
ctx = _make_ctx(stream_chunks=["x"], rich_message=None)
# Act
@ -78,10 +108,10 @@ class TestLoadBuildStageProcess:
assert ctx.outbound_payload.attachments == ()
@pytest.mark.asyncio
async def test_empty_stream_chunks(self):
# Arrange
stage = LoadBuildStage()
ctx = _make_ctx(stream_chunks=[])
async def test_empty_stream_chunks_without_agent_run_id(self):
# Arrange — agent_run_id 为空时不触发内容加载stream_chunks 保持空
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
ctx = _make_ctx(agent_run_id="", stream_chunks=[])
# Act
ok = await stage.process(ctx)
@ -92,10 +122,10 @@ class TestLoadBuildStageProcess:
@pytest.mark.asyncio
async def test_rich_message_without_media_urls(self):
# Arrangerich_message 无 image_url / video_url附件仅含原始
stage = LoadBuildStage()
# Arrange
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
rich = RichMessage(text="text-only")
ctx = _make_ctx(rich_message=rich)
ctx = _make_ctx(stream_chunks=["x"], rich_message=rich)
# Act
ok = await stage.process(ctx)
@ -119,32 +149,38 @@ class TestLoadBuildStageContract:
"""阶段契约属性StageContract"""
def test_reads_agent_run_and_payload_inputs(self):
stage = LoadBuildStage()
assert stage.reads == ("agent_run_id", "stream_chunks", "rich_message", "attachments")
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
assert stage.reads == (
"agent_run_id",
"stream_chunks",
"rich_message",
"attachments",
"delivery_mode",
)
def test_writes_outbound_payload(self):
stage = LoadBuildStage()
assert stage.writes == ("outbound_payload",)
def test_writes_outbound_payload_and_stream_chunks(self):
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
assert stage.writes == ("outbound_payload", "stream_chunks")
def test_is_idempotent(self):
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
assert stage.idempotent is True
def test_is_thread_safe(self):
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
assert stage.thread_safe is True
def test_failure_strategy_is_terminate(self):
# failure=TERMINATE装配失败时终止管道INV-8 agent_run_id 非空约束)
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
assert stage.failure is FailureStrategy.TERMINATE
def test_compensate_is_none(self):
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
assert stage.compensate is None
def test_condition_is_none(self):
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
assert stage.condition is None
@ -158,9 +194,9 @@ class TestLoadBuildStageEdgeCases:
@pytest.mark.asyncio
async def test_image_only_attachment_appended(self):
# Arrange — rich_message 仅含 image_url应转换为 image 附件并入
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
rich = RichMessage(text="img-only", image_url="http://img")
ctx = _make_ctx(rich_message=rich)
ctx = _make_ctx(stream_chunks=["x"], rich_message=rich)
# Act
await stage.process(ctx)
@ -174,9 +210,9 @@ class TestLoadBuildStageEdgeCases:
@pytest.mark.asyncio
async def test_video_only_attachment_appended(self):
# Arrange — rich_message 仅含 video_url应转换为 video 附件并入
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
rich = RichMessage(text="vid-only", video_url="http://vid")
ctx = _make_ctx(rich_message=rich)
ctx = _make_ctx(stream_chunks=["x"], rich_message=rich)
# Act
await stage.process(ctx)
@ -190,13 +226,13 @@ class TestLoadBuildStageEdgeCases:
@pytest.mark.asyncio
async def test_image_before_video_in_attachment_order(self):
# Arrange — image_url 与 video_url 同时存在时image 在前 video 在后
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
rich = RichMessage(
text="both",
image_url="http://img",
video_url="http://vid",
)
ctx = _make_ctx(rich_message=rich)
ctx = _make_ctx(stream_chunks=["x"], rich_message=rich)
# Act
await stage.process(ctx)
@ -208,7 +244,7 @@ class TestLoadBuildStageEdgeCases:
@pytest.mark.asyncio
async def test_stream_chunks_list_converted_to_tuple(self):
# Arrange — list 输入应转为 tuple 以保证 OutboundPayload 不可变语义
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
ctx = _make_ctx(stream_chunks=["a", "b", "c"])
# Act
@ -221,8 +257,9 @@ class TestLoadBuildStageEdgeCases:
@pytest.mark.asyncio
async def test_attachments_tuple_type_preserved(self):
# Arrange — 原始 attachments 为 tuple输出仍为 tuple
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
ctx = _make_ctx(
stream_chunks=["x"],
attachments=(Attachment(type="file", url="http://file"),),
)
@ -235,9 +272,10 @@ class TestLoadBuildStageEdgeCases:
@pytest.mark.asyncio
async def test_original_attachments_prepend_media_attachments(self):
# Arrange — 原始附件在前image/video 在后
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
rich = RichMessage(text="x", image_url="http://img", video_url="http://vid")
ctx = _make_ctx(
stream_chunks=["x"],
rich_message=rich,
attachments=(Attachment(type="file", url="http://file"),),
)
@ -252,9 +290,9 @@ class TestLoadBuildStageEdgeCases:
@pytest.mark.asyncio
async def test_rich_message_fields_wraps_rich_message(self):
# Arrange — rich_message 应被 RichMessageFields 包装
stage = LoadBuildStage()
stage = LoadBuildStage(agent_run_port=_make_agent_run_port())
rich = RichMessage(text="wrap")
ctx = _make_ctx(rich_message=rich)
ctx = _make_ctx(stream_chunks=["x"], rich_message=rich)
# Act
await stage.process(ctx)
@ -262,3 +300,102 @@ class TestLoadBuildStageEdgeCases:
# Assert
assert ctx.outbound_payload.rich_message_fields is not None
assert ctx.outbound_payload.rich_message_fields.rich_message is rich
# ─── 持久化模式 AgentRun 内容加载 ────────────────────────────────────────────
@pytest.mark.unit
class TestLoadBuildStagePersistentContentLoad:
"""持久化模式下 AgentRun 内容加载。"""
@pytest.mark.asyncio
async def test_persistent_mode_loads_agent_run_output(self):
# Arrange — 持久化模式 + agent_run_id 非空 + stream_chunks 为空 → 触发加载
port = _make_agent_run_port(
stream_events=["evt-1", "evt-2"],
final_output=Some("full response text"),
)
stage = LoadBuildStage(agent_run_port=port)
ctx = _make_ctx(delivery_mode="persistent", agent_run_id="run-001", stream_chunks=[])
# Act
ok = await stage.process(ctx)
# Assert
assert ok is True
# streamAgentRun 被调用(阻塞消费至完成)
port.streamAgentRun.assert_called_once()
# getAgentRunFinalOutput 被调用(获取完整文本)
port.getAgentRunFinalOutput.assert_called_once()
# 完整文本填充到 stream_chunks
assert ctx.stream_chunks == ["full response text"]
assert ctx.outbound_payload.stream_chunks == ("full response text",)
@pytest.mark.asyncio
async def test_persistent_mode_final_output_nothing_keeps_empty(self):
# Arrange — getAgentRunFinalOutput 返回 NothingAgentRun 无文本输出)
port = _make_agent_run_port(
stream_events=["evt-1"],
final_output=Nothing(),
)
stage = LoadBuildStage(agent_run_port=port)
ctx = _make_ctx(delivery_mode="persistent", agent_run_id="run-001", stream_chunks=[])
# Act
ok = await stage.process(ctx)
# Assert — stream_chunks 保持空,不抛异常
assert ok is True
assert ctx.stream_chunks == []
assert ctx.outbound_payload.stream_chunks == ()
@pytest.mark.asyncio
async def test_streaming_mode_skips_agent_run_load(self):
# Arrange — 流式模式不触发内容加载(由 stream_chunk_stage 负责)
port = _make_agent_run_port()
stage = LoadBuildStage(agent_run_port=port)
ctx = _make_ctx(delivery_mode="streaming", agent_run_id="run-001", stream_chunks=[])
# Act
ok = await stage.process(ctx)
# Assert
assert ok is True
port.streamAgentRun.assert_not_called()
port.getAgentRunFinalOutput.assert_not_called()
@pytest.mark.asyncio
async def test_pre_filled_stream_chunks_skips_agent_run_load(self):
# Arrange — stream_chunks 已预填充(如静默命令响应路径)不触发加载
port = _make_agent_run_port()
stage = LoadBuildStage(agent_run_port=port)
ctx = _make_ctx(
delivery_mode="persistent",
agent_run_id="run-001",
stream_chunks=["pre-filled"],
)
# Act
ok = await stage.process(ctx)
# Assert
assert ok is True
port.streamAgentRun.assert_not_called()
port.getAgentRunFinalOutput.assert_not_called()
assert ctx.outbound_payload.stream_chunks == ("pre-filled",)
@pytest.mark.asyncio
async def test_empty_agent_run_id_skips_agent_run_load(self):
# Arrange — agent_run_id 为空(如管理员消息)不触发加载
port = _make_agent_run_port()
stage = LoadBuildStage(agent_run_port=port)
ctx = _make_ctx(delivery_mode="persistent", agent_run_id="", stream_chunks=["manual"])
# Act
ok = await stage.process(ctx)
# Assert
assert ok is True
port.streamAgentRun.assert_not_called()
port.getAgentRunFinalOutput.assert_not_called()

View File

@ -123,11 +123,11 @@ class TestPrefixStageContract:
stage = PrefixStage()
assert stage.thread_safe is True
def test_failure_strategy_is_skip(self):
# failure=SKIPtrusted_message 缺失时返回 False管道跳过本阶段但不
# 终止,下游阶段继续执行(消息仍可投递,仅无前缀)
def test_failure_strategy_is_terminate(self):
# failure=TERMINATE非流式路径下 content 必须非空fail-closed
# ValidationError 终止管道而非吞掉空消息继续投递
stage = PrefixStage()
assert stage.failure is FailureStrategy.SKIP
assert stage.failure is FailureStrategy.TERMINATE
def test_compensate_is_none(self):
stage = PrefixStage()

View File

@ -4,18 +4,26 @@
- ``process(context)``适配器不存在降级正常流式投递TTL 超时停止
投递异常降级预填充分块遍历
- ``_iterChunkContents()``agent_run_id 非空走 streamAgentRun为空遍历 stream_chunks
- ``_extractTextFromStreamEvent()``event_type 过滤两层 payload 穿透
stream_event.content/response 回退
- 降级场景TTL/异常降级后 formatted_message 重新格式化
"""
from __future__ import annotations
import itertools
from unittest.mock import AsyncMock, MagicMock
import pytest
from yuxi.channels.application.context.outbound_context import OutboundContext
from yuxi.channels.application.pipeline.outbound.stream_chunk_stage import (
StreamChunkStage,
_extractTextFromStreamEvent,
)
from yuxi.channels.contract.dtos.channel import ChannelType
from yuxi.channels.contract.dtos.common import MessageFormat
from yuxi.channels.contract.dtos.outbound import FormattedMessage
from yuxi.channels.contract.dtos.stream_event import StreamEvent
from yuxi.channels.contract.dtos.streaming import StreamingConfig
from yuxi.channels.contract.errors.base import Error
from yuxi.channels.contract.errors.domain import ChannelDegradedError
@ -77,6 +85,40 @@ def _make_async_iter(items: list):
return _Iter()
def _make_stream_event(content: str, *, event_type: str = "messages") -> StreamEvent:
"""构造 messages 类型的 StreamEventpayload 含单 item。
item 结构与 run_worker 写入 Redis Stream 的一致
{"stream_event": {"content": ...}, "response": ...}
"""
return StreamEvent(
event_type=event_type,
payload={
"event": event_type,
"payload": {
"items": [
{
"stream_event": {"content": content},
"response": content,
"status": "loading",
}
]
},
},
seq="1-0",
)
def _make_outbound_adapter(*, formatted_content: str = "formatted") -> AsyncMock:
"""构造出站适配器桩formatOutbound 返回指定内容。"""
adapter = AsyncMock()
adapter.formatOutbound.return_value = FormattedMessage(
content=formatted_content,
format=MessageFormat.TEXT,
)
return adapter
@pytest.mark.unit
class TestStreamChunkStageProcess:
@pytest.mark.asyncio
@ -86,6 +128,7 @@ class TestStreamChunkStageProcess:
streaming_adapter_registry={},
config_port=_make_config_port(),
agent_run_port=AsyncMock(),
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx()
@ -101,16 +144,21 @@ class TestStreamChunkStageProcess:
@pytest.mark.asyncio
async def test_normal_streaming_delivery_with_agent_run(self):
# Arrange
# ArrangestreamAgentRun 返回真实 StreamEvent 对象
adapter = AsyncMock()
adapter.sendChunk.return_value = MagicMock(success=True)
agent_run_port = AsyncMock()
# streamAgentRun 在源码中为同步调用(未 await返回异步迭代器
agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["chunk-1", "chunk-2"]))
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter([
_make_stream_event("chunk-1"),
_make_stream_event("chunk-2"),
])
)
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(min_interval_ms=0, ttl_ms=60000),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx(agent_run_id="run-001")
@ -134,6 +182,7 @@ class TestStreamChunkStageProcess:
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(),
agent_run_port=AsyncMock(),
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx(agent_run_id="", stream_chunks=["a", "b", "c"])
@ -152,12 +201,19 @@ class TestStreamChunkStageProcess:
adapter = AsyncMock()
adapter.sendChunk.return_value = MagicMock(success=True)
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["chunk-1", "chunk-2", "chunk-3"]))
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter([
_make_stream_event("chunk-1"),
_make_stream_event("chunk-2"),
_make_stream_event("chunk-3"),
])
)
logger = AsyncMock()
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(ttl_ms=0),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=logger,
)
ctx = _make_ctx()
@ -165,8 +221,9 @@ class TestStreamChunkStageProcess:
# monkeypatch time.monotonic 返回递增值,避免 Windows 时钟精度导致
# elapsed=0.0 > ttl_s=0.0 为 False 的 flaky 失败。
# 序列说明started_at=0.0,首次循环 elapsed=0.0(不超时,发送 chunk-1
# 第二次循环 elapsed=0.002 > 0超时breakresults 含 1 个分块。
time_values = iter([0.0, 0.0, 0.002, 0.003])
# 第二次循环 elapsed=0.001 > 0超时break
# 使用 chain+count 提供无限序列,避免 teardown 时迭代器耗尽。
time_values = itertools.chain([0.0, 0.0], itertools.count(0.001, 0.001))
monkeypatch.setattr(
"yuxi.channels.application.pipeline.outbound.stream_chunk_stage.time.monotonic",
lambda: next(time_values),
@ -185,11 +242,14 @@ class TestStreamChunkStageProcess:
adapter = AsyncMock()
adapter.sendChunk.side_effect = Error("send failed")
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["x"]))
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter([_make_stream_event("x")])
)
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx()
@ -203,19 +263,22 @@ class TestStreamChunkStageProcess:
assert ctx.stream_aborted_at_chunk == 0
@pytest.mark.asyncio
async def test_send_chunk_raises_asyncio_timeout_error_degrades(self):
# Arrange原生 TimeoutError 不属于 Error 子类,
# 修复前会被 except Error 漏掉导致消息丢失C2-O3
async def test_send_chunk_raises_timeout_error_degrades(self):
# Arrange原生 TimeoutError 不属于 Error 子类C2-O3
adapter = AsyncMock()
adapter.sendChunk.side_effect = TimeoutError()
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter(["chunk-a", "chunk-b"])
return_value=_make_async_iter([
_make_stream_event("chunk-a"),
_make_stream_event("chunk-b"),
])
)
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx()
@ -237,12 +300,16 @@ class TestStreamChunkStageProcess:
adapter.sendChunk.side_effect = ConnectionError("network down")
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter(["chunk-a", "chunk-b"])
return_value=_make_async_iter([
_make_stream_event("chunk-a"),
_make_stream_event("chunk-b"),
])
)
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx()
@ -258,18 +325,23 @@ class TestStreamChunkStageProcess:
assert exc_info.value.__cause__ is not None
@pytest.mark.asyncio
async def test_send_chunk_raises_native_error_after_partial_send_sets_position(self):
async def test_send_chunk_raises_error_after_partial_send_sets_position(self):
# Arrange第二个分块投递失败stream_aborted_at_chunk 应记录已投递数量
adapter = AsyncMock()
adapter.sendChunk.side_effect = [MagicMock(success=True), ConnectionError("down")]
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter(["chunk-a", "chunk-b", "chunk-c"])
return_value=_make_async_iter([
_make_stream_event("chunk-a"),
_make_stream_event("chunk-b"),
_make_stream_event("chunk-c"),
])
)
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx()
@ -290,6 +362,7 @@ class TestStreamChunkStageProcess:
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(),
agent_run_port=AsyncMock(),
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx(agent_run_id="", stream_chunks=[])
@ -300,9 +373,6 @@ class TestStreamChunkStageProcess:
# Assert
assert ok is True
assert ctx.chunk_results == []
# 空分块流不构造 StreamingCompletedtotal_chunks=0 违反 DTO 不变量,
# 且 truncation_check 依赖 streaming_completed is not None 作为激活条件,
# 空流无需截断检测)
assert ctx.streaming_completed is None
@ -310,13 +380,19 @@ class TestStreamChunkStageProcess:
class TestIterChunkContents:
@pytest.mark.asyncio
async def test_agent_run_id_non_empty_uses_stream_agent_run(self):
# Arrange
# ArrangestreamAgentRun 返回 StreamEvent 对象
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(return_value=_make_async_iter(["a", "b"]))
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter([
_make_stream_event("a"),
_make_stream_event("b"),
])
)
stage = StreamChunkStage(
streaming_adapter_registry={},
config_port=_make_config_port(),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx(agent_run_id="run-001", stream_chunks=[])
@ -326,7 +402,7 @@ class TestIterChunkContents:
async for content in stage._iterChunkContents(ctx):
contents.append(content)
# Assert分块内容追加到 stream_chunks
# Assert从 StreamEvent 提取的文本追加到 stream_chunks
assert contents == ["a", "b"]
assert ctx.stream_chunks == ["a", "b"]
agent_run_port.streamAgentRun.assert_called_once()
@ -339,6 +415,7 @@ class TestIterChunkContents:
streaming_adapter_registry={},
config_port=_make_config_port(),
agent_run_port=agent_run_port,
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx(agent_run_id="", stream_chunks=["x", "y"])
@ -351,3 +428,253 @@ class TestIterChunkContents:
# Assert不调用 streamAgentRun直接遍历预填充列表
assert contents == ["x", "y"]
agent_run_port.streamAgentRun.assert_not_awaited()
@pytest.mark.unit
class TestExtractTextFromStreamEvent:
"""StreamEvent 内容提取逻辑。"""
def test_messages_event_extracts_stream_event_content(self):
# Arrangestream_event.content 存在时优先取
event = _make_stream_event("hello")
# Act
contents = list(_extractTextFromStreamEvent(event))
# Assert
assert contents == ["hello"]
def test_non_messages_event_skipped(self):
# Arrangemetadata/custom/error/end 等事件不提取内容
event = StreamEvent(
event_type="metadata",
payload={"payload": {"items": [{"response": "should-be-skipped"}]}},
seq="1-0",
)
# Act
contents = list(_extractTextFromStreamEvent(event))
# Assert
assert contents == []
def test_falls_back_to_response_when_no_stream_event_content(self):
# Arrangestream_event.content 缺失时回退到 response
event = StreamEvent(
event_type="messages",
payload={
"payload": {
"items": [{"response": "fallback-text"}]
}
},
seq="1-0",
)
# Act
contents = list(_extractTextFromStreamEvent(event))
# Assert
assert contents == ["fallback-text"]
def test_multiple_items_extracted_in_order(self):
# Arrange一个 messages 事件含多个 item
event = StreamEvent(
event_type="messages",
payload={
"payload": {
"items": [
{"stream_event": {"content": "part-1"}},
{"stream_event": {"content": "part-2"}},
{"response": "part-3"},
]
}
},
seq="1-0",
)
# Act
contents = list(_extractTextFromStreamEvent(event))
# Assert
assert contents == ["part-1", "part-2", "part-3"]
def test_empty_content_skipped(self):
# Arrange空字符串 content 不产出
event = StreamEvent(
event_type="messages",
payload={
"payload": {
"items": [
{"stream_event": {"content": ""}},
{"response": ""},
]
}
},
seq="1-0",
)
# Act
contents = list(_extractTextFromStreamEvent(event))
# Assert
assert contents == []
def test_non_dict_item_skipped(self):
# Arrange非 dict item 跳过
event = StreamEvent(
event_type="messages",
payload={"payload": {"items": ["not-a-dict", {"response": "valid"}]}},
seq="1-0",
)
# Act
contents = list(_extractTextFromStreamEvent(event))
# Assert
assert contents == ["valid"]
def test_empty_payload_returns_nothing(self):
# Arrangepayload 为空 dict
event = StreamEvent(event_type="messages", payload={}, seq="1-0")
# Act
contents = list(_extractTextFromStreamEvent(event))
# Assert
assert contents == []
@pytest.mark.unit
class TestReformatAfterDegrade:
"""降级后 formatted_message 重新格式化。"""
@pytest.mark.asyncio
async def test_ttl_degrade_reformats_formatted_message(self, monkeypatch):
# ArrangeTTL 超时降级后stream_chunks 已填充formatted_message 应更新
adapter = AsyncMock()
adapter.sendChunk.return_value = MagicMock(success=True)
outbound_adapter = _make_outbound_adapter(formatted_content="reformatted")
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter([
_make_stream_event("chunk-1"),
_make_stream_event("chunk-2"),
])
)
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(ttl_ms=0),
agent_run_port=agent_run_port,
outbound_adapter_registry={ChannelType("feishu"): outbound_adapter},
logger=AsyncMock(),
)
ctx = _make_ctx()
ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT)
# TTL=0 使首块后超时started_at=0.0,首块 elapsed=0.0(不超时),
# 第二块 elapsed=0.001 > 0超时break
time_values = itertools.chain([0.0, 0.0], itertools.count(0.001, 0.001))
monkeypatch.setattr(
"yuxi.channels.application.pipeline.outbound.stream_chunk_stage.time.monotonic",
lambda: next(time_values),
)
# Act
await stage.process(ctx)
# Assert降级后 formatted_message 已更新content 非空
assert ctx.delivery_mode == "persistent"
assert ctx.formatted_message.content == "reformatted"
outbound_adapter.formatOutbound.assert_awaited_once()
@pytest.mark.asyncio
async def test_exception_degrade_reformats_formatted_message(self):
# Arrange投递异常降级后stream_chunks 已填充formatted_message 应更新
adapter = AsyncMock()
adapter.sendChunk.side_effect = ConnectionError("down")
outbound_adapter = _make_outbound_adapter(formatted_content="reformatted")
agent_run_port = AsyncMock()
agent_run_port.streamAgentRun = MagicMock(
return_value=_make_async_iter([_make_stream_event("chunk-1")])
)
stage = StreamChunkStage(
streaming_adapter_registry={ChannelType("feishu"): adapter},
config_port=_make_config_port(),
agent_run_port=agent_run_port,
outbound_adapter_registry={ChannelType("feishu"): outbound_adapter},
logger=AsyncMock(),
)
ctx = _make_ctx()
ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT)
# Act
with pytest.raises(ChannelDegradedError):
await stage.process(ctx)
# Assert降级后抛出前formatted_message 已更新
assert ctx.delivery_mode == "persistent"
assert ctx.formatted_message.content == "reformatted"
outbound_adapter.formatOutbound.assert_awaited_once()
@pytest.mark.asyncio
async def test_adapter_not_registered_degrade_reformats(self):
# Arrange适配器未注册降级时stream_chunks 已预填充
outbound_adapter = _make_outbound_adapter(formatted_content="reformatted")
stage = StreamChunkStage(
streaming_adapter_registry={},
config_port=_make_config_port(),
agent_run_port=AsyncMock(),
outbound_adapter_registry={ChannelType("feishu"): outbound_adapter},
logger=AsyncMock(),
)
ctx = _make_ctx(stream_chunks=["pre-filled"], agent_run_id="")
ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT)
# Act
await stage.process(ctx)
# Assert
assert ctx.delivery_mode == "persistent"
assert ctx.formatted_message.content == "reformatted"
@pytest.mark.asyncio
async def test_degrade_with_empty_chunks_skips_reformat(self):
# Arrange降级时 stream_chunks 为空,不调用 formatOutbound
outbound_adapter = _make_outbound_adapter()
stage = StreamChunkStage(
streaming_adapter_registry={},
config_port=_make_config_port(),
agent_run_port=AsyncMock(),
outbound_adapter_registry={ChannelType("feishu"): outbound_adapter},
logger=AsyncMock(),
)
ctx = _make_ctx(stream_chunks=[], agent_run_id="")
ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT)
# Act
await stage.process(ctx)
# Assert空 chunks 不重新格式化
outbound_adapter.formatOutbound.assert_not_awaited()
assert ctx.formatted_message.content == ""
@pytest.mark.asyncio
async def test_degrade_without_outbound_adapter_falls_back_to_text(self):
# Arrange出站适配器未注册时直接构造 TEXT 格式
stage = StreamChunkStage(
streaming_adapter_registry={},
config_port=_make_config_port(),
agent_run_port=AsyncMock(),
outbound_adapter_registry={},
logger=AsyncMock(),
)
ctx = _make_ctx(stream_chunks=["fallback-text"], agent_run_id="")
ctx.formatted_message = FormattedMessage(content="", format=MessageFormat.TEXT)
# Act
await stage.process(ctx)
# Assert无出站适配器时直接用原始 content 构造 TEXT
assert ctx.delivery_mode == "persistent"
assert ctx.formatted_message.content == "fallback-text"
assert ctx.formatted_message.format == MessageFormat.TEXT

View File

@ -410,6 +410,8 @@ class TestApproveMerge:
user_identity_repo.getUserIdentityByIdentityId = AsyncMock(
side_effect=lambda iid: {"canonical-1": canonical, "sibling-1": sibling}.get(iid)
)
# updateUserIdentity 返回输入 DTO模拟 DB 更新后返回version 与聚合根一致)
user_identity_repo.updateUserIdentity = AsyncMock(side_effect=lambda dto, **kw: dto)
session_repo = AsyncMock()
session_repo.findSessionsByFilter = AsyncMock(return_value=sessions)
session_repo.updateChannelSession = AsyncMock(return_value=sessions[0])
@ -431,7 +433,7 @@ class TestApproveMerge:
assert result.canonical_identity_id == _CANONICAL_ID
assert result.merged_identity_id == _SIBLING_ID
assert result.canonical_version == 6
# sibling.version: 3 (DB) +1 (clearPendingReview) +1 (unbindUser, M-7) = 5
# sibling.version: 3 (DB) → +1 (clearPendingReview) 持久化 → +1 (unbindUser) 持久化 = 5
assert result.sibling_version == 5
# 校验 canonical 持久化expected_version 用合并前 DB 版本 5
@ -445,10 +447,11 @@ class TestApproveMerge:
bindings_arg = call_kwargs.args[1]
assert _PEER_SIBLING in bindings_arg.get("feishu", [])
# 校验 sibling pending_review 清除
user_identity_repo.updateUserIdentity.assert_awaited_once()
sibling_dto = user_identity_repo.updateUserIdentity.call_args.args[0]
assert sibling_dto.pending_review is False
# 校验 sibling 两次持久化clearPendingReview + unbindUser 拆分持久化
assert user_identity_repo.updateUserIdentity.await_count == 2
# 最后一次调用是 unbindUser 后的 DTOpending_review 已在第一次清除
last_call_dto = user_identity_repo.updateUserIdentity.await_args.args[0]
assert last_call_dto.pending_review is False
# 校验会话迁移sibling 名下的会话迁移到 canonical
# findSessionsByFilter 用 sibling_id 过滤,返回 2 条会话

View File

@ -161,7 +161,7 @@ class TestOutboxEntryMarkSent:
# Assert
assert entry.status == OutboxStatus.SENT
assert entry.channel_msg_id == "channel-msg-1"
assert entry.version == original_version + 1
assert entry.version == original_version
assert entry.funnel_node == "sent"
assert entry.latency_ms is not None
assert entry.latency_ms >= 0
@ -223,7 +223,7 @@ class TestOutboxEntryMarkSuppressed:
assert entry.status == OutboxStatus.SUPPRESSED
assert entry.last_error == "stale reply"
assert entry.funnel_node == "suppressed"
assert entry.version == original_version + 1
assert entry.version == original_version
@pytest.mark.parametrize(
"status",
@ -301,7 +301,7 @@ class TestOutboxEntryMarkFailed:
assert entry.next_retry_at is not None
assert entry.last_error == "network error"
assert entry.funnel_node == "failed"
assert entry.version == original_version + 1
assert entry.version == original_version
assert entry.last_retry_at is not None
def test_mark_failed_does_not_auto_transition_to_dead(self):
@ -365,7 +365,7 @@ class TestOutboxEntryMarkDead:
assert entry.next_retry_at is None
assert entry.last_error == "retry exhausted"
assert entry.funnel_node == "dead"
assert entry.version == original_version + 1
assert entry.version == original_version
@pytest.mark.parametrize(
"status",
@ -455,7 +455,7 @@ class TestOutboxEntryMarkPartialFailure:
assert entry.status == OutboxStatus.SENT_UNCONFIRMED
assert entry.partial_failure is True
assert entry.last_error == "partial failure on parts: [1, 3]"
assert entry.version == original_version + 1
assert entry.version == original_version
@pytest.mark.parametrize(
"status",
@ -492,7 +492,7 @@ class TestOutboxEntryRequeueForRetry:
assert entry.status == OutboxStatus.PENDING
assert entry.next_retry_at is None
assert entry.funnel_node == "enter"
assert entry.version == original_version + 1
assert entry.version == original_version
# retry_count 保持不变
assert entry.retry_count == 2
@ -656,14 +656,14 @@ class TestOutboxEntryRevive:
# Assert
assert entry.funnel_node == "failed"
def test_revive_increments_version(self):
def test_revive_preserves_version(self):
# Arrange
entry = self._make_dead_entry()
original_version = entry.version
# Act
entry.revive()
# Assert
assert entry.version == original_version + 1
assert entry.version == original_version
def test_revive_from_pending_raises_rule_violation_error(self):
# Arrange

View File

@ -223,7 +223,7 @@ class TestMarkPartialFailure:
assert entry.partial_failure is True
assert "2" in entry.last_error
assert "3" in entry.last_error
assert entry.version == original_version + 1
assert entry.version == original_version
def test_partial_failure_does_not_change_funnel_node(self):
"""markPartialFailure 不修改 funnel_node仍是 SENT_UNCONFIRMED 路径)。"""
@ -326,7 +326,7 @@ class TestRequeueForRetry:
assert entry.status == OutboxStatus.PENDING
assert entry.next_retry_at is None
assert entry.funnel_node == "enter"
assert entry.version == original_version + 1
assert entry.version == original_version
def test_retry_count_unchanged_after_requeue(self):
"""requeueForRetry 不重置 retry_count重试预算已由 markFailed 递增)。"""

View File

@ -3,7 +3,7 @@
覆盖
- ``_fill_adapter_registry``纯函数逻辑最后注册者覆盖 / duplicate_rule 校验
- 8 ``create_channel_*_handler_dependencies`` 工厂函数验证返回字典
键集与 ``*_repo_factory`` 闭包产物类型确保 worker 进程装配契约稳定
键集与 ``xxx_repo`` 端口实例类型确保 worker 进程装配契约稳定
不依赖运行中的 Docker 服务纯单元测试
"""
@ -25,7 +25,6 @@ from yuxi.channels.adapters.structured_logger_adapter import (
from yuxi.channels.contract.dtos.channel import ChannelType
from yuxi.channels.contract.dtos.outbox import OutboxConfig
from yuxi.channels.contract.errors import RuleViolationError
from yuxi.channels.contract.ports.driven.logger_port import LoggerPort
from yuxi.channels.infrastructure.channel_use_cases import (
_fill_adapter_registry,
create_channel_audit_log_retention_handler_dependencies,
@ -156,7 +155,7 @@ class TestFillAdapterRegistry:
@pytest.fixture
def mock_session_factory() -> MagicMock:
"""worker 进程级 session_factory 桩"""
"""worker 进程级 session_factory 桩``Callable[[], AsyncSession]``"""
return MagicMock(name="session_factory")
@ -167,9 +166,9 @@ def mock_arq_pool() -> MagicMock:
@pytest.fixture
def mock_db() -> MagicMock:
"""AsyncSession 桩,传入 *_repo_factory 构造适配器"""
return MagicMock(name="db")
def mock_cache_port() -> MagicMock:
"""worker 进程级 CachePort 桩,供 handler 分布式锁使用"""
return MagicMock(name="cache_port")
def _assert_logger_port(logger: object) -> None:
@ -186,105 +185,94 @@ def _assert_outbox_config(outbox_config: object) -> None:
class TestSessionInactiveCleanupHandlerDependencies:
"""create_channel_session_inactive_cleanup_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(self, mock_session_factory):
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_cache_port
):
# Act
deps = create_channel_session_inactive_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"session_repo_factory",
"logger",
}
assert deps["session_factory"] is mock_session_factory
assert set(deps.keys()) == {"session_repo", "cache_port", "logger"}
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
def test_session_repo_factory_creates_persistence_adapter(
self, mock_session_factory, mock_db
def test_session_repo_is_persistence_adapter(
self, mock_session_factory, mock_cache_port
):
# Arrange
# Act
deps = create_channel_session_inactive_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Act
repo = deps["session_repo_factory"](mock_db)
# Assert: fat adapter 实现 ChannelSessionRepositoryPort
assert isinstance(repo, ChannelPersistenceAdapter)
assert isinstance(deps["session_repo"], ChannelPersistenceAdapter)
@pytest.mark.unit
class TestOutboxTerminalCleanupHandlerDependencies:
"""create_channel_outbox_terminal_cleanup_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(self, mock_session_factory):
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_cache_port
):
# Act
deps = create_channel_outbox_terminal_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"outbox_repo_factory",
"outbox_repo",
"event_publisher",
"cache_port",
"logger",
}
assert deps["session_factory"] is mock_session_factory
# event_publisher 为 Noneworker 进程未注入事件发布端口best-effort
assert deps["event_publisher"] is None
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
def test_outbox_repo_factory_creates_persistence_adapter(
self, mock_session_factory, mock_db
def test_outbox_repo_is_persistence_adapter(
self, mock_session_factory, mock_cache_port
):
# Arrange
# Act
deps = create_channel_outbox_terminal_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Act
repo = deps["outbox_repo_factory"](mock_db)
# Assert
assert isinstance(repo, ChannelPersistenceAdapter)
assert isinstance(deps["outbox_repo"], ChannelPersistenceAdapter)
@pytest.mark.unit
class TestPairingTerminalCleanupHandlerDependencies:
"""create_channel_pairing_terminal_cleanup_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(self, mock_session_factory):
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_cache_port
):
# Act
deps = create_channel_pairing_terminal_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"pairing_repo_factory",
"logger",
}
assert deps["session_factory"] is mock_session_factory
assert set(deps.keys()) == {"pairing_repo", "cache_port", "logger"}
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
def test_pairing_repo_factory_creates_persistence_adapter(
self, mock_session_factory, mock_db
def test_pairing_repo_is_persistence_adapter(
self, mock_session_factory, mock_cache_port
):
# Arrange
# Act
deps = create_channel_pairing_terminal_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Act
repo = deps["pairing_repo_factory"](mock_db)
# Assert
assert isinstance(repo, ChannelPersistenceAdapter)
assert isinstance(deps["pairing_repo"], ChannelPersistenceAdapter)
@pytest.mark.unit
@ -292,176 +280,160 @@ class TestOutboxRecoveryHandlerDependencies:
"""create_channel_outbox_recovery_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_arq_pool
self, mock_session_factory, mock_arq_pool, mock_cache_port
):
# Act
deps = create_channel_outbox_recovery_handler_dependencies(
mock_session_factory, arq_pool=mock_arq_pool
mock_session_factory, arq_pool=mock_arq_pool, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"outbox_repo_factory",
"audit_log_repo_factory",
"outbox_repo",
"audit_log_repo",
"queue_port",
"outbox_config",
"cache_port",
"logger",
}
assert deps["session_factory"] is mock_session_factory
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
_assert_outbox_config(deps["outbox_config"])
def test_outbox_and_audit_repo_factories_create_persistence_adapter(
self, mock_session_factory, mock_arq_pool, mock_db
def test_outbox_and_audit_repos_are_persistence_adapters(
self, mock_session_factory, mock_arq_pool, mock_cache_port
):
"""outbox_repo_factory 与 audit_log_repo_factory 都构造 fat adapter"""
# Arrange
"""outbox_repo 与 audit_log_repo 都为 fat adapter 实例"""
# Act
deps = create_channel_outbox_recovery_handler_dependencies(
mock_session_factory, arq_pool=mock_arq_pool
mock_session_factory, arq_pool=mock_arq_pool, cache_port=mock_cache_port
)
# Act
outbox_repo = deps["outbox_repo_factory"](mock_db)
audit_repo = deps["audit_log_repo_factory"](mock_db)
# Assert: 两者都是 ChannelPersistenceAdapterfat adapter 实现多端口)
assert isinstance(outbox_repo, ChannelPersistenceAdapter)
assert isinstance(audit_repo, ChannelPersistenceAdapter)
assert isinstance(deps["outbox_repo"], ChannelPersistenceAdapter)
assert isinstance(deps["audit_log_repo"], ChannelPersistenceAdapter)
@pytest.mark.unit
class TestPairingExpirationHandlerDependencies:
"""create_channel_pairing_expiration_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(self, mock_session_factory):
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_cache_port
):
# Act
deps = create_channel_pairing_expiration_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"pairing_repo_factory",
"logger",
}
assert set(deps.keys()) == {"pairing_repo", "cache_port", "logger"}
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
def test_pairing_repo_factory_creates_persistence_adapter(
self, mock_session_factory, mock_db
def test_pairing_repo_is_persistence_adapter(
self, mock_session_factory, mock_cache_port
):
# Arrange
# Act
deps = create_channel_pairing_expiration_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Act
repo = deps["pairing_repo_factory"](mock_db)
# Assert
assert isinstance(repo, ChannelPersistenceAdapter)
assert isinstance(deps["pairing_repo"], ChannelPersistenceAdapter)
@pytest.mark.unit
class TestAuditLogRetentionHandlerDependencies:
"""create_channel_audit_log_retention_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(self, mock_session_factory):
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_cache_port
):
# Act
deps = create_channel_audit_log_retention_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"audit_log_repo_factory",
"logger",
}
assert set(deps.keys()) == {"audit_log_repo", "cache_port", "logger"}
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
def test_audit_log_repo_factory_creates_persistence_adapter(
self, mock_session_factory, mock_db
def test_audit_log_repo_is_persistence_adapter(
self, mock_session_factory, mock_cache_port
):
# Arrange
# Act
deps = create_channel_audit_log_retention_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Act
repo = deps["audit_log_repo_factory"](mock_db)
# Assert
assert isinstance(repo, ChannelPersistenceAdapter)
assert isinstance(deps["audit_log_repo"], ChannelPersistenceAdapter)
@pytest.mark.unit
class TestContentReviewRetentionHandlerDependencies:
"""create_channel_content_review_retention_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(self, mock_session_factory):
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_cache_port
):
# Act
deps = create_channel_content_review_retention_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"content_review_repo_factory",
"content_review_repo",
"cache_port",
"logger",
}
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
def test_content_review_repo_factory_creates_review_adapter(
self, mock_session_factory, mock_db
def test_content_review_repo_is_review_adapter(
self, mock_session_factory, mock_cache_port
):
"""与其它 7 个 handler 不同content_review 使用独立的 Adapter。"""
# Arrange
# Act
deps = create_channel_content_review_retention_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Act
repo = deps["content_review_repo_factory"](mock_db)
# Assert
assert isinstance(repo, ContentReviewRepositoryAdapter)
assert not isinstance(repo, ChannelPersistenceAdapter)
assert isinstance(deps["content_review_repo"], ContentReviewRepositoryAdapter)
assert not isinstance(deps["content_review_repo"], ChannelPersistenceAdapter)
@pytest.mark.unit
class TestIdempotencyCleanupHandlerDependencies:
"""create_channel_idempotency_cleanup_handler_dependencies 测试。"""
def test_returns_dict_with_required_keys(self, mock_session_factory):
def test_returns_dict_with_required_keys(
self, mock_session_factory, mock_cache_port
):
# Act
deps = create_channel_idempotency_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Assert
assert set(deps.keys()) == {
"session_factory",
"idempotency_repo_factory",
"logger",
}
assert set(deps.keys()) == {"idempotency_repo", "cache_port", "logger"}
assert deps["cache_port"] is mock_cache_port
_assert_logger_port(deps["logger"])
def test_idempotency_repo_factory_creates_persistence_adapter(
self, mock_session_factory, mock_db
def test_idempotency_repo_is_persistence_adapter(
self, mock_session_factory, mock_cache_port
):
# Arrange
# Act
deps = create_channel_idempotency_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
)
# Act
repo = deps["idempotency_repo_factory"](mock_db)
# Assert
assert isinstance(repo, ChannelPersistenceAdapter)
assert isinstance(deps["idempotency_repo"], ChannelPersistenceAdapter)
# ---------------------------------------------------------------------------
@ -474,34 +446,36 @@ class TestHandlerDependenciesConsistency:
"""所有 worker 进程 handler 依赖工厂的一致性约束。"""
def test_all_factories_produce_logger_port_compliant_instance(
self, mock_session_factory, mock_arq_pool
self, mock_session_factory, mock_arq_pool, mock_cache_port
):
"""8 个工厂的 logger 字段都为 StructuredLoggerAdapter满足 LoggerPort"""
# Act
all_deps = [
create_channel_session_inactive_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
),
create_channel_outbox_terminal_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
),
create_channel_pairing_terminal_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
),
create_channel_outbox_recovery_handler_dependencies(
mock_session_factory, arq_pool=mock_arq_pool
mock_session_factory,
arq_pool=mock_arq_pool,
cache_port=mock_cache_port,
),
create_channel_pairing_expiration_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
),
create_channel_audit_log_retention_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
),
create_channel_content_review_retention_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
),
create_channel_idempotency_cleanup_handler_dependencies(
mock_session_factory
mock_session_factory, cache_port=mock_cache_port
),
]
@ -509,5 +483,5 @@ class TestHandlerDependenciesConsistency:
for deps in all_deps:
assert "logger" in deps
_assert_logger_port(deps["logger"])
# 每个工厂都共享同一 session_factory
assert deps["session_factory"] is mock_session_factory
# 每个工厂都共享同一 cache_port
assert deps["cache_port"] is mock_cache_port

View File

@ -332,9 +332,7 @@ class TestRegisterBuiltinEventSubscribers:
assert event_type in event_bus._subscriptions
assert len(event_bus._subscriptions[event_type]) == 1
def test_registers_config_event_pair_for_whitelist_and_route_match(
self, fake_logger
):
def test_registers_config_event_pair_for_whitelist_and_route_match(self, fake_logger):
"""ConfigChanged/ConfigRollback 各注册 WhitelistConfigHandler +
RouteMatchCacheHandler 2 个订阅"""
# Arrange
@ -375,9 +373,7 @@ class TestRegisterBuiltinEventSubscribers:
assert "PairingApproved" in event_bus._subscriptions
assert len(event_bus._subscriptions["PairingApproved"]) == 1
def test_registers_channel_degraded_and_recovered_subscribers(
self, fake_logger
):
def test_registers_channel_degraded_and_recovered_subscribers(self, fake_logger):
# Arrange
event_bus = self._make_event_bus(fake_logger)
@ -396,9 +392,7 @@ class TestRegisterBuiltinEventSubscribers:
assert len(event_bus._subscriptions["ChannelDegraded"]) == 1
assert len(event_bus._subscriptions["ChannelRecovered"]) == 1
def test_registers_three_channel_event_broadcaster_subscribers(
self, fake_logger
):
def test_registers_three_channel_event_broadcaster_subscribers(self, fake_logger):
"""ChannelSessionUpdated/ChannelMessageReceived/ChannelMessageSent → 广播器。"""
# Arrange
event_bus = self._make_event_bus(fake_logger)
@ -501,9 +495,7 @@ def _make_plugin_loading_core() -> _PluginLoadingCore:
class TestRegisterDiSingletons:
"""``_register_di_singletons`` 注册全量 DI 单例到容器。"""
def test_registers_core_registries_and_shared_dependencies(
self, fake_logger
):
def test_registers_core_registries_and_shared_dependencies(self, fake_logger):
# Arrange
di = DependencyInjectionContainer()
core = _make_plugin_loading_core()
@ -513,7 +505,7 @@ class TestRegisterDiSingletons:
di_container=di,
core=core,
persistence_port=MagicMock(spec=PersistencePort),
persistence_db=MagicMock(),
session_factory=MagicMock(),
queue_port=MagicMock(spec=QueuePort),
config_port=MagicMock(spec=ConfigPort),
stage_slot_injector=MagicMock(spec=StageSlotInjector),
@ -539,9 +531,7 @@ class TestRegisterDiSingletons:
assert di.resolve(SensitiveFieldRegistry) is core.sensitive_registry
assert di.resolve(ConfigScopeRegistry) is core.config_scope_registry
def test_registers_command_and_identity_resolver_registries(
self, fake_logger
):
def test_registers_command_and_identity_resolver_registries(self, fake_logger):
"""命令注册表与身份解析器注册表由本函数构造并注册。"""
# Arrange
di = DependencyInjectionContainer()
@ -552,7 +542,7 @@ class TestRegisterDiSingletons:
di_container=di,
core=core,
persistence_port=MagicMock(spec=PersistencePort),
persistence_db=MagicMock(),
session_factory=MagicMock(),
queue_port=MagicMock(spec=QueuePort),
config_port=MagicMock(spec=ConfigPort),
stage_slot_injector=MagicMock(spec=StageSlotInjector),
@ -568,9 +558,7 @@ class TestRegisterDiSingletons:
identity_registry = di.resolve(IdentityResolverRegistry)
assert isinstance(identity_registry, IdentityResolverRegistry)
def test_registers_app_level_ports_and_orchestration_components(
self, fake_logger
):
def test_registers_app_level_ports_and_orchestration_components(self, fake_logger):
# Arrange
di = DependencyInjectionContainer()
core = _make_plugin_loading_core()
@ -587,7 +575,7 @@ class TestRegisterDiSingletons:
di_container=di,
core=core,
persistence_port=persistence_port,
persistence_db=MagicMock(),
session_factory=MagicMock(),
queue_port=queue_port,
config_port=config_port,
stage_slot_injector=stage_slot_injector,
@ -613,14 +601,14 @@ class TestRegisterDiSingletons:
# Arrange
di = DependencyInjectionContainer()
core = _make_plugin_loading_core()
persistence_db = MagicMock()
session_factory = MagicMock()
# Act
_register_di_singletons(
di_container=di,
core=core,
persistence_port=MagicMock(spec=PersistencePort),
persistence_db=persistence_db,
session_factory=session_factory,
queue_port=MagicMock(spec=QueuePort),
config_port=MagicMock(spec=ConfigPort),
stage_slot_injector=MagicMock(spec=StageSlotInjector),
@ -639,9 +627,7 @@ class TestRegisterDiSingletons:
# ContentReviewRepositoryPort 与 ContentReviewRepositoryAdapter 绑定到同一实例
assert review_port is review_adapter
def test_registers_redis_arq_and_execution_port_singletons(
self, fake_logger
):
def test_registers_redis_arq_and_execution_port_singletons(self, fake_logger):
# Arrange
di = DependencyInjectionContainer()
core = _make_plugin_loading_core()
@ -651,7 +637,7 @@ class TestRegisterDiSingletons:
di_container=di,
core=core,
persistence_port=MagicMock(spec=PersistencePort),
persistence_db=MagicMock(),
session_factory=MagicMock(),
queue_port=MagicMock(spec=QueuePort),
config_port=MagicMock(spec=ConfigPort),
stage_slot_injector=MagicMock(spec=StageSlotInjector),
@ -665,8 +651,8 @@ class TestRegisterDiSingletons:
# (由 core 直接注册)
assert di.resolve(AgentRunExecutionPort) is core.execution_port_impl
def test_persistence_db_not_registered_as_singleton(self, fake_logger):
"""persistence_db 是请求级资源,不应注册为 AsyncSession 单例"""
def test_async_session_not_registered_as_singleton(self, fake_logger):
"""AsyncSession 是请求级资源不应注册为单例session_factory 注册为单例供按需创建"""
# Arrange
di = DependencyInjectionContainer()
core = _make_plugin_loading_core()
@ -677,7 +663,7 @@ class TestRegisterDiSingletons:
di_container=di,
core=core,
persistence_port=MagicMock(spec=PersistencePort),
persistence_db=MagicMock(),
session_factory=MagicMock(),
queue_port=MagicMock(spec=QueuePort),
config_port=MagicMock(spec=ConfigPort),
stage_slot_injector=MagicMock(spec=StageSlotInjector),

View File

@ -229,9 +229,7 @@ class TestTeardownPipelines:
# Arrange
stage_slot_registry = MagicMock(spec=StageSlotRegistry)
# 模拟 inbound 有 2 个槽位outbound 1 个control 0 个
stage_slot_registry.findByPipeline = MagicMock(
side_effect=[["slot1", "slot2"], ["slot3"], []]
)
stage_slot_registry.findByPipeline = MagicMock(side_effect=[["slot1", "slot2"], ["slot3"], []])
shutdown = _make_shutdown(fake_logger, stage_slot_registry=stage_slot_registry)
# Act
@ -239,9 +237,7 @@ class TestTeardownPipelines:
# Assert: info 日志携带槽位数量
info_calls = fake_logger.info.call_args_list
teardown_call = next(
c for c in info_calls if "管道已拆除" in c.args[0]
)
teardown_call = next(c for c in info_calls if "管道已拆除" in c.args[0])
assert teardown_call.kwargs.get("inbound_slot_count") == 2
assert teardown_call.kwargs.get("outbound_slot_count") == 1
assert teardown_call.kwargs.get("control_slot_count") == 0
@ -251,9 +247,7 @@ class TestTeardownPipelines:
"""clear 抛异常时记录告警并继续(不中断 shutdown"""
# Arrange
stage_slot_registry = MagicMock(spec=StageSlotRegistry)
stage_slot_registry.findByPipeline = MagicMock(
side_effect=RuntimeError("registry corrupted")
)
stage_slot_registry.findByPipeline = MagicMock(side_effect=RuntimeError("registry corrupted"))
shutdown = _make_shutdown(fake_logger, stage_slot_registry=stage_slot_registry)
# Act: 不抛异常即证明异常被吞并
@ -339,9 +333,7 @@ class TestStopPlugins:
plugin_registry = MagicMock(spec=PluginRegistry)
# STARTED 返回 [a, b]PAUSED 返回 []
plugin_registry.listPluginsByState = MagicMock(
side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b]
if state == LifecycleState.STARTED
else []
side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b] if state == LifecycleState.STARTED else []
)
# resolver 返回 [a, b]a 先b 后),反向后为 [b, a]
@ -349,9 +341,7 @@ class TestStopPlugins:
resolver.resolve = MagicMock(return_value=[manifest_a, manifest_b])
plugin_lifecycle = MagicMock(spec=PluginLifecycleManager)
plugin_lifecycle.stop = AsyncMock(
return_value=LifecycleResult(state="stopped")
)
plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped"))
shutdown = _make_shutdown(
fake_logger,
@ -382,9 +372,7 @@ class TestStopPlugins:
plugin_registry = MagicMock(spec=PluginRegistry)
plugin_registry.listPluginsByState = MagicMock(
side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b]
if state == LifecycleState.STARTED
else []
side_effect=lambda state: [plugin_manifest_a, plugin_manifest_b] if state == LifecycleState.STARTED else []
)
# resolver.resolve 抛异常
@ -392,9 +380,7 @@ class TestStopPlugins:
resolver.resolve = MagicMock(side_effect=RuntimeError("cycle detected"))
plugin_lifecycle = MagicMock(spec=PluginLifecycleManager)
plugin_lifecycle.stop = AsyncMock(
return_value=LifecycleResult(state="stopped")
)
plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped"))
shutdown = _make_shutdown(
fake_logger,
@ -424,18 +410,14 @@ class TestStopPlugins:
plugin_registry = MagicMock(spec=PluginRegistry)
# STARTED 返回 []PAUSED 返回 [paused-plugin]
plugin_registry.listPluginsByState = MagicMock(
side_effect=lambda state: [plugin_manifest_paused]
if state == LifecycleState.PAUSED
else []
side_effect=lambda state: [plugin_manifest_paused] if state == LifecycleState.PAUSED else []
)
resolver = MagicMock(spec=PluginDependencyResolver)
resolver.resolve = MagicMock(return_value=[manifest_paused])
plugin_lifecycle = MagicMock(spec=PluginLifecycleManager)
plugin_lifecycle.stop = AsyncMock(
return_value=LifecycleResult(state="stopped")
)
plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped"))
shutdown = _make_shutdown(
fake_logger,
@ -460,9 +442,7 @@ class TestStopSinglePlugin:
# Arrange
manifest = _make_manifest("ok-plugin")
plugin_lifecycle = MagicMock(spec=PluginLifecycleManager)
plugin_lifecycle.stop = AsyncMock(
return_value=LifecycleResult(state="stopped")
)
plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="stopped"))
shutdown = _make_shutdown(fake_logger, plugin_lifecycle=plugin_lifecycle)
# Act
@ -481,11 +461,7 @@ class TestStopSinglePlugin:
# Arrange
manifest = _make_manifest("bad-plugin")
plugin_lifecycle = MagicMock(spec=PluginLifecycleManager)
plugin_lifecycle.stop = AsyncMock(
return_value=LifecycleResult(
state="failed", error="cleanup failed"
)
)
plugin_lifecycle.stop = AsyncMock(return_value=LifecycleResult(state="failed", error="cleanup failed"))
shutdown = _make_shutdown(fake_logger, plugin_lifecycle=plugin_lifecycle)
# Act
@ -563,8 +539,9 @@ class TestCloseDrivenAdapters:
# Act
await shutdown._closeAllDrivenAdapters()
# Assert: 记录跳过日志,不调用 aclose
fake_logger.info.assert_any_call(
# Assert: 记录跳过日志DEBUG 级别,因跳过是预期行为而非异常事件),
# 不调用 aclose
fake_logger.debug.assert_any_call(
"应用级被驱动适配器未实现 CloseablePort跳过",
adapter_type=type(non_closeable).__name__,
)
@ -601,9 +578,7 @@ class TestCloseDrivenAdapters:
plugin_adapters = MagicMock()
plugin_adapters.close = AsyncMock()
plugin_registry = MagicMock(spec=PluginRegistry)
plugin_registry.listPluginAdapters = MagicMock(
return_value=[("feishu", plugin_adapters)]
)
plugin_registry.listPluginAdapters = MagicMock(return_value=[("feishu", plugin_adapters)])
shutdown = _make_shutdown(
fake_logger,
plugin_registry=plugin_registry,

View File

@ -86,7 +86,9 @@ def patched_scheduler(monkeypatch):
"""替换 ``scheduler`` 模块的全部外部依赖,返回 ``_PatchTracker``。
替换内容
- ``pg_manager.get_async_session_context``返回 sentinel mock
- ``pg_manager.AsyncSession``返回 sentinel mockscheduler 通过
``pg_manager.AsyncSession`` 获取 ``async_sessionmaker`` 实例作为
``session_factory`` 透传给各 handler 工厂
- ``get_arq_pool``AsyncMock 返回 sentinel mock
- 8 ``create_channel_*_handler_dependencies`` 工厂返回空字典
调用参数记录到 ``factory_calls``
@ -94,16 +96,19 @@ def patched_scheduler(monkeypatch):
"""
tracker = _PatchTracker()
# 替换 pg_manager.get_async_session_context
# 替换 pg_manager.AsyncSessionscheduler 源码通过
# ``session_factory = pg_manager.AsyncSession`` 获取 ``async_sessionmaker``
# 实例callable调用返回 ``AsyncSession``),符合 ``Callable[[], AsyncSession]`` 协议
monkeypatch.setattr(
scheduler_module.pg_manager,
"get_async_session_context",
"AsyncSession",
tracker.session_factory,
)
# 替换 get_arq_pool
async def _fake_get_arq_pool():
return tracker.arq_pool
monkeypatch.setattr(scheduler_module, "get_arq_pool", _fake_get_arq_pool)
# 替换 8 个 handler 类:实例化时返回带正确 name 的 mock
@ -112,18 +117,23 @@ def patched_scheduler(monkeypatch):
def _make_side_effect(h_name: str):
def _instantiate(**kwargs):
return _make_handler_mock(h_name)
return _instantiate
class_mock = MagicMock(side_effect=_make_side_effect(handler_name))
tracker.handler_class_mocks[class_name] = class_mock
monkeypatch.setattr(scheduler_module, class_name, class_mock)
# 替换 8 个工厂函数:返回空字典,记录调用参数
for factory_name, handler_name in _FACTORY_TO_NAME.items():
def _make_factory(h_name: str):
def _factory(*args, **kwargs):
tracker.factory_calls[h_name].append({"args": args, "kwargs": kwargs})
return {}
return _factory
monkeypatch.setattr(scheduler_module, factory_name, _make_factory(handler_name))
return tracker
@ -179,9 +189,7 @@ class TestRegisterSchedulerHandlers:
# Assert: 每个工厂的首个位置参数为 session_factory
for handler_name in EXPECTED_HANDLER_NAMES:
factory_calls = patched_scheduler.factory_calls[handler_name]
assert len(factory_calls) == 1, (
f"factory for {handler_name} called {len(factory_calls)} times"
)
assert len(factory_calls) == 1, f"factory for {handler_name} called {len(factory_calls)} times"
assert factory_calls[0]["args"][0] is patched_scheduler.session_factory
@pytest.mark.asyncio
@ -240,9 +248,7 @@ class TestRegisterSchedulerHandlers:
# Assert: 每个 handler 类的 mock 被调用一次
for class_name in _HANDLER_CLASS_TO_NAME:
class_mock = patched_scheduler.handler_class_mocks[class_name]
assert class_mock.call_count == 1, (
f"{class_name} instantiated {class_mock.call_count} times"
)
assert class_mock.call_count == 1, f"{class_name} instantiated {class_mock.call_count} times"
# 工厂返回空字典,所以实例化参数为空 kwargs
call_kwargs = class_mock.call_args.kwargs
assert call_kwargs == {}

View File

@ -88,7 +88,7 @@ def _make_text_payload() -> dict:
"type": 1,
"content": "hello",
"create_time": "1000",
"session_type": "single",
"session_type": "p2p",
}
@ -102,7 +102,7 @@ def _make_image_payload() -> dict:
"type": 3,
"content": "",
"create_time": "1000",
"session_type": "single",
"session_type": "p2p",
}
@ -668,6 +668,38 @@ class TestNormalizeInboundMessageTypes:
await adapter.normalizeInbound(raw)
assert exc_info.value.field == "talker"
@pytest.mark.asyncio
async def test_p2p_session_type_is_accepted(self):
# 回归测试bridge 返回的单聊 session_type 为 "p2p"(非 "single"
# 必须通过校验,不应抛出 session_type_invalid。
client = _make_client()
adapter = WeChatWocInboundAdapter(client, _make_logger())
payload = _make_text_payload()
payload["session_type"] = "p2p"
raw = _make_raw_event(payload, account_id="acct-1")
# Act
content = await adapter.normalizeInbound(raw)
# Assert
assert content.metadata["session_type"] == "p2p"
@pytest.mark.asyncio
async def test_invalid_session_type_raises_validation_error(self):
# Arrange - session_type 不在合法枚举内
client = _make_client()
adapter = WeChatWocInboundAdapter(client, _make_logger())
payload = _make_text_payload()
payload["session_type"] = "unknown"
raw = _make_raw_event(payload, account_id="acct-1")
# Act / Assert
from yuxi.channels.contract.errors import ValidationError
with pytest.raises(ValidationError) as exc_info:
await adapter.normalizeInbound(raw)
assert exc_info.value.field == "session_type"
# ------------------------------------------------------------------
# normalizeInbound: trace_id 通过 ContextVar 隐式传递P1-10
@ -706,7 +738,7 @@ class TestNormalizeInboundTraceId:
msg_type=1,
sender="sender-1",
is_sender=0,
session_type="single",
session_type="p2p",
text_preview="hello",
attachment_count=0,
trace_id="ctx-trace-id-123",