from __future__ import annotations from unittest.mock import AsyncMock, MagicMock, patch import pytest from yuxi.channel.exceptions import ( ChannelErrorClassification, ChannelPermanentError, ChannelRateLimitedError, ChannelRetryableError, ChannelValidationError, ) from yuxi.channel.plugins.builders import _BaseChannelPlugin, create_channel_plugin_base, define_channel_plugin_entry from yuxi.channel.plugins.protocol import ( ChannelCapability, ChannelMeta, ChannelPlugin, DeliveryCapabilities, DeliveryMode, InboundMessage, OutboundMessage, TransportType, ) from yuxi.channel.ports import ( AuthPort, BindingPort, ConfigPort, InboundPort, LifecyclePort, MetaPort, OutboundPort, SecurityPort, SessionPort, StatusPort, ThreadingPort, TransportPort, ) class TestCreateChannelPluginBase: def test_returns_channel_plugin_instance(self): plugin = create_channel_plugin_base( channel_type="test", display_name="Test", config_schema={"type": "object"}, ) assert isinstance(plugin, ChannelPlugin) def test_meta_fields_populated(self): plugin = create_channel_plugin_base( channel_type="test", display_name="Test", config_schema={}, aliases=["t"], capabilities=ChannelCapability.TEXT | ChannelCapability.IMAGE, delivery_mode=DeliveryMode.HYBRID, transport_type=TransportType.POLLING, sort_weight=5, icon="icon.svg", description="desc", ) meta = plugin.get_meta() assert meta.channel_type == "test" assert meta.display_name == "Test" assert meta.aliases == ["t"] assert meta.capabilities == (ChannelCapability.TEXT | ChannelCapability.IMAGE) assert meta.delivery_mode == DeliveryMode.HYBRID assert meta.transport_type == TransportType.POLLING assert meta.sort_weight == 5 assert meta.icon == "icon.svg" assert meta.description == "desc" def test_adapter_methods_override_defaults(self): async def custom_normalize(request): return InboundMessage(channel_type="test", account_id="acc") plugin = create_channel_plugin_base( channel_type="test", display_name="Test", config_schema={}, normalize_inbound=custom_normalize, ) assert plugin.normalize_inbound is custom_normalize def test_non_callable_adapters_are_ignored(self): plugin = create_channel_plugin_base( channel_type="test", display_name="Test", config_schema={}, not_a_method=123, ) assert not hasattr(plugin, "not_a_method") class TestBasePluginDefaults: @pytest.fixture def plugin(self): return create_channel_plugin_base( channel_type="test", display_name="Test", config_schema={}, ) def test_list_account_ids_with_explicit_id(self, plugin): assert plugin.list_account_ids({"account_id": "acc"}) == ["acc"] def test_list_account_ids_empty_when_missing(self, plugin): assert plugin.list_account_ids({}) == [] def test_resolve_account_returns_config_for_known_account(self, plugin): config = {"account_id": "acc"} assert plugin.resolve_account(config) is config def test_resolve_account_raises_for_unknown_account(self, plugin): config = {"account_id": "acc"} with pytest.raises(ValueError, match="Unknown account"): plugin.resolve_account(config, account_id="other") def test_validate_config_default(self, plugin): assert plugin.validate_config({}) == (True, []) def test_chunk_text_short(self, plugin): assert plugin.chunk_text("hello", 10) == ["hello"] def test_chunk_text_splits(self, plugin): assert plugin.chunk_text("abcdef", 2) == ["ab", "cd", "ef"] def test_get_delivery_capabilities_default(self, plugin): caps = plugin.get_delivery_capabilities() assert isinstance(caps, DeliveryCapabilities) assert caps.supports_markdown is False def test_parse_session_key(self, plugin): inbound = InboundMessage( channel_type="test", account_id="acc", chat_type="private", peer_id="peer", ) assert plugin.parse_session_key(inbound) == "test:acc:private:peer" def test_parse_session_key_with_thread(self, plugin): inbound = InboundMessage( channel_type="test", account_id="acc", chat_type="private", peer_id="peer", thread_id="thread", ) assert plugin.parse_session_key(inbound) == "test:acc:private:peer:thread" def test_parse_session_key_omits_none_parts(self, plugin): inbound = InboundMessage(channel_type="test", account_id="acc") assert plugin.parse_session_key(inbound) == "test:acc" async def test_transform_inbound_passes_through(self, plugin): event = {"a": 1} assert await plugin.transform_inbound(event) is event async def test_format_outbound_default(self, plugin): msg = OutboundMessage(content="hello") assert await plugin.format_outbound(msg) == {"content": "hello"} async def test_send_message_not_implemented(self, plugin): with pytest.raises(NotImplementedError): await plugin.send_message("session", {}, {}) async def test_normalize_inbound_not_implemented(self, plugin): from yuxi.channel.common.fake_request import FakeRequest request = FakeRequest(config={}, body={}) with pytest.raises(NotImplementedError): await plugin.normalize_inbound(request) def test_resolve_dm_policy(self, plugin): policy = plugin.resolve_dm_policy( {"dm_policy": {"mode": "allow_from", "allow_list": ["u1"], "deny_list": ["u2"]}} ) assert policy.mode == "allow_from" assert policy.allow_list == ["u1"] assert policy.deny_list == ["u2"] def test_resolve_dm_policy_defaults(self, plugin): policy = plugin.resolve_dm_policy({}) assert policy.mode == "open" assert policy.allow_list == [] assert policy.deny_list == [] def test_resolve_group_policy(self, plugin): policy = plugin.resolve_group_policy( {"group_policy": {"require_mention": True, "allow_groups": ["g1"], "deny_groups": ["g2"]}} ) assert policy.require_mention is True assert policy.allow_groups == ["g1"] assert policy.deny_groups == ["g2"] def test_resolve_rate_limit_policy(self, plugin): policy = plugin.resolve_rate_limit_policy( {"rate_limit": {"max_requests_per_minute": 120, "max_concurrent_agent_runs": 10}} ) assert policy.max_requests_per_minute == 120 assert policy.max_concurrent_agent_runs == 10 def test_resolve_rate_limit_policy_returns_none_when_missing(self, plugin): assert plugin.resolve_rate_limit_policy({}) is None def test_supports_content_type_text(self, plugin): assert plugin.supports_content_type("text") is True def test_supports_content_type_defaults(self, plugin): assert plugin.supports_content_type("markdown") is False assert plugin.supports_content_type("image") is False assert plugin.supports_content_type("interactive") is False assert plugin.supports_content_type("unknown") is False def test_supports_content_type_with_capabilities(self): plugin = create_channel_plugin_base( channel_type="test", display_name="Test", config_schema={}, capabilities=ChannelCapability.MARKDOWN | ChannelCapability.IMAGE | ChannelCapability.INTERACTIVE, ) # capabilities does not directly affect get_delivery_capabilities default assert plugin.supports_content_type("markdown") is False plugin.get_delivery_capabilities = lambda: DeliveryCapabilities( supports_markdown=True, supports_media=True, supports_interactive=True ) assert plugin.supports_content_type("markdown") is True assert plugin.supports_content_type("image") is True assert plugin.supports_content_type("interactive") is True def test_classify_error_validation(self, plugin): result = plugin.classify_error(ChannelValidationError("bad")) assert result == (ChannelErrorClassification.PERMANENT, None) def test_classify_error_permanent(self, plugin): result = plugin.classify_error(ChannelPermanentError("bad")) assert result == (ChannelErrorClassification.PERMANENT, None) def test_classify_error_rate_limited(self, plugin): exc = ChannelRateLimitedError("slow", retry_after=30) result = plugin.classify_error(exc) assert result == (ChannelErrorClassification.RATE_LIMITED, 30) def test_classify_error_retryable(self, plugin): result = plugin.classify_error(ChannelRetryableError("try again")) assert result == (ChannelErrorClassification.RETRYABLE, None) def test_classify_error_unknown(self, plugin): result = plugin.classify_error(RuntimeError("unknown")) assert result == (ChannelErrorClassification.RETRYABLE, None) async def test_send_batch_fallback(self, plugin): plugin.send_message = AsyncMock(side_effect=["id-1", "id-2"]) results = await plugin.send_batch("session", [{"a": 1}, {"b": 2}]) assert results == ["id-1", "id-2"] assert plugin.send_message.await_count == 2 async def test_send_batch_catches_exceptions(self, plugin): plugin.send_message = AsyncMock(side_effect=["id-1", RuntimeError("boom")]) results = await plugin.send_batch("session", [{"a": 1}, {"b": 2}]) assert results == ["id-1", None] async def test_health_check_default(self, plugin): status = await plugin.health_check({"enabled": True}, "acc") assert status.healthy is True assert status.state == "unknown" assert status.enabled is True async def test_lifecycle_hooks_noop(self, plugin): await plugin.on_config_changed({}, "acc") await plugin.on_account_removed("acc") await plugin.on_channel_enabled({}, "acc") await plugin.on_channel_disabled({}, "acc") await plugin.run_startup_maintenance({}, "acc") async def test_webhook_management_defaults(self, plugin): assert await plugin.setup_webhook({}, "url") is False assert await plugin.delete_webhook({}) is False async def test_edit_and_delete_message_defaults(self, plugin): assert await plugin.edit_message("s", "m", {}) is False assert await plugin.delete_message("s", "m") is False class TestDefineChannelPluginEntry: def test_registers_plugin_in_full_mode(self): plugin = MagicMock() registry = MagicMock() with patch("yuxi.channel.plugins.builders.get_registry", return_value=registry): define_channel_plugin_entry(plugin, register_mode="full") registry.register.assert_called_once_with(plugin) registry.register_primary.assert_not_called() def test_registers_plugin_primary_in_discovery_mode(self): plugin = MagicMock() registry = MagicMock() with patch("yuxi.channel.plugins.builders.get_registry", return_value=registry): define_channel_plugin_entry(plugin, register_mode="discovery") registry.register_primary.assert_called_once_with(plugin) registry.register.assert_not_called() def test_unknown_register_mode_defaults_to_full(self): plugin = MagicMock() registry = MagicMock() with patch("yuxi.channel.plugins.builders.get_registry", return_value=registry): define_channel_plugin_entry(plugin, register_mode="unknown") registry.register.assert_called_once_with(plugin) class TestBaseChannelPluginPorts: @pytest.fixture def plugin(self): meta = ChannelMeta(channel_type="test", display_name="Test") return _BaseChannelPlugin(meta=meta, adapters={}) def test_is_instance_of_channel_plugin(self, plugin): assert isinstance(plugin, ChannelPlugin) def test_is_instance_of_meta_port(self, plugin): assert isinstance(plugin, MetaPort) def test_is_instance_of_config_port(self, plugin): assert isinstance(plugin, ConfigPort) def test_is_instance_of_inbound_port(self, plugin): assert isinstance(plugin, InboundPort) def test_is_instance_of_outbound_port(self, plugin): assert isinstance(plugin, OutboundPort) def test_is_instance_of_transport_port(self, plugin): assert isinstance(plugin, TransportPort) def test_is_instance_of_security_port(self, plugin): assert isinstance(plugin, SecurityPort) def test_is_instance_of_status_port(self, plugin): assert isinstance(plugin, StatusPort) def test_is_instance_of_lifecycle_port(self, plugin): assert isinstance(plugin, LifecyclePort) def test_is_instance_of_session_port(self, plugin): assert isinstance(plugin, SessionPort) def test_is_instance_of_auth_port(self, plugin): assert isinstance(plugin, AuthPort) def test_is_instance_of_binding_port(self, plugin): assert isinstance(plugin, BindingPort) def test_is_instance_of_threading_port(self, plugin): assert isinstance(plugin, ThreadingPort) def test_adapter_override_still_works(self): async def custom_normalize(request): return InboundMessage(channel_type="test", account_id="acc") meta = ChannelMeta(channel_type="test", display_name="Test") plugin = _BaseChannelPlugin(meta=meta, adapters={"normalize_inbound": custom_normalize}) assert plugin.normalize_inbound is custom_normalize