ForcePilot/backend/test/unit/channels/test_yuanbao_enhancements.py
Kris 3264900bc9 test: 新增多渠道单元测试用例并配置测试环境变量
新增了Twitch、Telegram、Discord、Slack、Mattermost、WeChat、Zalo等多渠道的单元测试用例,覆盖了令牌处理、速率限制、消息去重、会话解析、格式转换、安全策略等模块
同时在测试配置中添加了测试用的OpenAI API密钥环境变量
2026-05-12 00:56:47 +08:00

515 lines
18 KiB
Python

from __future__ import annotations
import json
import struct
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from yuxi.channels.adapters.yuanbao.config_ui_hints import (
ConfigFieldHint,
YUANBAO_CONFIG_HINTS,
get_config_hints,
get_config_hints_by_section,
get_config_hint,
get_required_config_keys,
SECTION_LABELS,
)
from yuxi.channels.adapters.yuanbao.dispatch import (
DispatchAction,
InteractiveDispatcher,
DispatchContext,
)
from yuxi.channels.adapters.yuanbao.monitor import YuanbaoMonitor
from yuxi.channels.adapters.yuanbao.outbound_queue import OutboundQueue
from yuxi.channels.adapters.yuanbao.proto_codec import ProtoCodec, is_protobuf_message
from yuxi.channels.adapters.yuanbao.security_audit import SecurityAuditLogger
from yuxi.channels.adapters.yuanbao.send_cache import SendMessageCache, SentMessageEntry
from yuxi.channels.adapters.yuanbao.setup import (
generate_config_snippet,
generate_config_snippet_with_accounts,
)
from yuxi.channels.adapters.yuanbao.template import (
TemplateMessageBuilder,
TemplateButton,
TemplateCard,
ActionSelector,
SelectorOption,
)
from yuxi.channels.models import (
ChannelIdentity,
ChannelResponse,
ChannelType,
DeliveryResult,
)
YUANBAO = ChannelType.YUANBAO
def _make_identity(
channel_chat_id: str = "user_openid_001",
) -> ChannelIdentity:
return ChannelIdentity(
channel_id="yuanbao",
channel_type=YUANBAO,
channel_user_id=channel_chat_id,
channel_chat_id=channel_chat_id,
)
# ==================== OutboundQueue Identity Preservation Tests ====================
class TestOutboundQueueIdentityFix:
@pytest.mark.asyncio
async def test_flush_preserves_identity(self):
sent_responses = []
async def send_fn(response: ChannelResponse) -> DeliveryResult:
sent_responses.append(response)
return DeliveryResult(success=True, message_id="msg_001")
queue = OutboundQueue({"outboundQueueStrategy": "merge-text"}, send_fn)
identity = _make_identity("test_chat")
response = ChannelResponse(identity=identity, content="Hello", metadata={"key": "value"})
await queue.enqueue(response)
result = await queue.flush()
assert result is not None
assert result.success is True
assert len(sent_responses) == 1
assert sent_responses[0].identity.channel_chat_id == "test_chat"
assert sent_responses[0].metadata.get("key") == "value"
assert sent_responses[0].content == "Hello"
@pytest.mark.asyncio
async def test_flush_merges_multiple_messages_with_identity(self):
sent_responses = []
async def send_fn(response: ChannelResponse) -> DeliveryResult:
sent_responses.append(response)
return DeliveryResult(success=True)
queue = OutboundQueue({"outboundQueueStrategy": "merge-text"}, send_fn)
identity = _make_identity("chat_abc")
for text in ["A", "B", "C"]:
await queue.enqueue(ChannelResponse(identity=identity, content=text))
await queue.flush()
assert len(sent_responses) == 1
assert sent_responses[0].identity.channel_chat_id == "chat_abc"
@pytest.mark.asyncio
async def test_immediate_mode_passes_through(self):
sent_responses = []
async def send_fn(response: ChannelResponse) -> DeliveryResult:
sent_responses.append(response)
return DeliveryResult(success=True)
queue = OutboundQueue({"outboundQueueStrategy": "immediate"}, send_fn)
identity = _make_identity("chat_x")
response = ChannelResponse(identity=identity, content="Direct")
result = await queue.enqueue(response)
assert len(sent_responses) == 1
assert sent_responses[0].identity.channel_chat_id == "chat_x"
# ==================== Template Message Tests ====================
class TestTemplateMessageBuilder:
def test_builder_empty(self):
builder = TemplateMessageBuilder()
metadata = builder.build_metadata()
assert metadata == {}
assert builder.build_content() == ""
def test_builder_with_content(self):
builder = TemplateMessageBuilder()
builder.set_content("Hello World")
assert builder.build_content() == "Hello World"
def test_builder_with_buttons(self):
builder = TemplateMessageBuilder()
builder.add_button("Click me", "url", "https://example.com")
builder.add_button("Cancel", "callback", "cancel_action")
metadata = builder.build_metadata()
assert "buttons" in metadata
assert len(metadata["buttons"]) == 2
assert metadata["buttons"][0]["text"] == "Click me"
assert metadata["buttons"][1]["type"] == "callback"
def test_builder_with_card(self):
card = TemplateCard(
title="Test Card",
content="Card content",
buttons=[TemplateButton("OK")],
)
builder = TemplateMessageBuilder()
builder.set_card(card)
metadata = builder.build_metadata()
assert "card" in metadata
assert metadata["card"]["title"] == "Test Card"
def test_builder_with_selector(self):
builder = TemplateMessageBuilder()
builder.add_selector(
placeholder="Choose",
options=[
SelectorOption(label="Option A", value="a"),
SelectorOption(label="Option B", value="b"),
],
selector_id="sel_001",
)
metadata = builder.build_metadata()
assert "selectors" in metadata
assert len(metadata["selectors"]) == 1
assert metadata["selectors"][0]["type"] == "action_select"
assert len(metadata["selectors"][0]["options"]) == 2
def test_builder_combined(self):
builder = TemplateMessageBuilder()
builder.set_content("Welcome")
builder.add_button("Start", action_type="callback", value="start")
card = TemplateCard(title="Info Card", content="Details")
builder.set_card(card)
metadata = builder.build_metadata()
assert "buttons" in metadata
assert "card" in metadata
def test_template_button_to_dict(self):
btn = TemplateButton(text="Hello", action_type="url", value="x")
d = btn.to_dict()
assert d == {"text": "Hello", "type": "url", "value": "x", "style": "default"}
def test_selector_option_to_dict(self):
opt = SelectorOption(label="L", value="V", description="desc")
d = opt.to_dict()
assert d == {"label": "L", "value": "V", "description": "desc"}
def test_builder_reset(self):
builder = TemplateMessageBuilder()
builder.set_content("x")
builder.add_button("B")
builder.reset()
assert builder.build_content() == ""
assert builder.build_metadata() == {}
# ==================== Dispatch Read Receipt Tests ====================
class TestReadReceiptDispatch:
def test_read_receipt_classified_correctly(self):
event = {"type": "read_receipt", "msg_id": "msg_001", "open_id": "user_123"}
action = InteractiveDispatcher.classify_event(event)
assert action == DispatchAction.READ_RECEIPT
def test_typing_still_system_event(self):
event = {"type": "typing", "open_id": "user_123"}
action = InteractiveDispatcher.classify_event(event)
assert action == DispatchAction.SYSTEM_EVENT
def test_system_event_still_system_event(self):
event = {"type": "system_event", "data": {}}
action = InteractiveDispatcher.classify_event(event)
assert action == DispatchAction.SYSTEM_EVENT
# ==================== Monitor Health Tests ====================
class TestMonitorHealth:
def test_healthy_when_not_connected(self):
def dummy_token():
pass
async def dummy_event(event):
pass
monitor = YuanbaoMonitor(
ws_url="wss://test/ws",
token_provider=dummy_token,
on_event=dummy_event,
)
assert not monitor.healthy
def test_reconnect_count_initial(self):
def dummy_token():
pass
async def dummy_event(event):
pass
monitor = YuanbaoMonitor(
ws_url="wss://test/ws",
token_provider=dummy_token,
on_event=dummy_event,
)
assert monitor.reconnect_count == 0
# ==================== Config UI Hints Tests ====================
class TestConfigUiHints:
def test_hints_count(self):
hints = get_config_hints()
assert len(hints) >= 40
def test_hints_by_section(self):
auth_hints = get_config_hints_by_section("authentication")
assert len(auth_hints) > 0
assert any(h.key == "app_key" for h in auth_hints)
def test_get_config_hint(self):
hint = get_config_hint("dm_policy")
assert hint is not None
assert hint.key == "dm_policy"
assert hint.field_type == "select"
assert len(hint.choices) == 4
def test_get_config_hint_not_found(self):
assert get_config_hint("nonexistent_key") is None
def test_required_keys(self):
required = get_required_config_keys()
assert "app_key" in required
assert "app_secret" in required
def test_section_labels(self):
assert "authentication" in SECTION_LABELS
assert "policy" in SECTION_LABELS
assert SECTION_LABELS["authentication"] == "认证配置"
def test_all_hints_have_valid_types(self):
valid_types = {"string", "password", "number", "boolean", "select", "array"}
for hint in YUANBAO_CONFIG_HINTS:
assert hint.field_type in valid_types, f"Invalid type for {hint.key}: {hint.field_type}"
# ==================== Send Cache Redis Tests ====================
class TestSendMessageCache:
def test_add_and_get(self):
cache = SendMessageCache()
cache.add("msg_1", "chat_abc", "Hello")
entry = cache.get("msg_1")
assert entry is not None
assert entry.chat_id == "chat_abc"
assert entry.status == "sent"
def test_mark_edited(self):
cache = SendMessageCache()
cache.add("msg_1", "chat_abc", "Hello")
assert cache.mark_edited("msg_1", "New content")
entry = cache.get("msg_1")
assert entry.status == "edited"
def test_mark_deleted(self):
cache = SendMessageCache()
cache.add("msg_1", "chat_abc", "Hello")
assert cache.mark_deleted("msg_1")
entry = cache.get("msg_1")
assert entry.status == "deleted"
def test_redis_backend_noop(self):
cache = SendMessageCache()
cache.set_redis_backend(None)
cache.add("msg_1", "chat_abc")
entry = cache.get("msg_1")
assert entry is not None
def test_entry_to_dict(self):
entry = SentMessageEntry(msg_id="m1", chat_id="c1", content_preview="Hello", status="sent")
d = entry.to_dict()
assert d["msg_id"] == "m1"
assert d["chat_id"] == "c1"
assert d["status"] == "sent"
def test_entry_from_dict(self):
d = {"msg_id": "m1", "chat_id": "c1", "content_preview": "H", "sent_at": 100.0, "status": "sent"}
entry = SentMessageEntry.from_dict(d)
assert entry.msg_id == "m1"
assert entry.chat_id == "c1"
assert entry.sent_at == 100.0
def test_get_by_status(self):
cache = SendMessageCache()
cache.add("m1", "c1", "A")
cache.add("m2", "c1", "B")
cache.mark_deleted("m2")
deleted = cache.get_by_status("deleted")
assert len(deleted) == 1
assert deleted[0].msg_id == "m2"
def test_update_status_non_existent(self):
cache = SendMessageCache()
assert not cache.update_status("nonexistent", "deleted")
def test_redis_backend_is_settable(self):
cache = SendMessageCache()
mock_redis = MagicMock()
cache.set_redis_backend(mock_redis, prefix="test:", ttl=3600)
assert cache._redis is not None
assert cache._redis_prefix == "test:"
assert cache._redis_ttl == 3600
# ==================== Security Audit JSON Tests ====================
class TestSecurityAuditJson:
def test_set_output_format_json(self):
SecurityAuditLogger.set_output_format("json")
assert SecurityAuditLogger._output_format == "json"
def test_set_output_format_text(self):
SecurityAuditLogger.set_output_format("text")
assert SecurityAuditLogger._output_format == "text"
def test_set_output_format_unknown(self):
SecurityAuditLogger.set_output_format("xml")
assert SecurityAuditLogger._output_format == "xml"
def test_log_dm_blocked_text(self, caplog):
import logging
SecurityAuditLogger.set_output_format("text")
with caplog.at_level(logging.WARNING):
SecurityAuditLogger.log_dm_policy_blocked("user_1", "test_reason", "allowlist")
assert "dm_policy_blocked" in caplog.text
assert "user_1" in caplog.text
def test_log_dm_blocked_json(self, caplog):
import logging
SecurityAuditLogger.set_output_format("json")
with caplog.at_level(logging.WARNING):
SecurityAuditLogger.log_dm_policy_blocked("user_1", "test_reason", "allowlist")
text = caplog.text
assert "SECURITY_AUDIT" in text
parsed = json.loads(text.split("[SECURITY_AUDIT] ", 1)[1])
assert parsed["event_type"] == "dm_policy_blocked"
assert parsed["user_id"] == "user_1"
def test_all_log_methods_exist(self):
methods = [
"log_dm_policy_blocked",
"log_group_access_blocked",
"log_mention_required_blocked",
"log_unauthorized_access",
"log_rate_limit_exceeded",
"log_circuit_breaker_open",
"log_auth_failure",
]
for method in methods:
assert hasattr(SecurityAuditLogger, method)
# ==================== Setup Wizard Multi-Account Tests ====================
class TestSetupWizardMultiAccount:
def test_generate_config_snippet(self):
snippet = generate_config_snippet("app_key_1", "secret_1", "bot_1")
assert "app_key_1" in snippet
assert "bot_1" in snippet
def test_generate_config_snippet_with_accounts(self):
accounts = [
{"id": "acct1", "app_key": "key1", "app_secret": "sec1", "name": "Account 1"},
{"id": "acct2", "app_key": "key2", "app_secret": "sec2"},
]
snippet = generate_config_snippet_with_accounts("app_key_main", "secret_main", "bot_main", accounts)
assert "accounts.acct1" in snippet
assert "appKey: \"key1\"" in snippet
assert "Account 1" in snippet
assert "account_id" not in snippet.split("acct2")[1] # uses "id" as fallback name
def test_generate_config_snippet_without_accounts(self):
snippet = generate_config_snippet_with_accounts("app_key", "secret", "bot", None)
assert "defaultAccount" in snippet
assert "accounts." not in snippet
# ==================== Proto Codec Tests ====================
class TestProtoCodec:
def test_is_protobuf_with_magic(self):
frame = bytes([ProtoCodec.MAGIC_BYTE]) + b"more_data"
assert is_protobuf_message(frame) is True
def test_is_protobuf_without_magic(self):
data = b"\x00hello"
assert is_protobuf_message(data) is False
def test_is_protobuf_empty(self):
assert is_protobuf_message(b"") is False
def test_encode_biz(self):
payload = {"key": "value"}
frame = ProtoCodec.encode_biz(payload)
assert isinstance(frame, bytes)
assert frame[0] == ProtoCodec.MAGIC_BYTE
assert frame[1] == ProtoCodec.VERSION
decoded = ProtoCodec.decode(frame)
assert decoded is not None
assert decoded["payload"] == payload
def test_encode_heartbeat(self):
frame = ProtoCodec.encode_heartbeat()
assert isinstance(frame, bytes)
decoded = ProtoCodec.decode(frame)
assert decoded is not None
assert decoded["msg_type"] == "heartbeat"
def test_encode_auth(self):
frame = ProtoCodec.encode_auth("test_token")
decoded = ProtoCodec.decode(frame)
assert decoded is not None
assert decoded["msg_type"] == "auth"
assert decoded["payload"]["token"] == "test_token"
def test_decode_invalid_magic(self):
result = ProtoCodec.decode(b"\x00\x01\x02\x03\x04")
assert result is None
def test_decode_too_short(self):
result = ProtoCodec.decode(b"\xfe")
assert result is None
def test_build_frame_body_too_large(self):
with pytest.raises(ValueError, match="Body too large"):
ProtoCodec._build_frame(ProtoCodec.MSG_TYPE_BIZ, b"x" * 65536)
def test_msg_type_names(self):
assert ProtoCodec._msg_type_name(1) == "biz"
assert ProtoCodec._msg_type_name(2) == "conn"
assert ProtoCodec._msg_type_name(3) == "heartbeat"
assert ProtoCodec._msg_type_name(4) == "auth"
assert ProtoCodec._msg_type_name(99) == "unknown"
def test_encode_decode_roundtrip(self):
payload = {"text": "Hello World", "user_id": "user_123"}
frame = ProtoCodec.encode_biz(payload)
decoded = ProtoCodec.decode(frame)
assert decoded["msg_type"] == "biz"
assert decoded["payload"]["text"] == "Hello World"
assert decoded["payload"]["user_id"] == "user_123"
def test_conn_encode(self):
payload = {"ws_url": "wss://test/ws"}
frame = ProtoCodec.encode_conn(payload)
decoded = ProtoCodec.decode(frame)
assert decoded["msg_type"] == "conn"