新增了Twitch、Telegram、Discord、Slack、Mattermost、WeChat、Zalo等多渠道的单元测试用例,覆盖了令牌处理、速率限制、消息去重、会话解析、格式转换、安全策略等模块 同时在测试配置中添加了测试用的OpenAI API密钥环境变量
1901 lines
67 KiB
Python
1901 lines
67 KiB
Python
from __future__ import annotations
|
|
|
|
from datetime import datetime
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
from yuxi.channels.adapters.qqbot.adapter import QQBotAdapter
|
|
from yuxi.channels.adapters.qqbot.format import (
|
|
build_ark_payload,
|
|
build_embed_payload,
|
|
build_image_payload,
|
|
build_markdown_payload,
|
|
build_text_payload,
|
|
format_outbound,
|
|
)
|
|
from yuxi.channels.adapters.qqbot.security import (
|
|
check_dm_policy,
|
|
check_group_policy,
|
|
check_mention_required,
|
|
)
|
|
from yuxi.channels.adapters.qqbot.session import (
|
|
resolve_agent_route,
|
|
resolve_chat_type,
|
|
resolve_thread_key,
|
|
)
|
|
from yuxi.channels.adapters.qqbot.streaming import send_blocks_stream
|
|
from yuxi.channels.adapters.qqbot.token import QQBotTokenManager
|
|
from yuxi.channels.models import (
|
|
Attachment,
|
|
ChannelIdentity,
|
|
ChannelMessage,
|
|
ChannelResponse,
|
|
ChannelType,
|
|
ChatType,
|
|
EventType,
|
|
MessageType,
|
|
MentionsInfo,
|
|
)
|
|
|
|
|
|
def _make_identity(
|
|
channel_chat_id: str = "group_abc123",
|
|
channel_user_id: str = "user_001",
|
|
channel_message_id: str = "msg_001",
|
|
) -> ChannelIdentity:
|
|
return ChannelIdentity(
|
|
channel_id="qqbot",
|
|
channel_type=ChannelType.QQ_BOT,
|
|
channel_user_id=channel_user_id,
|
|
channel_chat_id=channel_chat_id,
|
|
channel_message_id=channel_message_id,
|
|
)
|
|
|
|
|
|
def _make_response(
|
|
content: str = "Hello from QQBot",
|
|
channel_chat_id: str = "group_abc123",
|
|
**kwargs,
|
|
) -> ChannelResponse:
|
|
identity = _make_identity(channel_chat_id=channel_chat_id)
|
|
return ChannelResponse(identity=identity, content=content, **kwargs)
|
|
|
|
|
|
# ==================== Token Manager Tests ====================
|
|
|
|
|
|
class TestQQBotTokenManager:
|
|
def test_token_initial_state(self):
|
|
tm = QQBotTokenManager(app_id="test_app", app_secret="test_secret")
|
|
assert tm._access_token is None
|
|
assert tm._is_expired() is True
|
|
|
|
def test_api_base_production(self):
|
|
tm = QQBotTokenManager(app_id="test_app", app_secret="test_secret", sandbox=False)
|
|
assert tm.api_base == "https://api.sgroup.qq.com"
|
|
|
|
def test_api_base_sandbox(self):
|
|
tm = QQBotTokenManager(app_id="test_app", app_secret="test_secret", sandbox=True)
|
|
assert tm.api_base == "https://sandbox.api.sgroup.qq.com"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_token_refreshes(self):
|
|
tm = QQBotTokenManager(app_id="test_app", app_secret="test_secret")
|
|
|
|
with patch.object(tm, "_refresh", new_callable=AsyncMock) as mock_refresh:
|
|
mock_refresh.return_value = None
|
|
tm._access_token = None
|
|
await tm.get_token()
|
|
mock_refresh.assert_called_once()
|
|
|
|
def test_is_expired_with_no_token(self):
|
|
tm = QQBotTokenManager(app_id="test_app", app_secret="test_secret")
|
|
assert tm._is_expired() is True
|
|
|
|
|
|
# ==================== Format Tests ====================
|
|
|
|
|
|
class TestFormatOutbound:
|
|
def test_text_payload_group(self):
|
|
resp = _make_response(content="Hello World", channel_chat_id="group_abc123")
|
|
payload = build_text_payload(resp)
|
|
assert payload["content"] == "Hello World"
|
|
assert payload["group_openid"] == "abc123"
|
|
assert "channel_id" not in payload
|
|
|
|
def test_text_payload_channel(self):
|
|
resp = _make_response(content="Hello", channel_chat_id="channel_xyz")
|
|
payload = build_text_payload(resp)
|
|
assert payload["channel_id"] == "channel_xyz"
|
|
|
|
def test_text_payload_truncation(self):
|
|
long_content = "A" * 3000
|
|
resp = _make_response(content=long_content)
|
|
payload = build_text_payload(resp)
|
|
assert len(payload["content"]) == 2000
|
|
|
|
def test_text_payload_with_reply(self):
|
|
resp = _make_response(content="Reply", reply_to_message_id="msg_orig")
|
|
payload = build_text_payload(resp)
|
|
assert payload["msg_id"] == "msg_orig"
|
|
|
|
def test_markdown_payload(self):
|
|
resp = _make_response(content="**Bold** text")
|
|
resp.metadata["markdown_template_id"] = "tmpl_001"
|
|
payload = build_markdown_payload(resp)
|
|
assert payload["msg_type"] == 2
|
|
assert payload["markdown"]["template_id"] == "tmpl_001"
|
|
|
|
def test_ark_payload(self):
|
|
resp = _make_response(content="Ark content")
|
|
resp.metadata["ark_template_id"] = "ark_23"
|
|
resp.metadata["ark_data"] = {"title": "Hello", "desc": "World"}
|
|
payload = build_ark_payload(resp)
|
|
assert payload["msg_type"] == 3
|
|
assert payload["ark"]["template_id"] == "ark_23"
|
|
|
|
def test_embed_payload(self):
|
|
resp = _make_response(content="Embed desc")
|
|
resp.metadata["embed"] = {"title": "Title", "fields": []}
|
|
payload = build_embed_payload(resp)
|
|
assert payload["msg_type"] == 4
|
|
assert payload["embed"]["title"] == "Title"
|
|
|
|
def test_image_payload(self):
|
|
resp = _make_response(content="Image caption")
|
|
resp.attachments = [Attachment(type="image", file_id="file_123")]
|
|
payload = build_image_payload(resp, "file_123")
|
|
assert payload["msg_type"] == 1
|
|
assert payload["image"] == "file_123"
|
|
|
|
def test_format_outbound_defaults_to_text(self):
|
|
resp = _make_response(content="Plain text")
|
|
payload = format_outbound(resp)
|
|
assert "content" in payload
|
|
|
|
def test_format_outbound_with_markdown_flag(self):
|
|
resp = _make_response(content="Markdown content")
|
|
payload = format_outbound(resp, use_markdown=True)
|
|
assert "markdown" in payload
|
|
|
|
|
|
# ==================== Normalize Inbound Tests ====================
|
|
|
|
|
|
class TestNormalizeInbound:
|
|
@pytest.fixture
|
|
def adapter(self):
|
|
config = {"app_id": "test_app", "app_secret": "test_secret", "dm_policy": "open"}
|
|
return QQBotAdapter(config=config)
|
|
|
|
def test_at_message_create(self, adapter):
|
|
raw = {
|
|
"event_type": "at_message_create",
|
|
"event": {
|
|
"id": "msg_001",
|
|
"content": "你好 @bot",
|
|
"author": {"id": "user_123"},
|
|
"group_openid": "group_abc",
|
|
"timestamp": "2026-05-08T12:00:00",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.identity.channel_chat_id == "group_group_abc"
|
|
assert msg.chat_type == ChatType.GROUP
|
|
assert msg.content == "你好 @bot"
|
|
assert msg.mentions is not None
|
|
assert msg.mentions.is_bot_mentioned is True
|
|
|
|
def test_direct_message_create(self, adapter):
|
|
raw = {
|
|
"event_type": "direct_message_create",
|
|
"event": {
|
|
"id": "msg_002",
|
|
"content": "Private message",
|
|
"author": {"id": "user_456"},
|
|
"timestamp": "2026-05-08T12:00:00",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.identity.channel_chat_id == "dm_user_456"
|
|
assert msg.chat_type == ChatType.DIRECT
|
|
assert msg.content == "Private message"
|
|
|
|
def test_message_create_guild(self, adapter):
|
|
raw = {
|
|
"event_type": "message_create",
|
|
"event": {
|
|
"id": "msg_003",
|
|
"content": "Channel message",
|
|
"author": {"id": "user_789"},
|
|
"channel_id": "ch_guild_001",
|
|
"guild_id": "guild_001",
|
|
"timestamp": "2026-05-08T12:00:00",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.identity.channel_chat_id == "ch_guild_001"
|
|
assert msg.chat_type == ChatType.GUILD_CHANNEL
|
|
assert msg.metadata.get("guild_id") == "guild_001"
|
|
|
|
def test_command_message(self, adapter):
|
|
raw = {
|
|
"event_type": "at_message_create",
|
|
"event": {
|
|
"id": "msg_004",
|
|
"content": "/reset",
|
|
"author": {"id": "user_123"},
|
|
"group_openid": "group_abc",
|
|
"timestamp": "2026-05-08T12:00:00",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.message_type == MessageType.COMMAND
|
|
|
|
def test_with_attachments(self, adapter):
|
|
raw = {
|
|
"event_type": "direct_message_create",
|
|
"event": {
|
|
"id": "msg_005",
|
|
"content": "Image",
|
|
"author": {"id": "user_456"},
|
|
"timestamp": "2026-05-08T12:00:00",
|
|
"attachments": [
|
|
{
|
|
"content_type": "image/png",
|
|
"url": "https://example.com/img.png",
|
|
"filename": "img.png",
|
|
},
|
|
{
|
|
"content_type": "application/pdf",
|
|
"url": "https://example.com/doc.pdf",
|
|
"filename": "doc.pdf",
|
|
"size": 1024,
|
|
},
|
|
],
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert len(msg.attachments) == 2
|
|
assert msg.attachments[0].type == "image"
|
|
assert msg.attachments[1].type == "file"
|
|
|
|
def test_unknown_event_type_raises(self, adapter):
|
|
from yuxi.channels.exceptions import MessageFormatError
|
|
|
|
raw = {"event_type": "unknown_event", "event": {}}
|
|
with pytest.raises(MessageFormatError):
|
|
adapter.normalize_inbound(raw)
|
|
|
|
|
|
# ==================== Security Tests ====================
|
|
|
|
|
|
class TestSecurityPolicy:
|
|
@pytest.mark.asyncio
|
|
async def test_dm_policy_open(self):
|
|
config = {"dm_policy": "open"}
|
|
result = await check_dm_policy("user_001", config)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dm_policy_disabled(self):
|
|
config = {"dm_policy": "disabled"}
|
|
result = await check_dm_policy("user_001", config)
|
|
assert result is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dm_policy_allowlist_match(self):
|
|
config = {"dm_policy": "allowlist", "allow_from": ["qq:user_001", "qq:user_002"]}
|
|
result = await check_dm_policy("user_001", config)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dm_policy_allowlist_no_match(self):
|
|
config = {"dm_policy": "allowlist", "allow_from": ["qq:user_002"]}
|
|
result = await check_dm_policy("user_001", config)
|
|
assert result is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_group_policy_open(self):
|
|
config = {"group_policy": "open"}
|
|
result = await check_group_policy("group_abc", "user_001", config)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_group_policy_disabled(self):
|
|
config = {"group_policy": "disabled"}
|
|
result = await check_group_policy("group_abc", "user_001", config)
|
|
assert result is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_group_policy_allowlist_global(self):
|
|
config = {"group_policy": "allowlist", "group_allow_from": ["qq:user_001"]}
|
|
result = await check_group_policy("group_abc", "user_001", config)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_group_policy_allowlist_per_group(self):
|
|
config = {
|
|
"group_policy": "allowlist",
|
|
"groups": {"group_abc": {"allow_from": ["qq:user_001"]}},
|
|
}
|
|
result = await check_group_policy("group_abc", "user_001", config)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mention_required_disabled(self):
|
|
config = {"group_require_mention": False}
|
|
msg = ChannelMessage(
|
|
identity=_make_identity(),
|
|
content="Hello",
|
|
chat_type=ChatType.GROUP,
|
|
)
|
|
result = await check_mention_required("group_abc", msg, config)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mention_required_with_bot_mentioned(self):
|
|
config = {"group_require_mention": True}
|
|
msg = ChannelMessage(
|
|
identity=_make_identity(),
|
|
content="Hello",
|
|
chat_type=ChatType.GROUP,
|
|
mentions=MentionsInfo(is_bot_mentioned=True),
|
|
)
|
|
result = await check_mention_required("group_abc", msg, config)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mention_required_not_mentioned(self):
|
|
config = {"group_require_mention": True}
|
|
msg = ChannelMessage(
|
|
identity=_make_identity(),
|
|
content="Hello",
|
|
chat_type=ChatType.GROUP,
|
|
)
|
|
result = await check_mention_required("group_abc", msg, config)
|
|
assert result is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mention_required_by_name(self):
|
|
config = {"group_require_mention": True}
|
|
msg = ChannelMessage(
|
|
identity=_make_identity(),
|
|
content="@mybot help",
|
|
chat_type=ChatType.GROUP,
|
|
)
|
|
result = await check_mention_required("group_abc", msg, config, bot_names=["mybot"])
|
|
assert result is True
|
|
|
|
|
|
# ==================== Session Tests ====================
|
|
|
|
|
|
class TestSessionRouting:
|
|
def test_resolve_thread_key_group(self):
|
|
identity = _make_identity(channel_chat_id="group_abc123")
|
|
thread_key = resolve_thread_key(identity)
|
|
assert thread_key == "qqbot:group:group_abc123"
|
|
|
|
def test_resolve_thread_key_direct(self):
|
|
identity = _make_identity(channel_chat_id="dm_user_openid_123")
|
|
thread_key = resolve_thread_key(identity)
|
|
assert thread_key == "qqbot:direct:dm_user_openid_123"
|
|
|
|
def test_resolve_thread_key_guild(self):
|
|
identity = _make_identity(channel_chat_id="ch_guild_001")
|
|
thread_key = resolve_thread_key(identity)
|
|
assert thread_key == "qqbot:guild:ch_guild_001"
|
|
|
|
def test_resolve_chat_type_group(self):
|
|
identity = _make_identity(channel_chat_id="group_abc")
|
|
assert resolve_chat_type(identity) == "group"
|
|
|
|
def test_resolve_chat_type_direct(self):
|
|
identity = _make_identity(channel_chat_id="user_123")
|
|
assert resolve_chat_type(identity) == "guild_channel"
|
|
|
|
def test_resolve_agent_route_group(self):
|
|
identity = _make_identity(channel_chat_id="group_abc123")
|
|
route = resolve_agent_route(identity, default_agent_id="my_bot")
|
|
assert route == "agent:my_bot:qqbot:group:group_abc123"
|
|
|
|
def test_resolve_agent_route_group_custom_agent(self):
|
|
identity = _make_identity(channel_chat_id="group_abc123")
|
|
groups_config = {"group_abc123": {"agent_id": "custom_agent"}}
|
|
route = resolve_agent_route(identity, default_agent_id="default", groups_config=groups_config)
|
|
assert route == "agent:custom_agent:qqbot:group:group_abc123"
|
|
|
|
|
|
# ==================== Adapter Tests ====================
|
|
|
|
|
|
class TestQQBotAdapter:
|
|
@pytest.fixture
|
|
def adapter(self):
|
|
config = {"app_id": "test_app", "app_secret": "test_secret", "dm_policy": "open"}
|
|
return QQBotAdapter(config=config)
|
|
|
|
def test_channel_id(self, adapter):
|
|
assert adapter.channel_id == "qqbot"
|
|
|
|
def test_channel_type(self, adapter):
|
|
assert adapter.channel_type == ChannelType.QQ_BOT
|
|
|
|
def test_text_chunk_limit(self, adapter):
|
|
assert adapter.text_chunk_limit == 2000
|
|
|
|
def test_supports_markdown(self, adapter):
|
|
assert adapter.supports_markdown is True
|
|
|
|
def test_supports_streaming(self, adapter):
|
|
assert adapter.supports_streaming is True
|
|
|
|
def test_streaming_modes(self, adapter):
|
|
assert "off" in adapter.streaming_modes
|
|
assert "block" in adapter.streaming_modes
|
|
|
|
def test_initial_status(self, adapter):
|
|
assert adapter.status == "disconnected"
|
|
|
|
def test_format_outbound_text(self, adapter):
|
|
resp = _make_response(content="Hello")
|
|
payload = adapter.format_outbound(resp)
|
|
assert payload["content"] == "Hello"
|
|
|
|
def test_format_outbound_markdown(self, adapter):
|
|
adapter.config["use_markdown"] = True
|
|
resp = _make_response(content="**Bold**")
|
|
payload = adapter.format_outbound(resp)
|
|
assert "markdown" in payload
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_health_check_without_token(self, adapter):
|
|
status = await adapter.health_check()
|
|
assert status.status == "unhealthy"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_refresh_token_no_manager(self, adapter):
|
|
result = await adapter._refresh_token_if_needed()
|
|
assert result is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_send_without_client(self, adapter):
|
|
resp = _make_response(content="Test")
|
|
result = await adapter.send(resp)
|
|
assert result.success is False
|
|
|
|
def test_receive_yields_nothing(self, adapter):
|
|
async def _collect():
|
|
collected = []
|
|
async for msg in adapter.receive():
|
|
collected.append(msg)
|
|
return collected
|
|
|
|
import asyncio
|
|
result = asyncio.run(_collect())
|
|
assert result == []
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_security_dm_open(self, adapter):
|
|
msg = ChannelMessage(
|
|
identity=_make_identity(channel_chat_id="user_123"),
|
|
content="Hello",
|
|
chat_type=ChatType.DIRECT,
|
|
)
|
|
result = await adapter._check_security(msg)
|
|
assert result is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_security_group_with_mention(self, adapter):
|
|
adapter.config["group_policy"] = "open"
|
|
adapter.config["group_require_mention"] = False
|
|
adapter._bot_info = {"username": "mybot"}
|
|
msg = ChannelMessage(
|
|
identity=_make_identity(channel_chat_id="group_abc"),
|
|
content="Hello",
|
|
chat_type=ChatType.GROUP,
|
|
)
|
|
result = await adapter._check_security(msg)
|
|
assert result is True
|
|
|
|
|
|
# ==================== Streaming Tests ====================
|
|
|
|
|
|
class TestStreaming:
|
|
@pytest.mark.asyncio
|
|
async def test_stream_blocks(self):
|
|
text = "Para 1\n\nPara 2\n\nPara 3"
|
|
sent_messages = []
|
|
|
|
async def mock_send(response):
|
|
sent_messages.append(response.content)
|
|
from yuxi.channels.models import DeliveryResult
|
|
return DeliveryResult(success=True, message_id=f"msg_{len(sent_messages)}")
|
|
|
|
await send_blocks_stream(
|
|
"chat_001", text, mock_send,
|
|
channel_id="qqbot", channel_type=ChannelType.QQ_BOT,
|
|
chunk_size=1,
|
|
)
|
|
assert len(sent_messages) > 0
|
|
full_text = " ".join(sent_messages)
|
|
assert "Para 1" in full_text
|
|
assert "Para 2" in full_text
|
|
assert "Para 3" in full_text
|
|
|
|
|
|
# ==================== Send URL Resolution Tests ====================
|
|
|
|
|
|
class TestSendUrlResolution:
|
|
def test_group_chat_prefix(self):
|
|
from yuxi.channels.adapters.qqbot.send import _resolve_send_url
|
|
|
|
url = _resolve_send_url("https://api.sgroup.qq.com", "group_openid123")
|
|
assert "/v2/groups/openid123/messages" in url
|
|
|
|
def test_channel_id(self):
|
|
from yuxi.channels.adapters.qqbot.send import _resolve_send_url
|
|
|
|
url = _resolve_send_url("https://api.sgroup.qq.com", "channel_xyz")
|
|
assert "/v2/channels/channel_xyz/messages" in url
|
|
|
|
def test_direct_c2c(self):
|
|
from yuxi.channels.adapters.qqbot.send import _resolve_send_url
|
|
|
|
url = _resolve_send_url("https://api.sgroup.qq.com", "")
|
|
assert "/v2/users/@me/messages" in url
|
|
|
|
|
|
# ==================== Token Manager with Shared Session ====================
|
|
|
|
|
|
class TestTokenManagerSharedSession:
|
|
@pytest.mark.asyncio
|
|
async def test_with_external_session(self):
|
|
mock_session = MagicMock()
|
|
mock_resp = MagicMock()
|
|
mock_resp.status = 200
|
|
mock_resp.json = AsyncMock(return_value={
|
|
"access_token": "token123",
|
|
"expires_in": 7200,
|
|
})
|
|
mock_session.post.return_value.__aenter__ = AsyncMock(return_value=mock_resp)
|
|
mock_session.post.return_value.__aexit__ = AsyncMock()
|
|
|
|
tm = QQBotTokenManager(
|
|
app_id="test_app",
|
|
app_secret="test_secret",
|
|
http_client=mock_session,
|
|
)
|
|
token = await tm.get_token()
|
|
assert token == "token123"
|
|
mock_session.post.assert_called_once()
|
|
|
|
|
|
# ==================== Webhook Ed25519 Tests ====================
|
|
|
|
|
|
class TestWebhookSignature:
|
|
def test_missing_headers_returns_false(self):
|
|
from yuxi.channels.adapters.qqbot.security import verify_webhook_ed25519
|
|
|
|
result = verify_webhook_ed25519({}, b"body", "secret")
|
|
assert result is False
|
|
|
|
def test_valid_signature(self):
|
|
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
|
|
|
|
from yuxi.channels.adapters.qqbot.security import verify_webhook_ed25519
|
|
|
|
private_key = Ed25519PrivateKey.generate()
|
|
seed = private_key.private_bytes_raw().hex()
|
|
|
|
timestamp = "1750407202"
|
|
body = b'{"d":{"plain_token":"test"}}'
|
|
message = timestamp.encode() + body
|
|
signature = private_key.sign(message)
|
|
|
|
headers = {
|
|
"x-signature-ed25519": signature.hex(),
|
|
"x-signature-timestamp": timestamp,
|
|
}
|
|
result = verify_webhook_ed25519(headers, body, seed)
|
|
assert result is True
|
|
|
|
def test_invalid_signature(self):
|
|
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
|
|
|
|
from yuxi.channels.adapters.qqbot.security import verify_webhook_ed25519
|
|
|
|
key1 = Ed25519PrivateKey.generate()
|
|
key2 = Ed25519PrivateKey.generate()
|
|
|
|
seed = key1.private_bytes_raw().hex()
|
|
timestamp = "1750407202"
|
|
body = b'{"d":{"plain_token":"test"}}'
|
|
message = timestamp.encode() + body
|
|
signature = key2.sign(message)
|
|
|
|
headers = {
|
|
"x-signature-ed25519": signature.hex(),
|
|
"x-signature-timestamp": timestamp,
|
|
}
|
|
result = verify_webhook_ed25519(headers, body, seed)
|
|
assert result is False
|
|
|
|
|
|
# ==================== Health Check Tests ====================
|
|
|
|
|
|
class TestHealthCheck:
|
|
@pytest.mark.asyncio
|
|
async def test_timeout_protection(self):
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
from yuxi.channels.adapters.qqbot.probe import health_check_dsm
|
|
|
|
with patch("aiohttp.ClientSession") as mock_session_cls:
|
|
mock_session = AsyncMock()
|
|
mock_session_cls.return_value.__aenter__.return_value = mock_session
|
|
mock_get = AsyncMock()
|
|
mock_get.__aenter__.side_effect = Exception("Connection timeout")
|
|
mock_session.get.return_value = mock_get
|
|
|
|
status = await health_check_dsm(
|
|
"https://api.sgroup.qq.com",
|
|
"token",
|
|
sandbox=False,
|
|
)
|
|
assert status.status == "unhealthy"
|
|
|
|
|
|
# ==================== Image Upload with Shared Session ====================
|
|
|
|
|
|
class TestImageUpload:
|
|
@pytest.mark.asyncio
|
|
async def test_upload_with_shared_session(self):
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
from yuxi.channels.adapters.qqbot.media_upload import upload_image
|
|
|
|
mock_session = MagicMock()
|
|
mock_resp = MagicMock()
|
|
mock_resp.status = 200
|
|
mock_resp.json = AsyncMock(return_value={"file_uuid": "file_001"})
|
|
mock_session.post.return_value.__aenter__ = AsyncMock(return_value=mock_resp)
|
|
mock_session.post.return_value.__aexit__ = AsyncMock()
|
|
|
|
file_id = await upload_image(
|
|
b"fake_image_data",
|
|
"token",
|
|
http_client=mock_session,
|
|
)
|
|
assert file_id == "file_001"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_download_with_shared_session(self):
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
from yuxi.channels.adapters.qqbot.media_upload import download_image
|
|
|
|
mock_session = MagicMock()
|
|
mock_resp = MagicMock()
|
|
mock_resp.status = 200
|
|
mock_resp.read = AsyncMock(return_value=b"image_bytes")
|
|
mock_session.get.return_value.__aenter__ = AsyncMock(return_value=mock_resp)
|
|
mock_session.get.return_value.__aexit__ = AsyncMock()
|
|
|
|
data = await download_image(
|
|
"file_001",
|
|
"token",
|
|
http_client=mock_session,
|
|
)
|
|
assert data == b"image_bytes"
|
|
|
|
|
|
# ==================== Message Dedup Tests ====================
|
|
|
|
|
|
class TestMessageDedup:
|
|
@pytest.fixture
|
|
def adapter(self):
|
|
return QQBotAdapter(config={"app_id": "test_app", "app_secret": "test_secret", "dm_policy": "open"})
|
|
|
|
def test_dedup_stores_msg_id(self, adapter):
|
|
import asyncio
|
|
|
|
async def run():
|
|
await adapter._dispatch_event("at_message_create", {
|
|
"id": "msg_001",
|
|
"content": "Hello",
|
|
"author": {"id": "user_001"},
|
|
"group_openid": "group_001",
|
|
"timestamp": "2026-05-08T12:00:00",
|
|
})
|
|
|
|
asyncio.run(run())
|
|
assert "msg_001" in adapter._recent_msg_ids
|
|
|
|
def test_dedup_prunes_expired_ids(self, adapter):
|
|
import time
|
|
adapter._recent_msg_ids = {
|
|
"old_msg": time.monotonic() - 3600,
|
|
"recent_msg": time.monotonic() - 10,
|
|
}
|
|
adapter._prune_old_msg_ids(time.monotonic())
|
|
assert "old_msg" not in adapter._recent_msg_ids
|
|
assert "recent_msg" in adapter._recent_msg_ids
|
|
|
|
def test_dedup_config_default_window(self):
|
|
adapter = QQBotAdapter(config={"app_id": "123", "app_secret": "abc"})
|
|
assert adapter._dedup_window_s == 60
|
|
|
|
def test_dedup_config_custom_window(self):
|
|
adapter = QQBotAdapter(config={"app_id": "123", "app_secret": "abc", "dedup_window_s": 120})
|
|
assert adapter._dedup_window_s == 120
|
|
|
|
|
|
# ==================== Edit/Delete Message Tests ====================
|
|
|
|
|
|
class TestEditDeleteMessage:
|
|
@pytest.fixture
|
|
def adapter(self):
|
|
return QQBotAdapter(config={"app_id": "test_app", "app_secret": "test_secret", "dm_policy": "open"})
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_edit_group_message_returns_error(self, adapter):
|
|
adapter._http_client = AsyncMock()
|
|
adapter._token_manager = AsyncMock()
|
|
adapter._token_manager.get_token.return_value = "mock_token"
|
|
adapter._token_manager.api_base = "https://api.test.com"
|
|
|
|
result = await adapter.edit_message("group_openid123", "msg_001", "edited content")
|
|
assert result.success is False
|
|
assert "group" in result.error.lower()
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_delete_group_message_returns_error(self, adapter):
|
|
adapter._http_client = AsyncMock()
|
|
adapter._token_manager = AsyncMock()
|
|
adapter._token_manager.get_token.return_value = "mock_token"
|
|
adapter._token_manager.api_base = "https://api.test.com"
|
|
|
|
result = await adapter.delete_message("group_openid123", "msg_001")
|
|
assert result.success is False
|
|
assert "group" in result.error.lower()
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_edit_without_client_returns_error(self, adapter):
|
|
result = await adapter.edit_message("channel_001", "msg_001", "content")
|
|
assert result.success is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_delete_without_client_returns_error(self, adapter):
|
|
result = await adapter.delete_message("channel_001", "msg_001")
|
|
assert result.success is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_edit_channel_message_success(self, adapter):
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
adapter._http_client = MagicMock()
|
|
adapter._token_manager = MagicMock()
|
|
adapter._token_manager.get_token = AsyncMock(return_value="mock_token")
|
|
adapter._token_manager.api_base = "https://api.test.com"
|
|
|
|
mock_resp = AsyncMock()
|
|
mock_resp.status = 200
|
|
adapter._http_client.patch.return_value.__aenter__.return_value = mock_resp
|
|
|
|
result = await adapter.edit_message("channel_001", "msg_001", "new content")
|
|
assert result.success is True
|
|
assert result.message_id == "msg_001"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_delete_channel_message_success(self, adapter):
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
adapter._http_client = MagicMock()
|
|
adapter._token_manager = MagicMock()
|
|
adapter._token_manager.get_token = AsyncMock(return_value="mock_token")
|
|
adapter._token_manager.api_base = "https://api.test.com"
|
|
|
|
mock_resp = AsyncMock()
|
|
mock_resp.status = 200
|
|
adapter._http_client.delete.return_value.__aenter__.return_value = mock_resp
|
|
|
|
result = await adapter.delete_message("channel_001", "msg_001")
|
|
assert result.success is True
|
|
assert result.message_id == "msg_001"
|
|
|
|
|
|
# ==================== Member Event Tests ====================
|
|
|
|
|
|
class TestMemberEvents:
|
|
@pytest.fixture
|
|
def adapter(self):
|
|
return QQBotAdapter(config={"app_id": "test_app", "app_secret": "test_secret", "dm_policy": "open"})
|
|
|
|
def test_guild_member_add_returns_joined_event(self, adapter):
|
|
raw = {
|
|
"event_type": "guild_member_add",
|
|
"event": {
|
|
"id": "evt_001",
|
|
"user": {"id": "user_new"},
|
|
"guild_id": "guild_001",
|
|
"channel_id": "ch_001",
|
|
"guild": {"name": "测试频道"},
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.event_type.value == "member.joined"
|
|
assert msg.metadata["guild_id"] == "guild_001"
|
|
assert msg.metadata["guild_name"] == "测试频道"
|
|
|
|
def test_guild_member_remove_returns_left_event(self, adapter):
|
|
raw = {
|
|
"event_type": "guild_member_remove",
|
|
"event": {
|
|
"id": "evt_002",
|
|
"user": {"id": "user_left"},
|
|
"guild_id": "guild_001",
|
|
"channel_id": "ch_001",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.event_type.value == "member.left"
|
|
|
|
def test_group_add_robot_returns_joined_event(self, adapter):
|
|
raw = {
|
|
"event_type": "group_add_robot",
|
|
"event": {
|
|
"id": "evt_003",
|
|
"op_user": {"id": "admin_001"},
|
|
"group_openid": "group_001",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.event_type.value == "member.joined"
|
|
assert msg.identity.channel_chat_id == "group_group_001"
|
|
|
|
def test_group_del_robot_returns_left_event(self, adapter):
|
|
raw = {
|
|
"event_type": "group_del_robot",
|
|
"event": {
|
|
"id": "evt_004",
|
|
"op_user": {"id": "admin_001"},
|
|
"group_openid": "group_001",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.event_type.value == "member.left"
|
|
|
|
def test_message_delete_returns_deleted_event(self, adapter):
|
|
raw = {
|
|
"event_type": "message_delete",
|
|
"event": {
|
|
"id": "evt_005",
|
|
"op_user": {"id": "op_001"},
|
|
"channel_id": "ch_001",
|
|
"guild_id": "guild_001",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.event_type.value == "message.deleted"
|
|
|
|
def test_member_event_has_correct_channel_user_id(self, adapter):
|
|
raw = {
|
|
"event_type": "guild_member_add",
|
|
"event": {
|
|
"id": "evt_006",
|
|
"user": {"id": "user_new", "username": "新成员"},
|
|
"guild_id": "guild_001",
|
|
},
|
|
}
|
|
msg = adapter.normalize_inbound(raw)
|
|
assert msg.identity.channel_user_id == "user_new"
|
|
|
|
|
|
# ==================== Reconnect Close Code Tests ====================
|
|
|
|
|
|
class TestCloseCodeClassification:
|
|
def test_classify_fatal_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import classify_close_code, CloseCodeCategory
|
|
|
|
assert classify_close_code(4013) == CloseCodeCategory.FATAL
|
|
assert classify_close_code(4014) == CloseCodeCategory.FATAL
|
|
assert classify_close_code(4100) == CloseCodeCategory.FATAL
|
|
assert classify_close_code(4101) == CloseCodeCategory.FATAL
|
|
|
|
def test_classify_recoverable_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import classify_close_code, CloseCodeCategory
|
|
|
|
assert classify_close_code(4000) == CloseCodeCategory.RECOVERABLE
|
|
assert classify_close_code(4001) == CloseCodeCategory.RECOVERABLE
|
|
assert classify_close_code(4002) == CloseCodeCategory.RECOVERABLE
|
|
assert classify_close_code(4012) == CloseCodeCategory.RECOVERABLE
|
|
|
|
def test_classify_rate_limited_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import classify_close_code, CloseCodeCategory
|
|
|
|
assert classify_close_code(4008) == CloseCodeCategory.RATE_LIMITED
|
|
assert classify_close_code(4009) == CloseCodeCategory.RATE_LIMITED
|
|
|
|
def test_classify_server_error_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import classify_close_code, CloseCodeCategory
|
|
|
|
assert classify_close_code(4900) == CloseCodeCategory.SERVER_ERROR
|
|
assert classify_close_code(4905) == CloseCodeCategory.SERVER_ERROR
|
|
assert classify_close_code(4913) == CloseCodeCategory.SERVER_ERROR
|
|
|
|
def test_classify_none_code_returns_abnormal(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import classify_close_code, CloseCodeCategory
|
|
|
|
assert classify_close_code(None) == CloseCodeCategory.ABNORMAL
|
|
|
|
def test_classify_unknown_4xxx_returns_server_side(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import classify_close_code, CloseCodeCategory
|
|
|
|
assert classify_close_code(4099) == CloseCodeCategory.SERVER_SIDE
|
|
assert classify_close_code(4200) == CloseCodeCategory.SERVER_SIDE
|
|
|
|
def test_classify_unknown_code_returns_abnormal(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import classify_close_code, CloseCodeCategory
|
|
|
|
assert classify_close_code(1000) == CloseCodeCategory.ABNORMAL
|
|
assert classify_close_code(3000) == CloseCodeCategory.ABNORMAL
|
|
assert classify_close_code(5000) == CloseCodeCategory.ABNORMAL
|
|
|
|
|
|
class TestReconnectManagerServerError:
|
|
@pytest.mark.asyncio
|
|
async def test_server_error_transitions_to_backoff(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager, ReconnectState
|
|
|
|
mgr = QQBotReconnectManager(max_retries=5, base_delay=0.01, max_delay=0.1)
|
|
await mgr.transition(ReconnectState.CONNECTED)
|
|
mgr._session_id = "session_001"
|
|
mgr._last_seq = 10
|
|
|
|
await mgr.on_disconnect(4900)
|
|
|
|
assert mgr.state == ReconnectState.IDENTIFYING
|
|
assert mgr._retry_count == 1
|
|
assert mgr._session_id is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_server_error_exhausts_retries(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager, ReconnectState
|
|
|
|
mgr = QQBotReconnectManager(max_retries=2, base_delay=0.01, max_delay=0.1)
|
|
await mgr.transition(ReconnectState.CONNECTED)
|
|
|
|
await mgr.on_disconnect(4900)
|
|
assert mgr.state == ReconnectState.IDENTIFYING
|
|
|
|
await mgr.transition(ReconnectState.CONNECTED)
|
|
await mgr.on_disconnect(4901)
|
|
assert mgr.state == ReconnectState.FROZEN
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_fatal_code_freezes_immediately(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager, ReconnectState
|
|
|
|
mgr = QQBotReconnectManager(max_retries=5)
|
|
await mgr.transition(ReconnectState.CONNECTED)
|
|
|
|
await mgr.on_disconnect(4100)
|
|
assert mgr.state == ReconnectState.FROZEN
|
|
|
|
|
|
# ==================== Credential Backup Tests ====================
|
|
|
|
|
|
class TestCredentialBackup:
|
|
@pytest.fixture
|
|
def tmp_backup_dir(self, tmp_path):
|
|
return str(tmp_path / "credential_backup")
|
|
|
|
def test_save_and_restore(self, tmp_backup_dir):
|
|
from yuxi.channels.adapters.qqbot.credential_backup import CredentialBackup, CredentialSnapshot
|
|
|
|
backup = CredentialBackup("test_app_id", tmp_backup_dir)
|
|
snapshot = CredentialSnapshot(
|
|
app_id="test_app_id",
|
|
app_secret="test_secret_123",
|
|
access_token="token_abc",
|
|
expires_at=999999999.0,
|
|
sandbox=True,
|
|
session_id="session_xyz",
|
|
)
|
|
|
|
assert backup.save(snapshot) is True
|
|
|
|
restored = backup.restore()
|
|
assert restored is not None
|
|
assert restored.app_id == "test_app_id"
|
|
assert restored.app_secret == "test_secret_123"
|
|
assert restored.access_token == "token_abc"
|
|
assert restored.sandbox is True
|
|
assert restored.session_id == "session_xyz"
|
|
|
|
def test_restore_nonexistent(self, tmp_backup_dir):
|
|
from yuxi.channels.adapters.qqbot.credential_backup import CredentialBackup
|
|
|
|
backup = CredentialBackup("nonexistent_app", tmp_backup_dir)
|
|
assert backup.restore() is None
|
|
|
|
def test_clear_removes_backup(self, tmp_backup_dir):
|
|
from yuxi.channels.adapters.qqbot.credential_backup import CredentialBackup, CredentialSnapshot
|
|
|
|
backup = CredentialBackup("test_app_id", tmp_backup_dir)
|
|
backup.save(CredentialSnapshot(app_id="test_app_id", app_secret="secret"))
|
|
|
|
assert backup.restore() is not None
|
|
assert backup.clear() is True
|
|
assert backup.restore() is None
|
|
|
|
def test_token_expired(self):
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.credential_backup import CredentialSnapshot
|
|
|
|
snapshot = CredentialSnapshot(
|
|
app_id="test", app_secret="secret",
|
|
access_token="tok", expires_at=time.monotonic() - 100,
|
|
)
|
|
assert snapshot.token_expired() is True
|
|
|
|
snapshot_valid = CredentialSnapshot(
|
|
app_id="test", app_secret="secret",
|
|
access_token="tok", expires_at=time.monotonic() + 7200,
|
|
)
|
|
assert snapshot_valid.token_expired() is False
|
|
|
|
def test_snapshot_invalid_without_credentials(self):
|
|
from yuxi.channels.adapters.qqbot.credential_backup import CredentialSnapshot
|
|
|
|
assert CredentialSnapshot().is_valid() is False
|
|
assert CredentialSnapshot(app_id="test").is_valid() is False
|
|
assert CredentialSnapshot(app_id="test", app_secret="secret").is_valid() is True
|
|
|
|
def test_cleanup_expired(self, tmp_backup_dir):
|
|
import os
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.credential_backup import CredentialBackup, CredentialSnapshot
|
|
|
|
backup = CredentialBackup("old_app", tmp_backup_dir)
|
|
backup.save(CredentialSnapshot(app_id="old_app", app_secret="old_secret"))
|
|
|
|
filepath = os.path.join(tmp_backup_dir, "old_app.json")
|
|
mtime_back = time.time() - 86400 * 30
|
|
os.utime(filepath, (mtime_back, mtime_back))
|
|
|
|
removed = CredentialBackup.cleanup_expired(tmp_backup_dir, max_age_s=86400 * 7)
|
|
assert removed == 1
|
|
assert not os.path.exists(filepath)
|
|
|
|
|
|
# ==================== Capability Declaration Tests ====================
|
|
|
|
|
|
class TestCapabilityDeclarations:
|
|
def test_default_capabilities_edit_unsend_false(self):
|
|
assert QQBotAdapter.capabilities.edit is False
|
|
assert QQBotAdapter.capabilities.unsend is False
|
|
|
|
def test_guild_channel_enables_edit_unsend(self):
|
|
adapter = QQBotAdapter(config={
|
|
"app_id": "test",
|
|
"app_secret": "test",
|
|
"chat_types": ["guild_channel"],
|
|
})
|
|
assert adapter.capabilities.edit is True
|
|
assert adapter.capabilities.unsend is True
|
|
|
|
def test_mixed_chat_types_with_guild_enables_edit_unsend(self):
|
|
adapter = QQBotAdapter(config={
|
|
"app_id": "test",
|
|
"app_secret": "test",
|
|
"chat_types": ["direct", "group", "guild_channel"],
|
|
})
|
|
assert adapter.capabilities.edit is True
|
|
assert adapter.capabilities.unsend is True
|
|
|
|
|
|
# ==================== Pipeline Context Tests ====================
|
|
|
|
|
|
class TestPipelineContext:
|
|
def test_stop_sets_reason(self):
|
|
from yuxi.channels.pipeline.context import PipelineContext
|
|
|
|
ctx = PipelineContext()
|
|
assert ctx.stopped is False
|
|
|
|
ctx.stop("test_reason")
|
|
assert ctx.stopped is True
|
|
assert ctx._stop_reason == "test_reason"
|
|
|
|
def test_debug_summary(self):
|
|
from yuxi.channels.pipeline.context import PipelineContext
|
|
|
|
ctx = PipelineContext(event_type="test", sender_id="user_001", msg_id="msg_001")
|
|
summary = ctx.debug_summary()
|
|
assert "test" in summary
|
|
assert "user_001" in summary
|
|
|
|
|
|
# ==================== Adapter Backup Integration Tests ====================
|
|
|
|
|
|
class TestAdapterCredentialBackup:
|
|
@pytest.fixture
|
|
def adapter(self):
|
|
return QQBotAdapter(config={"app_id": "test_app", "app_secret": "test_secret", "dm_policy": "open"})
|
|
|
|
def test_backup_credentials_creates_backup(self, adapter, tmp_path):
|
|
adapter.config["credential_backup_dir"] = str(tmp_path)
|
|
|
|
adapter._token_manager = MagicMock()
|
|
adapter._token_manager._access_token = "mock_token"
|
|
adapter._token_manager._expires_at = 999999999.0
|
|
adapter._session_id = "session_test"
|
|
|
|
adapter._backup_credentials()
|
|
|
|
assert adapter._credential_backup is not None
|
|
restored = adapter._credential_backup.restore()
|
|
assert restored is not None
|
|
assert restored.app_id == "test_app"
|
|
assert restored.access_token == "mock_token"
|
|
assert restored.session_id == "session_test"
|
|
|
|
def test_restore_credentials_no_backup(self, adapter, tmp_path):
|
|
adapter.config["credential_backup_dir"] = str(tmp_path)
|
|
|
|
result = adapter._restore_credentials()
|
|
assert result is False
|
|
|
|
def test_restore_credentials_with_valid_backup(self, adapter, tmp_path):
|
|
from yuxi.channels.adapters.qqbot.credential_backup import CredentialBackup, CredentialSnapshot
|
|
|
|
adapter.config["credential_backup_dir"] = str(tmp_path)
|
|
|
|
backup = CredentialBackup("test_app", str(tmp_path))
|
|
backup.save(CredentialSnapshot(
|
|
app_id="test_app",
|
|
app_secret="test_secret",
|
|
access_token="saved_token",
|
|
expires_at=999999999.0,
|
|
session_id="saved_session",
|
|
))
|
|
|
|
result = adapter._restore_credentials()
|
|
assert result is True
|
|
assert adapter._session_id == "saved_session"
|
|
|
|
|
|
# ==================== Session Store Tests ====================
|
|
|
|
|
|
class TestSessionStore:
|
|
@pytest.fixture
|
|
def tmp_store_dir(self, tmp_path):
|
|
return str(tmp_path / "session_store")
|
|
|
|
def test_save_and_load(self, tmp_store_dir):
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.session_store import SessionStore, SessionRecord
|
|
|
|
store = SessionStore("test_app_id", tmp_store_dir)
|
|
record = SessionRecord(
|
|
session_id="session_abc",
|
|
last_seq=42,
|
|
last_heartbeat=time.monotonic(),
|
|
identify_at=time.time(),
|
|
shard_id=0,
|
|
shard_count=1,
|
|
)
|
|
|
|
assert store.save(record) is True
|
|
|
|
loaded = store.load()
|
|
assert loaded is not None
|
|
assert loaded.session_id == "session_abc"
|
|
assert loaded.last_seq == 42
|
|
assert loaded.shard_id == 0
|
|
|
|
def test_load_nonexistent(self, tmp_store_dir):
|
|
from yuxi.channels.adapters.qqbot.session_store import SessionStore
|
|
|
|
store = SessionStore("nonexistent_app", tmp_store_dir)
|
|
assert store.load() is None
|
|
|
|
def test_clear_removes_session(self, tmp_store_dir):
|
|
from yuxi.channels.adapters.qqbot.session_store import SessionStore, SessionRecord
|
|
|
|
store = SessionStore("test_app", tmp_store_dir)
|
|
store.save(SessionRecord(session_id="session_xyz"))
|
|
|
|
assert store.load() is not None
|
|
assert store.clear() is True
|
|
assert store.load() is None
|
|
|
|
def test_to_dict_and_from_dict(self):
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.session_store import SessionRecord
|
|
|
|
record = SessionRecord(
|
|
session_id="sid_001",
|
|
last_seq=100,
|
|
shard_id=1,
|
|
shard_count=2,
|
|
metadata={"key": "value"},
|
|
)
|
|
data = record.to_dict()
|
|
restored = SessionRecord.from_dict(data)
|
|
|
|
assert restored.session_id == "sid_001"
|
|
assert restored.last_seq == 100
|
|
assert restored.shard_id == 1
|
|
assert restored.shard_count == 2
|
|
assert restored.metadata == {"key": "value"}
|
|
|
|
def test_cleanup_expired(self, tmp_store_dir):
|
|
import os
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.session_store import SessionStore, SessionRecord
|
|
|
|
store = SessionStore("old_app", tmp_store_dir)
|
|
store.save(SessionRecord(session_id="old_session"))
|
|
|
|
filepath = os.path.join(tmp_store_dir, "old_app_session.json")
|
|
mtime_back = time.time() - 86400 * 30
|
|
os.utime(filepath, (mtime_back, mtime_back))
|
|
|
|
removed = SessionStore.cleanup_expired(tmp_store_dir, max_age_s=86400 * 7)
|
|
assert removed == 1
|
|
assert not os.path.exists(filepath)
|
|
|
|
|
|
# ==================== Known User Tracker Tests ====================
|
|
|
|
|
|
class TestKnownUserTracker:
|
|
@pytest.fixture
|
|
def tmp_persist_dir(self, tmp_path):
|
|
return str(tmp_path / "known_users")
|
|
|
|
def test_record_new_user(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
record = tracker.record("user_001", "Alice", "group")
|
|
|
|
assert record.user_id == "user_001"
|
|
assert record.username == "Alice"
|
|
assert record.message_count == 1
|
|
assert "group" in record.chat_types
|
|
|
|
def test_record_existing_user_updates(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
tracker.record("user_001", "Alice", "group")
|
|
record = tracker.record("user_001", "Alice_v2", "dm")
|
|
|
|
assert record.message_count == 2
|
|
assert "group" in record.chat_types
|
|
assert "dm" in record.chat_types
|
|
|
|
def test_is_known(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
tracker.record("user_001", "Alice")
|
|
|
|
assert tracker.is_known("user_001") is True
|
|
assert tracker.is_known("user_999") is False
|
|
|
|
def test_get_user(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
tracker.record("user_001", "Alice")
|
|
|
|
u = tracker.get("user_001")
|
|
assert u is not None
|
|
assert u.username == "Alice"
|
|
|
|
def test_remove_user(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
tracker.record("user_001", "Alice")
|
|
|
|
assert tracker.remove("user_001") is True
|
|
assert tracker.is_known("user_001") is False
|
|
assert tracker.remove("user_001") is False
|
|
|
|
def test_max_users_eviction(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", max_users=3, persist_dir=tmp_persist_dir)
|
|
tracker.record("user_001", "A")
|
|
tracker.record("user_002", "B")
|
|
tracker.record("user_003", "C")
|
|
tracker.record("user_004", "D")
|
|
|
|
assert tracker.count == 3
|
|
assert tracker.is_known("user_001") is False
|
|
assert tracker.is_known("user_004") is True
|
|
|
|
def test_persist_and_restore(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker1 = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
tracker1.record("user_001", "Alice")
|
|
tracker1.persist()
|
|
|
|
tracker2 = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
assert tracker2.is_known("user_001") is True
|
|
u = tracker2.get("user_001")
|
|
assert u.username == "Alice"
|
|
|
|
def test_clear(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
tracker.record("user_001", "Alice")
|
|
tracker.clear()
|
|
|
|
assert tracker.count == 0
|
|
assert tracker.is_known("user_001") is False
|
|
|
|
def test_get_recent_users(self, tmp_persist_dir):
|
|
from yuxi.channels.adapters.qqbot.known_users import KnownUserTracker
|
|
|
|
tracker = KnownUserTracker("test_app", persist_dir=tmp_persist_dir)
|
|
tracker.record("user_001", "A")
|
|
tracker.record("user_002", "B")
|
|
tracker.record("user_003", "C")
|
|
|
|
recent = tracker.get_recent_users(limit=2)
|
|
assert len(recent) == 2
|
|
assert recent[0].user_id == "user_003"
|
|
assert recent[1].user_id == "user_002"
|
|
|
|
|
|
# ==================== GroupHistoryBuffer Tests ====================
|
|
|
|
|
|
class TestGroupHistoryBuffer:
|
|
def test_record_and_retrieve(self):
|
|
from yuxi.channels.adapters.qqbot.group_buffer import GroupHistoryBuffer, GroupMessage
|
|
|
|
buf = GroupHistoryBuffer(buffer_limit=10, ttl_seconds=3600)
|
|
msg = GroupMessage(
|
|
msg_id="msg_001", author_id="user_001",
|
|
author_name="Alice", content="Hello",
|
|
timestamp=1000.0, mentions_bot=True,
|
|
)
|
|
buf.record("group_abc", msg)
|
|
|
|
context = buf.recent_context("group_abc", count=5)
|
|
assert len(context) == 1
|
|
assert context[0].content == "Hello"
|
|
assert context[0].mentions_bot is True
|
|
|
|
def test_buffer_limit(self):
|
|
from yuxi.channels.adapters.qqbot.group_buffer import GroupHistoryBuffer, GroupMessage
|
|
|
|
buf = GroupHistoryBuffer(buffer_limit=3)
|
|
for i in range(5):
|
|
buf.record("group_abc", GroupMessage(
|
|
msg_id=f"msg_{i}", author_id="user_001",
|
|
author_name="A", content=f"Content {i}", timestamp=float(i),
|
|
))
|
|
|
|
context = buf.recent_context("group_abc", count=10)
|
|
assert len(context) == 3
|
|
assert context[0].msg_id == "msg_2"
|
|
assert context[-1].msg_id == "msg_4"
|
|
|
|
def test_expired_session(self):
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.group_buffer import GroupHistoryBuffer, GroupMessage
|
|
|
|
buf = GroupHistoryBuffer(ttl_seconds=0.001)
|
|
buf.record("group_abc", GroupMessage(
|
|
msg_id="msg_001", author_id="user_001",
|
|
author_name="A", content="Hi", timestamp=time.time(),
|
|
))
|
|
|
|
import asyncio
|
|
asyncio.run(asyncio.sleep(0.01))
|
|
|
|
context = buf.recent_context("group_abc")
|
|
assert len(context) == 0
|
|
|
|
def test_gc_cleans_expired(self):
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.group_buffer import GroupHistoryBuffer, GroupMessage
|
|
|
|
buf = GroupHistoryBuffer(ttl_seconds=0.001)
|
|
buf.record("group_abc", GroupMessage(
|
|
msg_id="msg_001", author_id="user_001",
|
|
author_name="A", content="Hi", timestamp=time.time(),
|
|
))
|
|
buf.record("group_xyz", GroupMessage(
|
|
msg_id="msg_002", author_id="user_002",
|
|
author_name="B", content="Hey", timestamp=time.time() - 3600,
|
|
))
|
|
|
|
import asyncio
|
|
asyncio.run(asyncio.sleep(0.01))
|
|
|
|
removed = buf.gc()
|
|
assert removed >= 1
|
|
|
|
def test_unknown_group_returns_empty(self):
|
|
from yuxi.channels.adapters.qqbot.group_buffer import GroupHistoryBuffer
|
|
|
|
buf = GroupHistoryBuffer()
|
|
assert buf.recent_context("unknown_group") == []
|
|
|
|
|
|
# ==================== Server Error Classification Tests ====================
|
|
|
|
|
|
class TestServerErrorClassification:
|
|
def test_classify_overload_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
classify_server_error_category, ServerErrorCategory,
|
|
)
|
|
|
|
assert classify_server_error_category(4901) == ServerErrorCategory.OVERLOAD
|
|
assert classify_server_error_category(4904) == ServerErrorCategory.OVERLOAD
|
|
assert classify_server_error_category(4912) == ServerErrorCategory.OVERLOAD
|
|
|
|
def test_classify_maintenance_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
classify_server_error_category, ServerErrorCategory,
|
|
)
|
|
|
|
assert classify_server_error_category(4902) == ServerErrorCategory.MAINTENANCE
|
|
assert classify_server_error_category(4905) == ServerErrorCategory.MAINTENANCE
|
|
|
|
def test_classify_network_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
classify_server_error_category, ServerErrorCategory,
|
|
)
|
|
|
|
assert classify_server_error_category(4903) == ServerErrorCategory.NETWORK
|
|
assert classify_server_error_category(4907) == ServerErrorCategory.NETWORK
|
|
|
|
def test_classify_unavailable_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
classify_server_error_category, ServerErrorCategory,
|
|
)
|
|
|
|
assert classify_server_error_category(4908) == ServerErrorCategory.UNAVAILABLE
|
|
|
|
def test_classify_timeout_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
classify_server_error_category, ServerErrorCategory,
|
|
)
|
|
|
|
assert classify_server_error_category(4913) == ServerErrorCategory.TIMEOUT
|
|
|
|
def test_classify_internal_codes(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
classify_server_error_category, ServerErrorCategory,
|
|
)
|
|
|
|
assert classify_server_error_category(4900) == ServerErrorCategory.INTERNAL
|
|
assert classify_server_error_category(4910) == ServerErrorCategory.INTERNAL
|
|
|
|
def test_classify_unknown_server_code(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
classify_server_error_category, ServerErrorCategory,
|
|
)
|
|
|
|
assert classify_server_error_category(4999) == ServerErrorCategory.UNKNOWN
|
|
|
|
def test_get_server_error_name(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import get_server_error_name
|
|
|
|
assert get_server_error_name(4900) == "server_internal_error"
|
|
assert get_server_error_name(4901) == "server_overload"
|
|
assert get_server_error_name(4913) == "server_timeout"
|
|
assert get_server_error_name(9999) == "server_error_9999"
|
|
|
|
def test_calc_delay_for_server_error(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import (
|
|
QQBotReconnectManager,
|
|
classify_server_error_category,
|
|
)
|
|
|
|
mgr = QQBotReconnectManager(base_delay=1.0, jitter=0)
|
|
|
|
cat = classify_server_error_category(4901)
|
|
delay = mgr._calc_delay_for_server_error(cat)
|
|
assert delay > mgr._calc_delay()
|
|
|
|
cat_maintenance = classify_server_error_category(4902)
|
|
delay_m = mgr._calc_delay_for_server_error(cat_maintenance)
|
|
assert delay_m > delay
|
|
|
|
|
|
# ==================== Rapid Disconnect Detection Tests ====================
|
|
|
|
|
|
class TestRapidDisconnectDetection:
|
|
@pytest.mark.asyncio
|
|
async def test_no_warning_on_normal_disconnect(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager, ReconnectState
|
|
|
|
mgr = QQBotReconnectManager(max_retries=5, base_delay=0.01)
|
|
mgr._last_connect_time = 100.0
|
|
|
|
await mgr.on_disconnect(4000)
|
|
assert mgr._rapid_disconnect_count == 0
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_detects_rapid_disconnect(self):
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager, ReconnectState
|
|
|
|
mgr = QQBotReconnectManager(max_retries=5, base_delay=0.01)
|
|
mgr._last_connect_time = time.monotonic()
|
|
|
|
await mgr.on_disconnect(4000)
|
|
assert mgr._rapid_disconnect_count == 1
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_rapid_disconnect_threshold_reset(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager
|
|
|
|
mgr = QQBotReconnectManager(max_retries=5)
|
|
mgr._last_connect_time = 100.0
|
|
|
|
mgr._check_rapid_disconnect(4000, None, 103.0)
|
|
assert mgr._rapid_disconnect_count == 1
|
|
|
|
mgr._check_rapid_disconnect(4000, None, 107.0)
|
|
assert mgr._rapid_disconnect_count == 0
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mark_connected_resets_counter(self):
|
|
import time
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager
|
|
|
|
mgr = QQBotReconnectManager()
|
|
mgr._rapid_disconnect_count = 5
|
|
mgr.mark_connected()
|
|
|
|
assert mgr._rapid_disconnect_count == 0
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_fatal_code_skips_rapid_detection(self):
|
|
from yuxi.channels.adapters.qqbot.reconnect import QQBotReconnectManager
|
|
|
|
mgr = QQBotReconnectManager()
|
|
mgr._last_connect_time = 100.0
|
|
|
|
mgr._check_rapid_disconnect(4100, None, 102.0)
|
|
assert mgr._rapid_disconnect_count == 0
|
|
|
|
|
|
# ==================== Audio Format Policy Tests ====================
|
|
|
|
|
|
class TestAudioFormatPolicy:
|
|
def test_default_policy(self):
|
|
from yuxi.channels.adapters.qqbot.audio import AudioFormatPolicy, AudioFormat
|
|
|
|
policy = AudioFormatPolicy()
|
|
assert policy.transcode_enabled is True
|
|
assert policy.upload_direct_formats == ["wav", "mp3"]
|
|
assert policy.stt_direct_formats == ["wav", "mp3", "pcm"]
|
|
assert policy.fallback_format == AudioFormat.MP3
|
|
|
|
def test_needs_transcode_direct_format(self):
|
|
from yuxi.channels.adapters.qqbot.audio import AudioFormatPolicy, AudioFormat
|
|
|
|
policy = AudioFormatPolicy()
|
|
assert policy.needs_transcode(AudioFormat.MP3) is False
|
|
assert policy.needs_transcode(AudioFormat.WAV) is False
|
|
|
|
def test_needs_transcode_indirect_format(self):
|
|
from yuxi.channels.adapters.qqbot.audio import AudioFormatPolicy, AudioFormat
|
|
|
|
policy = AudioFormatPolicy()
|
|
assert policy.needs_transcode(AudioFormat.SILK) is True
|
|
assert policy.needs_transcode(AudioFormat.AAC) is True
|
|
|
|
def test_needs_transcode_disabled(self):
|
|
from yuxi.channels.adapters.qqbot.audio import AudioFormatPolicy, AudioFormat
|
|
|
|
policy = AudioFormatPolicy(transcode_enabled=False)
|
|
assert policy.needs_transcode(AudioFormat.SILK) is False
|
|
|
|
def test_can_stt_direct(self):
|
|
from yuxi.channels.adapters.qqbot.audio import AudioFormatPolicy, AudioFormat
|
|
|
|
policy = AudioFormatPolicy()
|
|
assert policy.can_stt_direct(AudioFormat.WAV) is True
|
|
assert policy.can_stt_direct(AudioFormat.MP3) is True
|
|
assert policy.can_stt_direct(AudioFormat.PCM) is True
|
|
assert policy.can_stt_direct(AudioFormat.SILK) is False
|
|
|
|
def test_from_config(self):
|
|
from yuxi.channels.adapters.qqbot.audio import AudioFormatPolicy, AudioFormat
|
|
|
|
config = {
|
|
"audio_format_policy": {
|
|
"transcode_enabled": False,
|
|
"upload_direct_formats": ["wav"],
|
|
"stt_direct_formats": ["wav", "mp3"],
|
|
"fallback_format": "wav",
|
|
"sample_rate": 44100,
|
|
"channels": 2,
|
|
"bitrate": 64000,
|
|
}
|
|
}
|
|
policy = AudioFormatPolicy.from_config(config)
|
|
assert policy.transcode_enabled is False
|
|
assert policy.upload_direct_formats == ["wav"]
|
|
assert policy.stt_direct_formats == ["wav", "mp3"]
|
|
assert policy.fallback_format == AudioFormat.WAV
|
|
assert policy.sample_rate == 44100
|
|
assert policy.channels == 2
|
|
assert policy.bitrate == 64000
|
|
|
|
def test_from_config_empty(self):
|
|
from yuxi.channels.adapters.qqbot.audio import AudioFormatPolicy
|
|
|
|
policy = AudioFormatPolicy.from_config(None)
|
|
assert policy.transcode_enabled is True
|
|
policy2 = AudioFormatPolicy.from_config({})
|
|
assert policy2.transcode_enabled is True
|
|
|
|
|
|
# ==================== TTS Provider Tests ====================
|
|
|
|
|
|
class TestTTSProvider:
|
|
def test_default_config(self):
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider, AudioFormat
|
|
|
|
provider = TTSProvider()
|
|
assert provider._default_voice == "zh-CN-XiaoxiaoNeural"
|
|
assert provider._default_format == AudioFormat.MP3
|
|
|
|
def test_custom_config(self):
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider, AudioFormat
|
|
|
|
provider = TTSProvider(
|
|
default_voice="en-US-JennyNeural",
|
|
default_format=AudioFormat.WAV,
|
|
)
|
|
assert provider._default_voice == "en-US-JennyNeural"
|
|
assert provider._default_format == AudioFormat.WAV
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_synthesize_builtin_returns_bytes(self):
|
|
from unittest.mock import patch
|
|
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider
|
|
|
|
provider = TTSProvider()
|
|
|
|
with patch.object(provider, "_synthesize_builtin") as mock_synth:
|
|
mock_synth.return_value = b"mock_audio_bytes"
|
|
|
|
result = await provider.synthesize("你好世界")
|
|
assert result == b"mock_audio_bytes"
|
|
mock_synth.assert_called_once_with("你好世界")
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_synthesize_azure(self):
|
|
from unittest.mock import patch
|
|
|
|
import os
|
|
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider
|
|
|
|
with patch.dict(os.environ, {"QQBOT_TTS_PROVIDER": "azure"}, clear=False):
|
|
provider = TTSProvider()
|
|
|
|
with patch.object(provider, "_synthesize_azure") as mock_synth:
|
|
mock_synth.return_value = b"azure_audio_bytes"
|
|
|
|
result = await provider.synthesize("Hello")
|
|
assert result == b"azure_audio_bytes"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_synthesize_edge(self):
|
|
from unittest.mock import patch
|
|
|
|
import os
|
|
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider
|
|
|
|
with patch.dict(os.environ, {"QQBOT_TTS_PROVIDER": "edge"}, clear=False):
|
|
provider = TTSProvider()
|
|
|
|
with patch.object(provider, "_synthesize_edge") as mock_synth:
|
|
mock_synth.return_value = b"edge_audio_bytes"
|
|
|
|
result = await provider.synthesize("Hello")
|
|
assert result == b"edge_audio_bytes"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_synthesize_uses_cache(self):
|
|
from unittest.mock import patch
|
|
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider
|
|
|
|
provider = TTSProvider()
|
|
|
|
with patch.object(provider, "_synthesize_builtin") as mock_synth:
|
|
mock_synth.return_value = b"cached_bytes"
|
|
|
|
result1 = await provider.synthesize("test text")
|
|
result2 = await provider.synthesize("test text")
|
|
|
|
assert result1 == b"cached_bytes"
|
|
assert result2 == b"cached_bytes"
|
|
mock_synth.assert_called_once()
|
|
|
|
def test_clear_cache(self):
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider
|
|
|
|
provider = TTSProvider()
|
|
provider._cache["test_key"] = b"test_data"
|
|
provider.clear_cache()
|
|
assert len(provider._cache) == 0
|
|
|
|
def test_generate_silence(self):
|
|
from yuxi.channels.adapters.qqbot.audio import TTSProvider
|
|
|
|
silence = TTSProvider._generate_silence(1.0)
|
|
assert len(silence) == 32000
|
|
|
|
|
|
# ==================== STT Provider Tests ====================
|
|
|
|
|
|
class TestSTTProvider:
|
|
def test_default_config(self):
|
|
from yuxi.channels.adapters.qqbot.audio import STTProvider
|
|
|
|
provider = STTProvider()
|
|
assert provider._provider == "builtin"
|
|
|
|
def test_from_config(self):
|
|
from yuxi.channels.adapters.qqbot.audio import STTProvider
|
|
|
|
config = {
|
|
"stt": {
|
|
"provider": "azure",
|
|
"api_key": "test_key",
|
|
"region": "eastasia",
|
|
"model": "base",
|
|
}
|
|
}
|
|
provider = STTProvider.from_config(config)
|
|
assert provider._provider == "azure"
|
|
assert provider._api_key == "test_key"
|
|
assert provider._region == "eastasia"
|
|
assert provider._model == "base"
|
|
|
|
def test_from_config_empty(self):
|
|
from yuxi.channels.adapters.qqbot.audio import STTProvider
|
|
|
|
provider = STTProvider.from_config(None)
|
|
assert provider._provider == "builtin"
|
|
assert provider._api_key == ""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_transcribe_builtin_returns_string(self):
|
|
from unittest.mock import patch
|
|
|
|
from yuxi.channels.adapters.qqbot.audio import STTProvider
|
|
|
|
provider = STTProvider()
|
|
|
|
with patch.object(provider, "_transcribe_builtin") as mock_stt:
|
|
mock_stt.return_value = "你好世界"
|
|
|
|
result = await provider.transcribe(b"fake_audio")
|
|
assert result == "你好世界"
|
|
|
|
|
|
# ==================== Adapter TTS/STT Integration Tests ====================
|
|
|
|
|
|
class TestAdapterTTSIntegration:
|
|
@pytest.fixture
|
|
def adapter(self):
|
|
return QQBotAdapter(config={"app_id": "test_app", "app_secret": "test_secret", "dm_policy": "open"})
|
|
|
|
def test_tts_provider_initialized(self, adapter):
|
|
assert adapter._tts_provider is not None
|
|
assert adapter._tts_provider._default_voice == "zh-CN-XiaoxiaoNeural"
|
|
|
|
def test_stt_provider_initialized(self, adapter):
|
|
assert adapter._stt_provider is not None
|
|
assert adapter._stt_provider._provider == "builtin"
|
|
|
|
def test_audio_format_policy_initialized(self, adapter):
|
|
assert adapter._audio_format_policy is not None
|
|
assert adapter._audio_format_policy.transcode_enabled is True
|
|
|
|
def test_tts_custom_config(self):
|
|
adapter = QQBotAdapter(config={
|
|
"app_id": "test_app",
|
|
"app_secret": "test_secret",
|
|
"tts_default_voice": "en-US-JennyNeural",
|
|
"tts_default_format": "wav",
|
|
})
|
|
assert adapter._tts_provider._default_voice == "en-US-JennyNeural"
|
|
|
|
def test_stt_custom_config(self):
|
|
adapter = QQBotAdapter(config={
|
|
"app_id": "test_app",
|
|
"app_secret": "test_secret",
|
|
"stt_provider": "azure",
|
|
"stt_api_key": "key123",
|
|
"stt_region": "eastasia",
|
|
})
|
|
assert adapter._stt_provider._provider == "azure"
|
|
assert adapter._stt_provider._api_key == "key123"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_send_tts_voice_without_client(self, adapter):
|
|
result = await adapter.send_tts_voice("chat_001", "Hello voice")
|
|
assert result.success is False
|
|
assert "Client not initialized" in result.error
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_send_tts_voice_success(self, adapter):
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
from yuxi.channels.models import DeliveryResult
|
|
|
|
adapter._http_client = MagicMock()
|
|
adapter._token_manager = MagicMock()
|
|
adapter._token_manager.get_token = AsyncMock(return_value="mock_token")
|
|
adapter._token_manager.api_base = "https://api.test.com"
|
|
|
|
with patch.object(adapter._tts_provider, "_synthesize_builtin") as mock_synth:
|
|
mock_synth.return_value = b"fake_voice_data"
|
|
|
|
with patch("yuxi.channels.adapters.qqbot.voice_send.send_voice") as mock_send_voice:
|
|
mock_send_voice.return_value = DeliveryResult(success=True, message_id="msg_voice")
|
|
|
|
result = await adapter.send_tts_voice("group_openid123", "Hello")
|
|
assert result.success is True
|
|
assert result.message_id == "msg_voice"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_transcribe_voice_returns_text(self, adapter):
|
|
from unittest.mock import patch
|
|
|
|
with patch.object(adapter._stt_provider, "_transcribe_builtin") as mock_stt:
|
|
mock_stt.return_value = "转写结果"
|
|
|
|
result = await adapter.transcribe_voice(b"fake_audio")
|
|
assert result == "转写结果" |