"""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): """注册后无重名 handler(HandlerRegistry 重复注册抛 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