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

396 lines
14 KiB
Python

from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock
import pytest
from yuxi.channels.adapters.imessage.pairing import (
IMessagePairingManager,
PairingRequest,
)
from yuxi.channels.adapters.imessage.security import (
DmPolicy,
GroupPolicy,
IMessageSecurityPolicy,
SecurityCheckResult,
_match_entry,
_normalize_handle,
_strip_prefix,
resolve_dm_policy,
resolve_group_policy,
)
class TestPairingManager:
@pytest.fixture
def pm(self, tmp_path):
return IMessagePairingManager(storage_dir=str(tmp_path), account_id="test")
def test_generate_code_length(self, pm):
code = pm.generate_code()
assert len(code) == 8
def test_generate_code_charset(self, pm):
code = pm.generate_code()
from yuxi.channels.adapters.imessage.pairing import PAIRING_CODE_CHARS
for c in code:
assert c in PAIRING_CODE_CHARS
def test_generate_code_uniqueness(self, pm):
codes = {pm.generate_code() for _ in range(50)}
assert len(codes) >= 45
def test_request_pairing_returns_code(self, pm):
code = pm.request_pairing("+8613800138000", "user001", "chat_001")
assert code is not None
assert len(code) == 8
def test_request_pairing_rate_limit(self, pm):
pm.request_pairing("+8613800138000", "user001", "chat_001")
code2 = pm.request_pairing("+8613800138000", "user001", "chat_001")
assert code2 is None
def test_request_pairing_already_approved(self, pm):
pm._approved.add("user001")
code = pm.request_pairing("+8613800138000", "user001", "chat_001")
assert code is None
def test_approve_success(self, pm):
code = pm.request_pairing("+8613800138000", "user001", "chat_001")
result = pm.approve(code)
assert result is True
assert pm.is_approved("user001")
def test_approve_invalid_code(self, pm):
result = pm.approve("INVALID")
assert result is False
def test_reject_success(self, pm):
code = pm.request_pairing("+8613800138000", "user001", "chat_001")
result = pm.reject(code)
assert result is True
def test_reject_invalid_code(self, pm):
result = pm.reject("INVALID")
assert result is False
def test_list_pending(self, pm):
pm.request_pairing("+8613800138000", "user001", "chat_001")
pending = pm.list_pending()
assert len(pending) == 1
assert pending[0]["sender_handle"] == "+8613800138000"
def test_is_paired(self, pm):
assert pm.is_paired("user001") is False
pm._approved.add("user001")
assert pm.is_paired("user001") is True
def test_revoke_approval(self, pm):
pm._approved.add("user001")
assert pm.revoke_approval("user001") is True
assert pm.is_approved("user001") is False
def test_revoke_nonexistent(self, pm):
assert pm.revoke_approval("nobody") is False
def test_get_paired_handles(self, pm):
pm._approved.add("user001")
pm._approved.add("user002")
handles = pm.get_paired_handles()
assert "user001" in handles
assert "user002" in handles
@pytest.mark.asyncio
async def test_save_and_load_allowlist(self, pm):
pm._approved.add("user001")
pm._approved.add("user002")
await pm.save_allowlist()
pm2 = IMessagePairingManager(storage_dir=pm._storage_dir, account_id="test")
loaded = await pm2.load_allowlist()
assert "user001" in loaded
assert "user002" in loaded
@pytest.mark.asyncio
async def test_save_and_load_pending(self, pm):
code = pm.request_pairing("+8613800138000", "user001", "chat_001")
await pm.save_pending()
pm2 = IMessagePairingManager(storage_dir=pm._storage_dir, account_id="test")
await pm2.load_pending()
assert len(pm2._pending) == 1
@pytest.mark.asyncio
async def test_load_nonexistent_file(self, pm):
data = await pm.load_allowlist()
assert data == []
def test_config_writes_disabled(self, pm):
pm._config_writes = False
pm._approved.add("user001")
pm._schedule_save()
class TestSecurityPolicy:
@pytest.fixture
def policy_open(self):
return IMessageSecurityPolicy(
{"dmPolicy": "open", "groupPolicy": "open", "allowFrom": []},
config_writes=False,
)
@pytest.fixture
def policy_allowlist(self):
return IMessageSecurityPolicy(
{
"dmPolicy": "allowlist",
"groupPolicy": "allowlist",
"allowFrom": ["+8613800138000"],
"groupAllowFrom": ["iMessage;-;group001"],
},
config_writes=False,
)
@pytest.fixture
def policy_pairing(self):
return IMessageSecurityPolicy(
{"dmPolicy": "pairing", "allowFrom": []},
config_writes=False,
)
@pytest.fixture
def policy_disabled(self):
return IMessageSecurityPolicy(
{"dmPolicy": "disabled", "groupPolicy": "disabled", "allowFrom": []},
config_writes=False,
)
def test_check_dm_open(self, policy_open):
result = policy_open.check_dm_access("anyone")
assert result.allowed is True
def test_check_dm_disabled(self, policy_disabled):
result = policy_disabled.check_dm_access("anyone")
assert result.allowed is False
assert result.reject_reason == "dm_disabled"
def test_check_dm_allowlist_match(self, policy_allowlist):
result = policy_allowlist.check_dm_access("+8613800138000")
assert result.allowed is True
def test_check_dm_allowlist_no_match(self, policy_allowlist):
result = policy_allowlist.check_dm_access("+8613900139000")
assert result.allowed is False
assert result.reject_reason == "not_in_dm_allowlist"
def test_check_dm_pairing(self, policy_pairing):
result = policy_pairing.check_dm_access("anyone")
assert result.allowed is True
def test_check_group_open(self, policy_open):
result = policy_open.check_group_access("iMessage;-;group001")
assert result.allowed is True
def test_check_group_disabled(self, policy_disabled):
result = policy_disabled.check_group_access("iMessage;-;group001")
assert result.allowed is False
def test_check_group_allowlist_match(self, policy_allowlist):
result = policy_allowlist.check_group_access("iMessage;-;group001")
assert result.allowed is True
def test_check_group_allowlist_no_match(self, policy_allowlist):
result = policy_allowlist.check_group_access("iMessage;-;group999")
assert result.allowed is False
def test_check_group_override_disabled(self):
policy = IMessageSecurityPolicy(
{
"groupPolicy": "open",
"groups": {"iMessage;-;group001": {"enabled": False}},
},
config_writes=False,
)
result = policy.check_group_access("iMessage;-;group001")
assert result.allowed is False
def test_check_mention_not_required(self, policy_allowlist):
result = policy_allowlist.check_mention_required(
"iMessage;-;group001", None, "bot_handle"
)
assert result.allowed is True
def test_check_mention_required_matched(self):
policy = IMessageSecurityPolicy(
{"requireMention": True, "allowFrom": []},
config_writes=False,
)
mentions = MagicMock()
mentions.mentioned_user_ids = ["bot_handle"]
result = policy.check_mention_required(
"iMessage;-;group001", mentions, "+bot_handle"
)
assert result.allowed is True
def test_check_mention_required_not_matched(self):
policy = IMessageSecurityPolicy(
{"requireMention": True, "allowFrom": []},
config_writes=False,
)
mentions = MagicMock()
mentions.mentioned_user_ids = ["other_user"]
result = policy.check_mention_required(
"iMessage;-;group001", mentions, "bot_handle"
)
assert result.allowed is False
def test_collect_warnings_open_dm(self, policy_open):
warnings = policy_open.collect_warnings()
assert any("dmPolicy='open'" in w for w in warnings)
def test_collect_warnings_empty_allowlist(self):
policy = IMessageSecurityPolicy(
{"dmPolicy": "allowlist", "allowFrom": []},
config_writes=False,
)
warnings = policy.collect_warnings()
assert len(warnings) > 0
def test_add_to_allow_list(self, policy_allowlist):
policy_allowlist.add_to_allow_list("+8613900139000")
assert "8613900139000" in policy_allowlist.allow_list
def test_add_to_allow_list_duplicate(self, policy_allowlist):
policy_allowlist.add_to_allow_list("+8613800138000")
count = policy_allowlist.allow_list.count("8613800138000")
assert count == 1
def test_remove_from_allow_list(self, policy_allowlist):
result = policy_allowlist.remove_from_allow_list("+8613800138000")
assert result is True
assert "8613800138000" not in policy_allowlist.allow_list
def test_add_to_group_allow_list(self, policy_allowlist):
policy_allowlist.add_to_group_allow_list("iMessage;-;group002")
assert "iMessage;-;group002" in policy_allowlist.group_allow_list
def test_remove_from_group_allow_list(self, policy_allowlist):
result = policy_allowlist.remove_from_group_allow_list("iMessage;-;group001")
assert result is True
def test_allow_list_property(self, policy_allowlist):
assert "8613800138000" in policy_allowlist.allow_list
def test_require_mention_property(self, policy_allowlist):
assert policy_allowlist.require_mention is False
def test_resolve_runtime_group_policy(self, policy_allowlist):
result = policy_allowlist.resolve_runtime_group_policy("iMessage;-;group001")
assert result == GroupPolicy.ALLOWLIST
def test_resolve_runtime_group_policy_with_override(self):
policy = IMessageSecurityPolicy(
{
"groupPolicy": "open",
"groups": {"iMessage;-;group001": {"group_policy": "allowlist"}},
},
config_writes=False,
)
result = policy.resolve_runtime_group_policy("iMessage;-;group001")
assert result == GroupPolicy.ALLOWLIST
def test_evaluate_context_visibility_open(self, policy_open):
policy_open._context_visibility_mode = "open"
result = policy_open.evaluate_context_visibility(
"iMessage;-;group001", "anyone"
)
assert result is True
def test_evaluate_context_visibility_allowlist_match(self, policy_allowlist):
result = policy_allowlist.evaluate_context_visibility(
"iMessage;-;group001", "+8613800138000"
)
assert result is True
def test_evaluate_context_visibility_allowlist_no_match(self, policy_allowlist):
result = policy_allowlist.evaluate_context_visibility(
"iMessage;-;group999", "+8613900139000"
)
assert result is False
def test_resolve_pinned_dm_owner(self, policy_allowlist):
owner = policy_allowlist.resolve_pinned_dm_owner()
assert owner == "8613800138000"
def test_resolve_pinned_dm_owner_none(self, policy_open):
owner = policy_open.resolve_pinned_dm_owner()
assert owner is None
def test_should_update_last_route_owner(self, policy_allowlist):
assert policy_allowlist.should_update_last_route("+8613800138000") is True
def test_should_update_last_route_non_owner(self, policy_allowlist):
assert policy_allowlist.should_update_last_route("+8613900139000") is False
def test_should_update_last_route_no_owner(self, policy_open):
assert policy_open.should_update_last_route("anyone") is True
class TestSecurityHelpers:
def test_resolve_dm_policy_explicit(self):
assert resolve_dm_policy({"dmPolicy": "open"}) == DmPolicy.OPEN
def test_resolve_dm_policy_auto_open_star(self):
assert resolve_dm_policy({"allowFrom": ["*"]}) == DmPolicy.OPEN
def test_resolve_dm_policy_auto_allowlist(self):
assert resolve_dm_policy({"allowFrom": ["+8613800138000"]}) == DmPolicy.ALLOWLIST
def test_resolve_group_policy_explicit(self):
assert resolve_group_policy({"groupPolicy": "disabled"}) == GroupPolicy.DISABLED
def test_resolve_group_policy_default(self):
assert resolve_group_policy({}) == GroupPolicy.ALLOWLIST
def test_normalize_handle(self):
assert _normalize_handle(" +8613800138000 ") == "8613800138000"
assert _normalize_handle("test@example.com") == "test@example.com"
def test_strip_prefix_chat_id(self):
value, prefix = _strip_prefix("chat_id:12345")
assert value == "12345"
assert prefix == "chat_id"
def test_strip_prefix_no_match(self):
value, prefix = _strip_prefix("hello")
assert value == "hello"
assert prefix is None
def test_match_entry_wildcard(self):
assert _match_entry("anything", ["*"]) is True
def test_match_entry_chat_guid(self):
assert _match_entry("iMessage;-;group001", ["chat_guid:iMessage;-;group001"]) is True
def test_match_entry_handle(self):
assert _match_entry("+8613800138000", ["8613800138000"]) is True
def test_match_entry_no_match(self):
assert _match_entry("+8613900139000", ["8613800138000"]) is False
@pytest.mark.asyncio
async def test_save_and_load_allowlist(self, tmp_path):
policy = IMessageSecurityPolicy(
{
"allowFrom": ["+8613800138000"],
"groupAllowFrom": ["iMessage;-;group001"],
},
storage_dir=str(tmp_path),
)
await policy.save_allowlist()
policy2 = IMessageSecurityPolicy({}, storage_dir=str(tmp_path))
await policy2.load_allowlist()
assert "8613800138000" in policy2.allow_list
assert "iMessage;-;group001" in policy2.group_allow_list