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

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 == "转写结果"