ForcePilot/backend/test/unit/channels/infrastructure/test_scheduler.py

262 lines
11 KiB
Python
Raw Normal View History

"""yuxi.channels.infrastructure.scheduler 单元测试。
覆盖 ``register_scheduler_handlers`` 注册行为
- 注册全部 8 handler``registry.list_names`` 返回 8 个预期名称
- 每个 handler ``name`` 属性与预期一致
- 工厂函数被正确调用``session_factory`` 透传``arq_pool`` 仅传给 outbox_recovery
- ``get_arq_pool`` 在注册期间被调用一次
- handler 类用工厂返回的字典实例化``**deps``
不依赖运行中的 Docker / Redis / DB 服务通过 ``monkeypatch`` 替换
``scheduler`` 模块的外部依赖``pg_manager`` / ``get_arq_pool`` / 工厂函数 /
handler 使用真实 ``HandlerRegistry`` 验证注册结果
"""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from yuxi.channels.infrastructure import scheduler as scheduler_module
from yuxi.scheduler.framework.runtime.handler_registry import HandlerRegistry
pytestmark = pytest.mark.unit
# 预期注册的 8 个 handler 名称(与各 handler 类的 ``name`` 类变量对齐)
EXPECTED_HANDLER_NAMES = {
"channel_session_inactive_cleanup",
"channel_outbox_terminal_cleanup",
"channel_pairing_terminal_cleanup",
"channel_outbox_recovery",
"channel_pairing_expiration",
"channel_audit_log_retention",
"channel_content_review_retention",
"channel_idempotency_cleanup",
}
# handler 类名 → 预期 name 映射
_HANDLER_CLASS_TO_NAME = {
"ChannelSessionInactiveCleanupHandler": "channel_session_inactive_cleanup",
"ChannelOutboxTerminalCleanupHandler": "channel_outbox_terminal_cleanup",
"ChannelPairingTerminalCleanupHandler": "channel_pairing_terminal_cleanup",
"ChannelOutboxRecoveryHandler": "channel_outbox_recovery",
"ChannelPairingExpirationHandler": "channel_pairing_expiration",
"ChannelAuditLogRetentionHandler": "channel_audit_log_retention",
"ChannelContentReviewRetentionHandler": "channel_content_review_retention",
"ChannelIdempotencyCleanupHandler": "channel_idempotency_cleanup",
}
# 工厂函数名 → 对应 handler name
_FACTORY_TO_NAME = {
"create_channel_session_inactive_cleanup_handler_dependencies": "channel_session_inactive_cleanup",
"create_channel_outbox_terminal_cleanup_handler_dependencies": "channel_outbox_terminal_cleanup",
"create_channel_pairing_terminal_cleanup_handler_dependencies": "channel_pairing_terminal_cleanup",
"create_channel_outbox_recovery_handler_dependencies": "channel_outbox_recovery",
"create_channel_pairing_expiration_handler_dependencies": "channel_pairing_expiration",
"create_channel_audit_log_retention_handler_dependencies": "channel_audit_log_retention",
"create_channel_content_review_retention_handler_dependencies": "channel_content_review_retention",
"create_channel_idempotency_cleanup_handler_dependencies": "channel_idempotency_cleanup",
}
def _make_handler_mock(name: str) -> MagicMock:
"""构造带 ``name`` 属性的 handler 桩。
``MagicMock`` ``name`` 参数是构造参数用于 repr不能直接作为属性
读取因此先创建 mock 再显式赋值 ``name`` 属性
"""
mock = MagicMock()
mock.name = name
return mock
class _PatchTracker:
"""聚合 scheduler 模块桩的追踪数据。"""
def __init__(self) -> None:
self.session_factory: MagicMock = MagicMock(name="session_factory")
self.arq_pool: MagicMock = MagicMock(name="arq_pool")
self.factory_calls: dict[str, list[dict]] = {n: [] for n in EXPECTED_HANDLER_NAMES}
self.handler_class_mocks: dict[str, MagicMock] = {}
@pytest.fixture
def patched_scheduler(monkeypatch):
"""替换 ``scheduler`` 模块的全部外部依赖,返回 ``_PatchTracker``。
替换内容
- ``pg_manager.get_async_session_context``返回 sentinel mock
- ``get_arq_pool``AsyncMock 返回 sentinel mock
- 8 ``create_channel_*_handler_dependencies`` 工厂返回空字典
调用参数记录到 ``factory_calls``
- 8 handler ``side_effect`` 返回带正确 ``name`` 属性的 mock
"""
tracker = _PatchTracker()
# 替换 pg_manager.get_async_session_context
monkeypatch.setattr(
scheduler_module.pg_manager,
"get_async_session_context",
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
for class_name, handler_name in _HANDLER_CLASS_TO_NAME.items():
# 闭包捕获 handler_name 避免延迟绑定
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
# ---------------------------------------------------------------------------
# register_scheduler_handlers
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestRegisterSchedulerHandlers:
"""register_scheduler_handlers 注册行为测试。"""
@pytest.mark.asyncio
async def test_registers_all_8_handlers(self, patched_scheduler):
"""注册后 ``HandlerRegistry`` 包含全部 8 个 handler。"""
# Arrange
registry = HandlerRegistry()
# Act
await scheduler_module.register_scheduler_handlers(registry)
# Assert: 8 个 handler 全部注册
registered = set(registry.list_names())
assert registered == EXPECTED_HANDLER_NAMES
assert len(registered) == 8
@pytest.mark.asyncio
async def test_each_handler_has_correct_name(self, patched_scheduler):
"""每个注册的 handler ``name`` 属性与预期一致。"""
# Arrange
registry = HandlerRegistry()
# Act
await scheduler_module.register_scheduler_handlers(registry)
# Assert: 通过 registry.get(name) 查找每个 handler验证 name 属性
for expected_name in EXPECTED_HANDLER_NAMES:
handler = registry.get(expected_name)
assert handler is not None, f"handler {expected_name} not registered"
assert handler.name == expected_name
@pytest.mark.asyncio
async def test_session_factory_passed_to_all_factories(self, patched_scheduler):
"""``session_factory`` 透传给全部 8 个工厂函数。"""
# Arrange
registry = HandlerRegistry()
# Act
await scheduler_module.register_scheduler_handlers(registry)
# 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 factory_calls[0]["args"][0] is patched_scheduler.session_factory
@pytest.mark.asyncio
async def test_arq_pool_only_passed_to_outbox_recovery(self, patched_scheduler):
"""``arq_pool`` 仅传给 ``outbox_recovery`` 工厂,其余 7 个不传。"""
# Arrange
registry = HandlerRegistry()
# Act
await scheduler_module.register_scheduler_handlers(registry)
# Assert: outbox_recovery 工厂接收 arq_pool 关键字参数
outbox_recovery_calls = patched_scheduler.factory_calls["channel_outbox_recovery"]
assert len(outbox_recovery_calls) == 1
assert outbox_recovery_calls[0]["kwargs"].get("arq_pool") is patched_scheduler.arq_pool
# Assert: 其余 7 个工厂不接收 arq_pool
other_names = EXPECTED_HANDLER_NAMES - {"channel_outbox_recovery"}
for handler_name in other_names:
factory_calls = patched_scheduler.factory_calls[handler_name]
assert len(factory_calls) == 1
assert "arq_pool" not in factory_calls[0]["kwargs"], (
f"factory for {handler_name} should not receive arq_pool"
)
@pytest.mark.asyncio
async def test_get_arq_pool_called_once(self, patched_scheduler, monkeypatch):
"""``get_arq_pool`` 在注册期间仅调用一次。"""
# Arrange: 在已 patch 的基础上再包装一层计数
call_count = 0
original_get_arq_pool = scheduler_module.get_arq_pool
async def _counting_get_arq_pool():
nonlocal call_count
call_count += 1
return await original_get_arq_pool()
monkeypatch.setattr(scheduler_module, "get_arq_pool", _counting_get_arq_pool)
registry = HandlerRegistry()
# Act
await scheduler_module.register_scheduler_handlers(registry)
# Assert
assert call_count == 1
@pytest.mark.asyncio
async def test_handler_classes_instantiated_with_factory_deps(self, patched_scheduler):
"""每个 handler 类被实例化一次,参数为工厂返回的空字典展开。"""
# Arrange
registry = HandlerRegistry()
# Act
await scheduler_module.register_scheduler_handlers(registry)
# 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"
)
# 工厂返回空字典,所以实例化参数为空 kwargs
call_kwargs = class_mock.call_args.kwargs
assert call_kwargs == {}
@pytest.mark.asyncio
async def test_registry_has_no_duplicate_names(self, patched_scheduler):
"""注册后无重名 handlerHandlerRegistry 重复注册抛 DuplicateHandlerError"""
# Arrange
registry = HandlerRegistry()
# Act: 不抛异常即证明 8 个 handler 名称互不重复
await scheduler_module.register_scheduler_handlers(registry)
# Assert: list_names 返回 8 个唯一名称
names = registry.list_names()
assert len(names) == len(set(names)) == 8