ForcePilot/backend/test/unit/channel/plugins/test_protocol.py

530 lines
25 KiB
Python
Raw Normal View History

2026-07-15 12:30:58 +08:00
from __future__ import annotations
from typing import Any
import pytest
from yuxi.channel.exceptions import ChannelErrorClassification
from yuxi.channel.message.models import RunAcceptedResponse
from yuxi.channel.plugins.builders import _BaseChannelPlugin
from yuxi.channel.plugins.protocol import (
BindingConversationRef,
BindingRoute,
ChannelCapability,
ChannelHealthStatus,
ChannelMeta,
ChannelPlugin,
DeliveryCapabilities,
DeliveryMode,
InboundMedia,
InboundMessage,
InboundRequest,
OutboundMessage,
QRLoginSession,
SessionConversationRef,
Transport,
TransportState,
TransportType,
)
from yuxi.channel.ports import (
ConfigPort,
InboundPort,
LifecyclePort,
MetaPort,
OutboundPort,
SecurityPort,
SessionPort,
StatusPort,
TransportPort,
)
from yuxi.channel.security.models import DmPolicy, GroupPolicy, RateLimitPolicy
from yuxi.channel.security.policy import SecurityCheckResult
class TestEnumsAndCapabilities:
def test_channel_capability_flags(self):
caps = ChannelCapability.TEXT | ChannelCapability.MARKDOWN
assert ChannelCapability.TEXT in caps
assert ChannelCapability.IMAGE not in caps
def test_delivery_mode_values(self):
assert DeliveryMode.DIRECT == "direct"
assert DeliveryMode.GATEWAY == "gateway"
assert DeliveryMode.HYBRID == "hybrid"
def test_transport_type_values(self):
assert TransportType.WEBHOOK == "webhook"
assert TransportType.WEBSOCKET == "websocket"
assert TransportType.POLLING == "polling"
def test_transport_state_values(self):
assert TransportState.CONNECTED == "connected"
assert TransportState.DISCONNECTED == "disconnected"
def test_channel_meta_defaults(self):
meta = ChannelMeta(channel_type="test", display_name="Test")
assert meta.aliases == []
assert meta.capabilities == ChannelCapability.TEXT
assert meta.delivery_mode == DeliveryMode.DIRECT
assert meta.transport_type == TransportType.WEBHOOK
assert meta.config_schema == {}
assert meta.ui_hints == {}
assert meta.sort_weight == 0
assert meta.icon is None
assert meta.description == ""
def test_channel_meta_capability_matrix_defaults(self):
from yuxi.channel.capabilities import CapabilityMatrix
meta = ChannelMeta(channel_type="test", display_name="Test")
assert meta.capability_matrix is not None
assert isinstance(meta.capability_matrix, CapabilityMatrix)
assert meta.capability_matrix.text is True
assert meta.capability_matrix.markdown.value == "none"
def test_inbound_media_defaults(self):
media = InboundMedia(media_type="image")
assert media.url is None
assert media.file_key is None
assert media.file_name is None
assert media.size is None
assert media.mime_type is None
def test_outbound_message_defaults(self):
msg = OutboundMessage(content="hello")
assert msg.content_type == "text"
assert msg.media == []
assert msg.reply_to_channel_message_id is None
assert msg.thread_id is None
assert msg.extra == {}
def test_channel_health_status_defaults(self):
status = ChannelHealthStatus(healthy=True, state="ok", enabled=True)
assert status.last_connected_at is None
assert status.last_message_at is None
assert status.last_error is None
assert status.reconnect_attempts == 0
def test_delivery_capabilities_defaults(self):
caps = DeliveryCapabilities()
assert caps.max_text_length == 4000
assert caps.supports_markdown is False
assert caps.supports_interactive is False
assert caps.supports_media is False
class TestRuntimeCheckableProtocols:
def test_inbound_request_protocol_with_fake_request(self):
from yuxi.channel.common.fake_request import FakeRequest
request = FakeRequest(config={}, body={})
assert isinstance(request, InboundRequest)
def test_channel_plugin_protocol_with_port_based_minimal_implementation(self):
class MinimalPlugin(
MetaPort,
ConfigPort,
InboundPort,
OutboundPort,
TransportPort,
SecurityPort,
StatusPort,
LifecyclePort,
SessionPort,
):
def get_meta(self) -> ChannelMeta:
return ChannelMeta(channel_type="test", display_name="Test")
def validate_config(self, config: dict) -> tuple[bool, list[str]]:
return True, []
def list_account_ids(self, config: dict) -> list[str]:
return []
def resolve_account(self, config: dict, account_id: str | None = None) -> dict:
return config
async def validate_webhook(self, request: InboundRequest) -> bool:
return True
async def preprocess_event(
self, request: InboundRequest, config: dict, account_id: str
) -> tuple[bool, dict[str, Any] | None, dict | None]:
return True, None, {}
async def normalize_inbound(self, request: InboundRequest) -> InboundMessage:
return InboundMessage(channel_type="test", account_id="acc")
async def build_webhook_response(
self, request: InboundRequest, inbound: InboundMessage
) -> dict[str, Any] | None:
return None
async def resolve_account_id_from_webhook(self, request: InboundRequest) -> str | None:
return None
async def format_outbound(self, message: OutboundMessage, config: dict | None = None) -> dict:
return {}
async def send_message(
self, channel_session_id: str, payload: dict, config: dict | None = None
) -> str | None:
return None
def chunk_text(self, text: str, limit: int) -> list[str]:
return []
def get_delivery_capabilities(self) -> DeliveryCapabilities:
return DeliveryCapabilities()
async def enrich_outbound(
self, payload: dict, message: OutboundMessage, config: dict | None = None
) -> dict:
return payload
async def upload_media(self, media: InboundMedia | dict, config: dict | None = None) -> dict | None:
return None
def supports_batch_send(self, config: dict | None = None) -> bool:
return False
async def send_batch(
self, channel_session_id: str, payloads: list[dict], config: dict | None = None
) -> list[str | None]:
return []
async def send_typing_indicator(self, channel_session_id: str, config: dict | None = None) -> None:
return None
def supports_content_type(self, content_type: str, config: dict | None = None) -> bool:
return False
async def create_transport(
self, config: dict, account_id: str, credentials: dict | None = None
) -> Transport | None:
return None
async def on_transport_message(self, raw: bytes, config: dict, account_id: str) -> InboundMessage | None:
return None
def resolve_dm_policy(self, config: dict, account_id: str | None = None) -> DmPolicy | None:
return None
def resolve_group_policy(self, config: dict, account_id: str | None = None) -> GroupPolicy | None:
return None
def resolve_rate_limit_policy(self, config: dict, account_id: str | None = None) -> RateLimitPolicy | None:
return None
async def health_check(self, config: dict, account_id: str) -> ChannelHealthStatus:
return ChannelHealthStatus(healthy=True, state="ok", enabled=True)
async def on_config_changed(self, config: dict, account_id: str) -> None:
return None
async def on_account_removed(self, account_id: str) -> None:
return None
async def on_channel_enabled(self, config: dict, account_id: str) -> None:
return None
async def on_channel_disabled(self, config: dict, account_id: str) -> None:
return None
async def run_startup_maintenance(self, config: dict, account_id: str) -> None:
return None
def resolve_session_conversation(self, chat_type: str, raw_id: str) -> SessionConversationRef | None:
return None
def parse_session_key(self, inbound: InboundMessage) -> str:
return "key"
assert isinstance(MinimalPlugin(), ChannelPlugin)
def test_channel_plugin_protocol_backward_compatibility(self):
class LegacyMinimalPlugin:
def get_meta(self) -> ChannelMeta: ...
def validate_config(self, config: dict) -> tuple[bool, list[str]]: ...
async def create_transport(self, config: dict, account_id: str) -> Transport | None: ...
async def on_transport_message(
self, raw: bytes, config: dict, account_id: str
) -> InboundMessage | None: ...
async def validate_webhook(self, request: InboundRequest) -> bool: ...
async def preprocess_event(
self, request: InboundRequest, config: dict, account_id: str
) -> tuple[bool, dict[str, Any] | None, dict | None]: ...
async def normalize_inbound(self, request: InboundRequest) -> InboundMessage: ...
async def build_webhook_response(
self, request: InboundRequest, inbound: InboundMessage
) -> dict[str, Any] | None: ...
async def resolve_account_id_from_webhook(self, request: InboundRequest) -> str | None: ...
async def format_outbound(self, message: OutboundMessage, config: dict | None = None) -> dict: ...
async def send_message(
self, channel_session_id: str, payload: dict, config: dict | None = None
) -> str | None: ...
def chunk_text(self, text: str, limit: int) -> list[str]: ...
def get_delivery_capabilities(self) -> DeliveryCapabilities: ...
def resolve_session_conversation(self, chat_type: str, raw_id: str) -> SessionConversationRef | None: ...
def parse_session_key(self, inbound: InboundMessage) -> str: ...
def list_account_ids(self, config: dict) -> list[str]: ...
def resolve_account(self, config: dict, account_id: str | None = None) -> dict: ...
def resolve_dm_policy(self, config: dict, account_id: str | None = None) -> DmPolicy | None: ...
def resolve_group_policy(self, config: dict, account_id: str | None = None) -> GroupPolicy | None: ...
def resolve_rate_limit_policy(
self, config: dict, account_id: str | None = None
) -> RateLimitPolicy | None: ...
def compile_binding(self, config: dict, inbound: InboundMessage) -> Any | None: ...
def match_binding(self, config: dict, binding: Any, inbound: InboundMessage) -> bool: ...
def resolve_reply_to_mode(self, config: dict, inbound: InboundMessage) -> str | None: ...
def resolve_auto_thread_id(self, config: dict, inbound: InboundMessage) -> str | None: ...
def describe_actions(self, config: dict, inbound: InboundMessage) -> list[dict]: ...
def handle_action(self, action: dict, config: dict, inbound: InboundMessage) -> dict | None: ...
async def on_config_changed(self, config: dict, account_id: str) -> None: ...
async def on_account_removed(self, account_id: str) -> None: ...
async def on_channel_enabled(self, config: dict, account_id: str) -> None: ...
async def on_channel_disabled(self, config: dict, account_id: str) -> None: ...
async def run_startup_maintenance(self, config: dict, account_id: str) -> None: ...
def classify_error(
self, exc: Exception, payload: dict | None = None
) -> tuple[ChannelErrorClassification, int | None]: ...
async def upload_media(self, media: InboundMedia | dict, config: dict | None = None) -> dict | None: ...
async def send_typing_indicator(self, channel_session_id: str, config: dict | None = None) -> None: ...
def supports_content_type(self, content_type: str, config: dict | None = None) -> bool: ...
async def setup_webhook(self, config: dict, callback_url: str) -> bool: ...
async def delete_webhook(self, config: dict) -> bool: ...
async def edit_message(
self, channel_session_id: str, channel_message_id: str, new_payload: dict, config: dict | None = None
) -> bool: ...
async def delete_message(
self, channel_session_id: str, channel_message_id: str, config: dict | None = None
) -> bool: ...
def supports_batch_send(self, config: dict | None = None) -> bool: ...
async def send_batch(
self, channel_session_id: str, payloads: list[dict], config: dict | None = None
) -> list[str | None]: ...
async def resolve_user_profile(self, sender_id: str, config: dict | None = None) -> dict | None: ...
async def transform_inbound(self, raw_event: dict, config: dict | None = None) -> dict: ...
async def enrich_outbound(
self, payload: dict, message: OutboundMessage, config: dict | None = None
) -> dict: ...
async def health_check(self, config: dict, account_id: str) -> ChannelHealthStatus: ...
def supports_qr_login(self, config: dict | None = None) -> bool: ...
async def create_qr_login_session(self, config: dict, account_id: str) -> QRLoginSession: ...
async def poll_qr_login_status(
self, session: QRLoginSession, config: dict, account_id: str
) -> QRLoginSession: ...
async def load_qr_credentials(self, config: dict, account_id: str) -> QRLoginSession | None: ...
async def logout_qr_login(self, config: dict, account_id: str) -> bool: ...
def supports_scan_pairing(self, config: dict | None = None) -> bool: ...
async def normalize_scan_event(
self, request: Any, config: dict, account_id: str
) -> InboundMessage | None: ...
async def build_pairing_qr_reply(
self,
pairing_code: str,
qr_content: str,
config: dict,
account_id: str,
) -> OutboundMessage: ...
async def download_attachment(
self, inbound: InboundMessage, media: InboundMedia
) -> tuple[bytes, str] | None: ...
assert isinstance(LegacyMinimalPlugin(), ChannelPlugin)
def test_inbound_message_protocol_uses_dataclass_fields(self):
msg = InboundMessage(channel_type="test", account_id="acc")
assert msg.content == ""
assert msg.content_type == "text"
assert msg.media == []
assert msg.is_at_bot is False
assert msg.timestamp is None
def test_binding_route_dataclass(self):
route = BindingRoute(agent_id="agent", session_key="key", matched_by="rule")
assert route.binding_rule_hash is None
def test_binding_conversation_ref_defaults(self):
ref = BindingConversationRef(session_key="key", channel_type="test", account_id="acc")
assert ref.chat_type is None
assert ref.peer_id is None
assert ref.thread_id is None
assert ref.binding_rule_hash == ""
def test_session_conversation_ref_defaults(self):
ref = SessionConversationRef(session_key="key")
assert ref.chat_type is None
assert ref.channel_sender_id is None
assert ref.parent_conversation_candidates == []
assert ref.channel_metadata == {}
class TestScanPairingExtension:
class _ChannelPluginStub:
def get_meta(self) -> ChannelMeta: ...
def validate_config(self, config: dict) -> tuple[bool, list[str]]: ...
async def create_transport(
self, config: dict, account_id: str, credentials: dict | None = None
) -> Transport | None: ...
async def on_transport_message(self, raw: bytes, config: dict, account_id: str) -> InboundMessage | None: ...
async def validate_webhook(self, request: InboundRequest) -> bool: ...
async def preprocess_event(
self, request: InboundRequest, config: dict, account_id: str
) -> tuple[bool, dict[str, Any] | None, dict | None]: ...
async def normalize_inbound(self, request: InboundRequest) -> InboundMessage: ...
async def build_webhook_response(
self, request: InboundRequest, inbound: InboundMessage
) -> dict[str, Any] | None: ...
async def resolve_account_id_from_webhook(self, request: InboundRequest) -> str | None: ...
async def format_outbound(self, message: OutboundMessage, config: dict | None = None) -> dict: ...
async def send_message(
self, channel_session_id: str, payload: dict, config: dict | None = None
) -> str | None: ...
def chunk_text(self, text: str, limit: int) -> list[str]: ...
def get_delivery_capabilities(self) -> DeliveryCapabilities: ...
def resolve_session_conversation(self, chat_type: str, raw_id: str) -> SessionConversationRef | None: ...
def parse_session_key(self, inbound: InboundMessage) -> str: ...
def list_account_ids(self, config: dict) -> list[str]: ...
def resolve_account(self, config: dict, account_id: str | None = None) -> dict: ...
def resolve_dm_policy(self, config: dict, account_id: str | None = None) -> DmPolicy | None: ...
def resolve_group_policy(self, config: dict, account_id: str | None = None) -> GroupPolicy | None: ...
def resolve_rate_limit_policy(self, config: dict, account_id: str | None = None) -> RateLimitPolicy | None: ...
def compile_binding(self, config: dict, inbound: InboundMessage) -> Any | None: ...
def match_binding(self, config: dict, binding: Any, inbound: InboundMessage) -> bool: ...
def resolve_reply_to_mode(self, config: dict, inbound: InboundMessage) -> str | None: ...
def resolve_auto_thread_id(self, config: dict, inbound: InboundMessage) -> str | None: ...
def describe_actions(self, config: dict, inbound: InboundMessage) -> list[dict]: ...
def handle_action(self, action: dict, config: dict, inbound: InboundMessage) -> dict | None: ...
async def on_config_changed(self, config: dict, account_id: str) -> None: ...
async def on_account_removed(self, account_id: str) -> None: ...
async def on_channel_enabled(self, config: dict, account_id: str) -> None: ...
async def on_channel_disabled(self, config: dict, account_id: str) -> None: ...
async def run_startup_maintenance(self, config: dict, account_id: str) -> None: ...
def classify_error(
self, exc: Exception, payload: dict | None = None
) -> tuple[ChannelErrorClassification, int | None]: ...
async def upload_media(self, media: InboundMedia | dict, config: dict | None = None) -> dict | None: ...
async def send_typing_indicator(self, channel_session_id: str, config: dict | None = None) -> None: ...
def supports_content_type(self, content_type: str, config: dict | None = None) -> bool: ...
async def setup_webhook(self, config: dict, callback_url: str) -> bool: ...
async def delete_webhook(self, config: dict) -> bool: ...
async def edit_message(
self,
channel_session_id: str,
channel_message_id: str,
new_payload: dict,
config: dict | None = None,
) -> bool: ...
async def delete_message(
self,
channel_session_id: str,
channel_message_id: str,
config: dict | None = None,
) -> bool: ...
def supports_batch_send(self, config: dict | None = None) -> bool: ...
async def send_batch(
self, channel_session_id: str, payloads: list[dict], config: dict | None = None
) -> list[str | None]: ...
async def resolve_user_profile(self, sender_id: str, config: dict | None = None) -> dict | None: ...
async def transform_inbound(self, raw_event: dict, config: dict | None = None) -> dict: ...
async def enrich_outbound(
self, payload: dict, message: OutboundMessage, config: dict | None = None
) -> dict: ...
def supports_qr_login(self, config: dict | None = None) -> bool: ...
async def create_qr_login_session(self, config: dict, account_id: str) -> QRLoginSession: ...
async def poll_qr_login_status(
self, session: QRLoginSession, config: dict, account_id: str
) -> QRLoginSession: ...
async def load_qr_credentials(self, config: dict, account_id: str) -> QRLoginSession | None: ...
async def logout_qr_login(self, config: dict, account_id: str) -> bool: ...
async def health_check(self, config: dict, account_id: str) -> ChannelHealthStatus: ...
async def download_attachment(
self, inbound: InboundMessage, media: InboundMedia
) -> tuple[bytes, str] | None: ...
@pytest.mark.unit
def test_inbound_message_scan_event_default(self):
msg = InboundMessage(channel_type="test", account_id="acc")
assert msg.is_scan_event is False
scan_msg = InboundMessage(channel_type="test", account_id="acc", is_scan_event=True)
assert scan_msg.is_scan_event is True
@pytest.mark.unit
def test_run_accepted_response_scan_pairing_fields(self):
reply = OutboundMessage(content="scan me")
response = RunAcceptedResponse(
accepted=False,
pairing_code="CODE123",
qr_content="https://example.com/bind?code=CODE123",
qr_reply=reply,
)
assert response.qr_content == "https://example.com/bind?code=CODE123"
assert response.qr_reply is reply
@pytest.mark.unit
def test_security_check_result_scan_pairing_fields(self):
reply = OutboundMessage(content="scan me")
result = SecurityCheckResult(
allowed=False,
pairing_code="CODE123",
qr_content="https://example.com/bind?code=CODE123",
qr_reply=reply,
)
assert result.qr_content == "https://example.com/bind?code=CODE123"
assert result.qr_reply is reply
@pytest.mark.unit
def test_channel_plugin_with_scan_pairing_methods(self):
class PluginWithScanPairing(self._ChannelPluginStub):
def supports_scan_pairing(self, config: dict | None = None) -> bool:
return True
async def normalize_scan_event(
self,
request: Any,
config: dict,
account_id: str,
) -> InboundMessage | None:
return None
async def build_pairing_qr_reply(
self,
pairing_code: str,
qr_content: str,
config: dict,
account_id: str,
) -> OutboundMessage:
return OutboundMessage(content=qr_content)
assert isinstance(PluginWithScanPairing(), ChannelPlugin)
@pytest.mark.unit
def test_channel_plugin_without_scan_pairing_methods_passes_isinstance(self):
"""_BaseChannelPlugin 提供扫码配对默认实现,未重写这些方法仍满足 Protocol。"""
class PluginWithoutScanPairing(_BaseChannelPlugin):
pass
meta = ChannelMeta(channel_type="test", display_name="Test")
plugin = PluginWithoutScanPairing(meta=meta, adapters={})
assert isinstance(plugin, ChannelPlugin)
@pytest.mark.unit
async def test_base_channel_plugin_scan_pairing_defaults(self):
meta = ChannelMeta(channel_type="test", display_name="Test")
plugin = _BaseChannelPlugin(meta=meta, adapters={})
assert isinstance(plugin, ChannelPlugin)
assert plugin.supports_scan_pairing(None) is False
assert plugin.supports_scan_pairing({}) is False
assert plugin.supports_scan_pairing({"pairing": {"mode": "code"}}) is False
assert plugin.supports_scan_pairing({"pairing": {"mode": "qr"}}) is True
assert plugin.supports_scan_pairing({"pairing": {"mode": "both"}}) is True
normalized = await plugin.normalize_scan_event(None, {}, "acc")
assert normalized is None
reply = await plugin.build_pairing_qr_reply("CODE123", "qr-content", {}, "acc")
assert isinstance(reply, OutboundMessage)
assert reply.content == "请扫描下方二维码完成绑定CODE123"
assert reply.content_type == "text"