363 lines
14 KiB
Python
363 lines
14 KiB
Python
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
|