1. 移除Telegram格式化测试中未使用的导入项 2. 修复Teams测试用例,添加monkeypatch参数并配置通配符开关 3. 更新钉钉适配器测试,替换弃用的流属性检查 4. 修正Twitch规范化测试,更新ROOMSTATE测试逻辑 5. 重构会话映射测试,完善数据库执行结果模拟 6. 格式化Slack块构建测试的长参数调用 7. 修复LINE适配器测试,更新能力断言和异步锁使用 8. 修正Slack会话解析测试,修复聊天类型判断错误 9. 更新能力测试,补充缺失的字段检查 10. 修复Matrix适配器测试,修正位置参数和配置校验逻辑 11. 为飞书分析模块测试添加跳过标记 12. 新增微信能力、限流、链接格式、会话路由等模块的单元测试 13. 修复Twitch适配器导入路径和测试断言 14. 新增Discord Webhook、Nextcloud Talk、Signal多账户等模块的单元测试 15. 修复Manager阶段测试的导入路径 16. 新增iMessage异常和命令处理的单元测试 17. 新增Nostr健康检查和相关模块的单元测试 18. 新增Signal守护进程和SSE重连相关测试
1554 lines
57 KiB
Python
1554 lines
57 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import time
|
|
from unittest.mock import MagicMock, AsyncMock, patch
|
|
|
|
import pytest
|
|
from yuxi.channels.adapters.zalo_oa.commands import (
|
|
CommandRegistry,
|
|
extract_command,
|
|
get_command_list,
|
|
get_command_registry,
|
|
handle_command,
|
|
register_command,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.dedup import EventDeduplicator, EventState
|
|
from yuxi.channels.adapters.zalo_oa.security import (
|
|
DMPolicy,
|
|
check_dm_allowed,
|
|
collect_security_warnings,
|
|
load_allowlist,
|
|
normalize_allow_entry,
|
|
resolve_dm_policy,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.chunking import chunk_text
|
|
from yuxi.channels.adapters.zalo_oa.messaging import (
|
|
build_target_id,
|
|
looks_like_user_id,
|
|
normalize_messaging_target,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.cache import SentMessageCache
|
|
from yuxi.channels.adapters.zalo_oa.pairing import (
|
|
PairingStore,
|
|
build_pairing_message,
|
|
build_pairing_notification,
|
|
build_pairing_success_message,
|
|
PAIRING_CODE_TTL_SEC,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.webhook_ratelimit import (
|
|
WebhookRateLimiter,
|
|
build_rate_limit_key,
|
|
resolve_client_ip,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.approval import (
|
|
build_approval_request,
|
|
check_approval_required,
|
|
normalize_approver_id,
|
|
resolve_approvers,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.audit import AuditEventType, AuditLogger
|
|
from yuxi.channels.adapters.zalo_oa.accounts import (
|
|
get_default_account_id,
|
|
list_account_ids,
|
|
resolve_account,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.webhook_anomaly import (
|
|
WebhookAnomalyTracker,
|
|
ALERT_THRESHOLD_401,
|
|
ALERT_THRESHOLD_400,
|
|
ALERT_THRESHOLD_413,
|
|
ALERT_THRESHOLD_429,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.message_actions import (
|
|
SUPPORTED_ACTIONS,
|
|
build_action_handler,
|
|
describe_actions,
|
|
is_action_supported,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.signature import verify_zalo_oa_signature
|
|
from yuxi.channels.adapters.zalo_oa.session import (
|
|
SessionRouter,
|
|
resolve_thread_key,
|
|
resolve_session_params,
|
|
get_session_router,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.formatter import ZaloOAMessageFormatter
|
|
from yuxi.channels.adapters.zalo_oa.outbound_media import (
|
|
OutboundMediaHost,
|
|
cleanup_media_cache,
|
|
get_media,
|
|
resolve_attachment,
|
|
store_media,
|
|
)
|
|
from yuxi.channels.adapters.zalo_oa.config_schema import ZaloOAConfig, ZaloOAAccountConfig
|
|
from yuxi.channels.adapters.zalo_oa.polling import ZaloOAPoller
|
|
from yuxi.channels.adapters.zalo_oa.voice import ZaloOAVoice
|
|
from yuxi.channels.adapters.zalo_oa.normalizer import SkipMessageError, ZaloOAEventNormalizer
|
|
from yuxi.channels.models import (
|
|
ChannelIdentity,
|
|
ChannelMessage,
|
|
ChannelResponse,
|
|
ChannelType,
|
|
ChatType,
|
|
EventType,
|
|
MessageType,
|
|
)
|
|
|
|
# ==============================================================================
|
|
# Commands
|
|
# ==============================================================================
|
|
|
|
|
|
class TestCommandRegistry:
|
|
def setup_method(self):
|
|
self.registry = CommandRegistry()
|
|
|
|
def test_default_commands_registered(self):
|
|
assert self.registry.is_registered("help")
|
|
assert self.registry.is_registered("status")
|
|
assert self.registry.is_registered("ping")
|
|
|
|
def test_register_new_command(self):
|
|
self.registry.register("echo", "Echo text back")
|
|
assert self.registry.is_registered("echo")
|
|
|
|
def test_unregister_command(self):
|
|
self.registry.register("temp_cmd", "Temporary")
|
|
self.registry.unregister("temp_cmd")
|
|
assert not self.registry.is_registered("temp_cmd")
|
|
|
|
def test_cannot_unregister_builtin_commands(self):
|
|
self.registry.unregister("help")
|
|
assert self.registry.is_registered("help")
|
|
self.registry.unregister("status")
|
|
assert self.registry.is_registered("status")
|
|
self.registry.unregister("ping")
|
|
assert self.registry.is_registered("ping")
|
|
|
|
def test_get_list_returns_all_commands(self):
|
|
cmd_list = self.registry.get_list()
|
|
names = [c["command"] for c in cmd_list]
|
|
assert "/help" in names
|
|
assert "/status" in names
|
|
assert "/ping" in names
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_unknown_command(self):
|
|
result = await self.registry.handle("unknown", "", {})
|
|
assert result is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_help(self):
|
|
result = await self.registry.handle("help", "", {})
|
|
assert result
|
|
assert "Available commands:" in result
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_status(self):
|
|
context = {"oa_name": "TestOA", "status": "connected", "dm_policy": "open"}
|
|
result = await self.registry.handle("status", "", context)
|
|
assert "TestOA" in result
|
|
assert "connected" in result
|
|
assert "open" in result
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_ping(self):
|
|
result = await self.registry.handle("ping", "", {})
|
|
assert result == "pong"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_with_extra_handlers(self):
|
|
async def custom_handler(args, ctx):
|
|
return f"custom:{args}"
|
|
|
|
result = await self.registry.handle("custom_cmd", "hello", {}, {"custom_cmd": custom_handler})
|
|
assert result == "custom:hello"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_with_args(self):
|
|
self.registry.register("echo", "Echo text back")
|
|
result = await self.registry.handle("echo", "hello world", {})
|
|
assert result == "Command /echo executed"
|
|
|
|
def test_case_insensitive_lookup(self):
|
|
self.registry.register("UpperCase", "Test")
|
|
assert self.registry.is_registered("UpperCase")
|
|
assert self.registry.is_registered("UPPERCASE")
|
|
|
|
def test_case_insensitive_lookup_uppercase(self):
|
|
self.registry.register("echo", "Echo text back")
|
|
assert self.registry.is_registered("ECHO")
|
|
assert self.registry.is_registered("Echo")
|
|
|
|
|
|
class TestExtractCommand:
|
|
def test_extract_slash_command(self):
|
|
cmd, args = extract_command("/help")
|
|
assert cmd == "help"
|
|
assert args == ""
|
|
|
|
def test_extract_slash_command_with_args(self):
|
|
cmd, args = extract_command("/echo hello world")
|
|
assert cmd == "echo"
|
|
assert args == "hello world"
|
|
|
|
def test_extract_exclamation_command(self):
|
|
cmd, args = extract_command("!status")
|
|
assert cmd == "status"
|
|
assert args == ""
|
|
|
|
def test_extract_exclamation_with_args(self):
|
|
cmd, args = extract_command("!ping test arg")
|
|
assert cmd == "ping"
|
|
assert args == "test arg"
|
|
|
|
def test_no_command_prefix(self):
|
|
cmd, args = extract_command("hello world")
|
|
assert cmd is None
|
|
assert args == "hello world"
|
|
|
|
def test_no_command_prefix_leading_space(self):
|
|
cmd, args = extract_command(" hello")
|
|
assert cmd is None
|
|
assert args == "hello"
|
|
|
|
def test_empty_string(self):
|
|
cmd, args = extract_command("")
|
|
assert cmd is None
|
|
assert args == ""
|
|
|
|
def test_only_slash(self):
|
|
cmd, args = extract_command("/")
|
|
assert cmd == ""
|
|
assert args == ""
|
|
|
|
def test_only_slash_with_space(self):
|
|
cmd, args = extract_command("/ ")
|
|
assert cmd == ""
|
|
assert args == ""
|
|
|
|
def test_command_case_insensitive(self):
|
|
cmd, args = extract_command("/HELP")
|
|
assert cmd == "help"
|
|
|
|
|
|
class TestGlobalCommandFunctions:
|
|
@pytest.mark.asyncio
|
|
async def test_handle_command_global(self):
|
|
context = {"oa_name": "TestOA"}
|
|
result = await handle_command("ping", "", context)
|
|
assert result == "pong"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_command_list_global(self):
|
|
cmd_list = get_command_list()
|
|
assert len(cmd_list) >= 3
|
|
|
|
def test_register_command_global(self):
|
|
register_command("global_test", "Test global registration")
|
|
registry = get_command_registry()
|
|
assert registry.is_registered("global_test")
|
|
registry.unregister("global_test")
|
|
|
|
|
|
# ==============================================================================
|
|
# Deduplication
|
|
# ==============================================================================
|
|
|
|
|
|
class TestEventDeduplicator:
|
|
def setup_method(self):
|
|
self.dedup = EventDeduplicator(window_ms=60000)
|
|
|
|
def test_claim_first_time(self):
|
|
payload = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
assert self.dedup.claim(payload) is True
|
|
|
|
def test_claim_duplicate_within_window(self):
|
|
payload = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m2"}}
|
|
assert self.dedup.claim(payload) is True
|
|
assert self.dedup.claim(payload) is False
|
|
|
|
def test_different_senders_not_deduplicated(self):
|
|
p1 = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
p2 = {"event_name": "user_send_text", "sender": {"id": "u2"}, "message": {"msg_id": "m1"}}
|
|
assert self.dedup.claim(p1) is True
|
|
assert self.dedup.claim(p2) is True
|
|
|
|
def test_different_msg_ids_not_deduplicated(self):
|
|
p1 = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
p2 = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m2"}}
|
|
assert self.dedup.claim(p1) is True
|
|
assert self.dedup.claim(p2) is True
|
|
|
|
def test_commit_changes_state(self):
|
|
payload = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
self.dedup.claim(payload)
|
|
self.dedup.commit(payload)
|
|
assert self.dedup.is_committed(payload) is True
|
|
|
|
def test_is_committed_before_commit(self):
|
|
payload = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
self.dedup.claim(payload)
|
|
assert self.dedup.is_committed(payload) is False
|
|
|
|
def test_release_removes_event(self):
|
|
payload = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
self.dedup.claim(payload)
|
|
self.dedup.release(payload)
|
|
assert self.dedup.is_committed(payload) is False
|
|
|
|
def test_claim_after_release(self):
|
|
payload = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
assert self.dedup.claim(payload) is True
|
|
self.dedup.release(payload)
|
|
assert self.dedup.claim(payload) is True
|
|
|
|
def test_window_expiry_allows_reclaim(self):
|
|
dedup = EventDeduplicator(window_ms=1)
|
|
payload = {"event_name": "user_send_text", "sender": {"id": "u1"}, "message": {"msg_id": "m1"}}
|
|
assert dedup.claim(payload) is True
|
|
time.sleep(0.002)
|
|
assert dedup.claim(payload) is True
|
|
|
|
def test_events_without_sender(self):
|
|
payload = {"event_name": "test_event", "sender": {}, "message": {"msg_id": ""}}
|
|
assert self.dedup.claim(payload) is True
|
|
|
|
def test_events_without_message(self):
|
|
payload = {"event_name": "test_event", "sender": {"id": "u1"}}
|
|
assert self.dedup.claim(payload) is True
|
|
|
|
|
|
# ==============================================================================
|
|
# Security
|
|
# ==============================================================================
|
|
|
|
|
|
class TestDMPolicy:
|
|
def test_open_policy_allows_all(self):
|
|
assert check_dm_allowed("any_user", DMPolicy.OPEN, set()) is True
|
|
|
|
def test_disabled_policy_denies_all(self):
|
|
assert check_dm_allowed("any_user", DMPolicy.DISABLED, set()) is False
|
|
|
|
def test_allowlist_policy_match(self):
|
|
assert check_dm_allowed("user1", DMPolicy.ALLOWLIST, {"user1"}) is True
|
|
|
|
def test_allowlist_policy_no_match(self):
|
|
assert check_dm_allowed("user2", DMPolicy.ALLOWLIST, {"user1"}) is False
|
|
|
|
def test_allowlist_policy_wildcard(self):
|
|
assert check_dm_allowed("any_user", DMPolicy.ALLOWLIST, {"*"}) is True
|
|
|
|
def test_pairing_policy_match(self):
|
|
assert check_dm_allowed("user1", DMPolicy.PAIRING, {"user1"}) is True
|
|
|
|
def test_pairing_policy_no_match(self):
|
|
assert check_dm_allowed("user2", DMPolicy.PAIRING, {"user1"}) is False
|
|
|
|
def test_pairing_policy_wildcard(self):
|
|
assert check_dm_allowed("any_user", DMPolicy.PAIRING, {"*"}) is True
|
|
|
|
def test_resolve_dm_policy_open(self):
|
|
assert resolve_dm_policy({"dm_policy": "open"}) == DMPolicy.OPEN
|
|
|
|
def test_resolve_dm_policy_pairing(self):
|
|
assert resolve_dm_policy({"dm_policy": "pairing"}) == DMPolicy.PAIRING
|
|
|
|
def test_resolve_dm_policy_allowlist(self):
|
|
assert resolve_dm_policy({"dm_policy": "allowlist"}) == DMPolicy.ALLOWLIST
|
|
|
|
def test_resolve_dm_policy_disabled(self):
|
|
assert resolve_dm_policy({"dm_policy": "disabled"}) == DMPolicy.DISABLED
|
|
|
|
def test_resolve_dm_policy_default(self):
|
|
assert resolve_dm_policy({}) == DMPolicy.OPEN
|
|
|
|
def test_resolve_dm_policy_unknown(self):
|
|
assert resolve_dm_policy({"dm_policy": "unknown"}) == DMPolicy.OPEN
|
|
|
|
def test_resolve_dm_policy_dmPolicy_fallback(self):
|
|
assert resolve_dm_policy({"dmPolicy": "disabled"}) == DMPolicy.DISABLED
|
|
|
|
|
|
class TestLoadAllowlist:
|
|
def test_normal_allowlist(self):
|
|
config = {"allowFrom": ["user1", "user2"]}
|
|
result = load_allowlist(config)
|
|
assert result == {"user1", "user2"}
|
|
|
|
def test_empty_allowlist(self):
|
|
result = load_allowlist({})
|
|
assert result == set()
|
|
|
|
def test_wildcard_only(self):
|
|
config = {"allowFrom": ["*"]}
|
|
result = load_allowlist(config)
|
|
assert result == {"*"}
|
|
|
|
def test_wildcard_precedence(self):
|
|
config = {"allowFrom": ["user1", "*", "user2"]}
|
|
result = load_allowlist(config)
|
|
assert result == {"*"}
|
|
|
|
def test_prefix_normalization_zoa(self):
|
|
config = {"allowFrom": ["zoa:user1"]}
|
|
result = load_allowlist(config)
|
|
assert "user1" in result
|
|
|
|
def test_prefix_normalization_zalo_oa(self):
|
|
config = {"allowFrom": ["zalo_oa:user2"]}
|
|
result = load_allowlist(config)
|
|
assert "user2" in result
|
|
|
|
def test_non_list_allow_from(self):
|
|
config = {"allowFrom": "not_a_list"}
|
|
result = load_allowlist(config)
|
|
assert result == set()
|
|
|
|
|
|
class TestNormalizeAllowEntry:
|
|
def test_zoa_prefix(self):
|
|
assert normalize_allow_entry("zoa:user1") == "user1"
|
|
|
|
def test_zalo_oa_prefix(self):
|
|
assert normalize_allow_entry("zalo_oa:user1") == "user1"
|
|
|
|
def test_no_prefix(self):
|
|
assert normalize_allow_entry("user1") == "user1"
|
|
|
|
def test_case_insensitive_prefix(self):
|
|
assert normalize_allow_entry("ZOA:User1") == "User1"
|
|
assert normalize_allow_entry("ZALO_OA:User1") == "User1"
|
|
|
|
|
|
class TestCollectSecurityWarnings:
|
|
def test_open_policy_warning(self):
|
|
warnings = collect_security_warnings({"dm_policy": "open"})
|
|
assert any(w["type"] == "dm_policy_open" for w in warnings)
|
|
|
|
def test_empty_allowlist_warning(self):
|
|
warnings = collect_security_warnings({"dm_policy": "allowlist", "allowFrom": []})
|
|
assert any(w["type"] == "empty_allowlist" for w in warnings)
|
|
|
|
def test_pairing_no_allowlist_info(self):
|
|
warnings = collect_security_warnings({"dm_policy": "pairing", "allowFrom": []})
|
|
assert any(w["type"] == "pairing_no_preset" for w in warnings)
|
|
|
|
def test_wildcard_info(self):
|
|
warnings = collect_security_warnings({"dm_policy": "allowlist", "allowFrom": ["*"]})
|
|
assert any(w["type"] == "allowlist_wildcard" for w in warnings)
|
|
|
|
def test_no_webhook_url_info(self):
|
|
warnings = collect_security_warnings({"dm_policy": "open"})
|
|
assert any(w["type"] == "webhook_not_configured" for w in warnings)
|
|
|
|
def test_no_mac_key_warning(self):
|
|
warnings = collect_security_warnings({"dm_policy": "open"})
|
|
assert any(w["type"] == "webhook_mac_key_missing" for w in warnings)
|
|
|
|
def test_with_webhook_and_mac_key_no_warnings(self):
|
|
warnings = collect_security_warnings(
|
|
{"dm_policy": "allowlist", "allowFrom": ["user1"], "webhook": {"url": "https://ex.com", "mac_key": "key1"}}
|
|
)
|
|
assert not any(w["type"] in ("webhook_not_configured", "webhook_mac_key_missing") for w in warnings)
|
|
|
|
|
|
# ==============================================================================
|
|
# Chunking
|
|
# ==============================================================================
|
|
|
|
|
|
class TestChunkText:
|
|
def test_short_text_no_chunking(self):
|
|
result = chunk_text("hello", 100)
|
|
assert result == ["hello"]
|
|
|
|
def test_chunk_by_paragraph_boundary(self):
|
|
text = "A" * 100 + "\n\n" + "B" * 100
|
|
result = chunk_text(text, 150)
|
|
assert len(result) == 2
|
|
assert result[0].strip() == "A" * 100
|
|
assert result[1].strip() == "B" * 100
|
|
|
|
def test_chunk_by_newline_and_comma(self):
|
|
text = "Hello, World. This is" + "X" * 300
|
|
result = chunk_text(text, 300)
|
|
assert len(result) >= 1
|
|
|
|
def test_chunk_by_space(self):
|
|
text = "A " * 200 + "B" * 200
|
|
result = chunk_text(text, 200)
|
|
assert len(result) > 1
|
|
|
|
def test_chunk_hard_split(self):
|
|
text = "X" * 500
|
|
result = chunk_text(text, 100)
|
|
assert len(result) == 5
|
|
for chunk in result:
|
|
assert len(chunk) <= 100
|
|
|
|
def test_empty_string(self):
|
|
result = chunk_text("", 100)
|
|
assert result == [""]
|
|
|
|
def test_max_chars_equals_length(self):
|
|
text = "hello world"
|
|
result = chunk_text(text, len(text))
|
|
assert result == ["hello world"]
|
|
|
|
def test_chunk_by_sentence_boundary(self):
|
|
text = "Hello there. How are you? " + "X" * 200
|
|
result = chunk_text(text, 200)
|
|
assert len(result) >= 1
|
|
|
|
def test_all_chunks_within_limit(self):
|
|
text = "A" * 2500
|
|
result = chunk_text(text, 500)
|
|
for chunk in result:
|
|
assert len(chunk) <= 500
|
|
|
|
def test_multiple_paragraphs(self):
|
|
paragraphs = ["P" * 100 for _ in range(5)]
|
|
text = "\n\n".join(paragraphs)
|
|
result = chunk_text(text, 250)
|
|
assert len(result) >= 3
|
|
|
|
|
|
# ==============================================================================
|
|
# Messaging
|
|
# ==============================================================================
|
|
|
|
|
|
class TestMessaging:
|
|
def test_normalize_zalo_oa_prefix(self):
|
|
assert normalize_messaging_target("zalo_oa:12345") == "12345"
|
|
|
|
def test_normalize_zoa_prefix(self):
|
|
assert normalize_messaging_target("zoa:12345") == "12345"
|
|
|
|
def test_normalize_zl_prefix(self):
|
|
assert normalize_messaging_target("zl:12345") == "12345"
|
|
|
|
def test_normalize_case_insensitive(self):
|
|
assert normalize_messaging_target("ZALO_OA:12345") == "12345"
|
|
|
|
def test_no_prefix(self):
|
|
assert normalize_messaging_target("12345") == "12345"
|
|
|
|
def test_empty_string(self):
|
|
assert normalize_messaging_target("") == ""
|
|
|
|
def test_looks_like_user_id_numeric(self):
|
|
assert looks_like_user_id("123456") is True
|
|
|
|
def test_looks_like_user_id_zero(self):
|
|
assert looks_like_user_id("0") is True
|
|
|
|
def test_looks_like_user_id_not_numeric(self):
|
|
assert looks_like_user_id("abc123") is False
|
|
|
|
def test_looks_like_user_id_empty(self):
|
|
assert looks_like_user_id("") is False
|
|
|
|
def test_build_target_id(self):
|
|
assert build_target_id("12345") == "zalo_oa:12345"
|
|
|
|
|
|
# ==============================================================================
|
|
# Sent Message Cache
|
|
# ==============================================================================
|
|
|
|
|
|
class TestSentMessageCache:
|
|
def setup_method(self):
|
|
self.cache = SentMessageCache(ttl_sec=3600)
|
|
|
|
def test_put_and_get(self):
|
|
self.cache.put("msg1", "user1", "hello")
|
|
entry = self.cache.get("msg1")
|
|
assert entry["message_id"] == "msg1"
|
|
assert entry["recipient"] == "user1"
|
|
assert entry["content"] == "hello"
|
|
|
|
def test_get_nonexistent(self):
|
|
assert self.cache.get("nonexistent") is None
|
|
|
|
def test_remove(self):
|
|
self.cache.put("msg1", "user1", "hello")
|
|
self.cache.remove("msg1")
|
|
assert self.cache.get("msg1") is None
|
|
|
|
def test_ttl_expiry(self):
|
|
cache = SentMessageCache(ttl_sec=0)
|
|
cache.put("msg1", "user1", "hello")
|
|
assert cache.get("msg1") is None
|
|
|
|
def test_metadata_storage(self):
|
|
self.cache.put("msg1", "user1", "hello", metadata={"key": "value"})
|
|
entry = self.cache.get("msg1")
|
|
assert entry["metadata"] == {"key": "value"}
|
|
|
|
def test_size_property(self):
|
|
assert self.cache.size == 0
|
|
self.cache.put("msg1", "user1", "hello")
|
|
assert self.cache.size == 1
|
|
|
|
def test_multiple_entries(self):
|
|
for i in range(5):
|
|
self.cache.put(f"msg{i}", f"user{i}", f"content{i}")
|
|
assert self.cache.size == 5
|
|
|
|
def test_cleanup_on_put(self):
|
|
cache = SentMessageCache(ttl_sec=0)
|
|
cache.put("msg1", "user1", "hello")
|
|
cache.put("msg2", "user2", "world")
|
|
assert cache.size == 0
|
|
|
|
|
|
# ==============================================================================
|
|
# Pairing
|
|
# ==============================================================================
|
|
|
|
|
|
class TestPairingStore:
|
|
def setup_method(self):
|
|
self.store = PairingStore()
|
|
|
|
def test_generate_code_length(self):
|
|
code = self.store.generate_code("user1")
|
|
assert len(code) == 8
|
|
assert code == code.upper()
|
|
|
|
def test_verify_correct_code(self):
|
|
code = self.store.generate_code("user1")
|
|
assert self.store.verify_code(code, "user1") is True
|
|
|
|
def test_verify_wrong_user(self):
|
|
code = self.store.generate_code("user1")
|
|
assert self.store.verify_code(code, "user2") is False
|
|
|
|
def test_verify_wrong_code(self):
|
|
self.store.generate_code("user1")
|
|
assert self.store.verify_code("WRONG001", "user1") is False
|
|
|
|
def test_code_cannot_be_used_twice(self):
|
|
code = self.store.generate_code("user1")
|
|
assert self.store.verify_code(code, "user1") is True
|
|
assert self.store.verify_code(code, "user1") is False
|
|
|
|
def test_expired_code(self):
|
|
code = self.store.generate_code("user1")
|
|
self.store._codes[code]["created_at"] = time.time() - PAIRING_CODE_TTL_SEC - 1
|
|
assert self.store.verify_code(code, "user1") is False
|
|
|
|
def test_cleanup_expired(self):
|
|
code1 = self.store.generate_code("user1")
|
|
code2 = self.store.generate_code("user2")
|
|
self.store._codes[code2]["created_at"] = time.time() - PAIRING_CODE_TTL_SEC - 1
|
|
self.store.cleanup_expired()
|
|
assert code2 not in self.store._codes
|
|
assert code1 in self.store._codes
|
|
|
|
def test_generate_multiple_codes(self):
|
|
codes = [self.store.generate_code(f"user{i}") for i in range(5)]
|
|
assert len(set(codes)) == 5
|
|
|
|
|
|
class TestPairingMessages:
|
|
def test_build_pairing_message(self):
|
|
msg = build_pairing_message("ABC12345", "TestOA")
|
|
assert "Hello from TestOA" in msg
|
|
assert "ABC12345" in msg
|
|
assert "expires" in msg
|
|
|
|
def test_build_pairing_message_no_oa_name(self):
|
|
msg = build_pairing_message("ABC12345")
|
|
assert "Welcome" in msg
|
|
assert "ABC12345" in msg
|
|
|
|
def test_build_pairing_notification(self):
|
|
msg = build_pairing_notification("user1", "TestOA")
|
|
assert "user1" in msg
|
|
assert "TestOA" in msg
|
|
|
|
def test_build_pairing_notification_no_name(self):
|
|
msg = build_pairing_notification("user1", "")
|
|
assert "user1" in msg
|
|
|
|
def test_build_pairing_success(self):
|
|
msg = build_pairing_success_message("TestOA")
|
|
assert "Pairing successful" in msg
|
|
assert "TestOA" in msg
|
|
|
|
|
|
# ==============================================================================
|
|
# Webhook Rate Limiter
|
|
# ==============================================================================
|
|
|
|
|
|
class TestWebhookRateLimiter:
|
|
def setup_method(self):
|
|
self.limiter = WebhookRateLimiter(window_ms=60000, max_requests=5)
|
|
|
|
def test_initial_allow(self):
|
|
for i in range(5):
|
|
assert self.limiter.is_allowed(f"key_{i}") is True
|
|
|
|
def test_rate_limit_exceeded(self):
|
|
key = "test_key"
|
|
for _ in range(5):
|
|
assert self.limiter.is_allowed(key) is True
|
|
assert self.limiter.is_allowed(key) is False
|
|
|
|
def test_different_keys_independent(self):
|
|
for _ in range(5):
|
|
assert self.limiter.is_allowed("key_a") is True
|
|
assert self.limiter.is_allowed("key_b") is True
|
|
|
|
def test_cleanup_removes_empty_buckets(self):
|
|
assert self.limiter.is_allowed("temp_key") is True
|
|
self.limiter.cleanup()
|
|
|
|
def test_window_expiry(self):
|
|
limiter = WebhookRateLimiter(window_ms=1, max_requests=3)
|
|
key = "fast_key"
|
|
for _ in range(3):
|
|
assert limiter.is_allowed(key) is True
|
|
assert limiter.is_allowed(key) is False
|
|
time.sleep(0.002)
|
|
assert limiter.is_allowed(key) is True
|
|
|
|
def test_large_max_requests(self):
|
|
limiter = WebhookRateLimiter(max_requests=1000)
|
|
for i in range(500):
|
|
assert limiter.is_allowed(f"key_{i}") is True
|
|
|
|
|
|
class TestRateLimitKeyAndIP:
|
|
def test_build_key_with_ip(self):
|
|
key = build_rate_limit_key("user1", "192.168.1.1")
|
|
assert key == "user1:192.168.1.1"
|
|
|
|
def test_build_key_without_ip(self):
|
|
key = build_rate_limit_key("user1")
|
|
assert key == "user1"
|
|
|
|
def test_build_key_empty_ip(self):
|
|
key = build_rate_limit_key("user1", "")
|
|
assert key == "user1"
|
|
|
|
def test_resolve_client_ip_none_headers(self):
|
|
assert resolve_client_ip(None) == ""
|
|
|
|
def test_resolve_client_ip_x_forwarded_for(self):
|
|
headers = {"x-forwarded-for": "10.0.0.1, 10.0.0.2"}
|
|
assert resolve_client_ip(headers) == "10.0.0.1"
|
|
|
|
def test_resolve_client_ip_x_real_ip(self):
|
|
headers = {"x-real-ip": "10.0.0.5"}
|
|
assert resolve_client_ip(headers) == "10.0.0.5"
|
|
|
|
def test_resolve_client_ip_no_headers(self):
|
|
headers = {}
|
|
assert resolve_client_ip(headers) == ""
|
|
|
|
def test_resolve_client_ip_with_trusted_proxies(self):
|
|
headers = {"x-forwarded-for": "10.0.0.1, 10.0.0.2"}
|
|
result = resolve_client_ip(headers, ["10.0.0.1"])
|
|
assert result == "10.0.0.2"
|
|
|
|
def test_resolve_client_ip_all_trusted(self):
|
|
headers = {"x-forwarded-for": "10.0.0.1"}
|
|
result = resolve_client_ip(headers, ["10.0.0.1"])
|
|
assert result == "10.0.0.1"
|
|
|
|
def test_resolve_client_ip_empty_forwarded(self):
|
|
headers = {"x-forwarded-for": ""}
|
|
assert resolve_client_ip(headers) == ""
|
|
|
|
|
|
# ==============================================================================
|
|
# Approval
|
|
# ==============================================================================
|
|
|
|
|
|
class TestApproval:
|
|
def test_check_approval_required_disabled(self):
|
|
config = {"requireExecApproval": False}
|
|
assert check_approval_required("send", config) is False
|
|
|
|
def test_check_approval_required_enabled(self):
|
|
config = {"requireExecApproval": True}
|
|
assert check_approval_required("send", config) is True
|
|
|
|
def test_check_approval_required_exempt(self):
|
|
config = {"requireExecApproval": True, "approvalExemptActions": ["send"]}
|
|
assert check_approval_required("send", config) is False
|
|
|
|
def test_check_approval_required_not_exempt(self):
|
|
config = {"requireExecApproval": True, "approvalExemptActions": ["ping"]}
|
|
assert check_approval_required("send", config) is True
|
|
|
|
def test_build_approval_request(self):
|
|
result = build_approval_request("send", "user1", {"key": "val"})
|
|
assert result["action"] == "send"
|
|
assert result["requester_id"] == "user1"
|
|
assert result["status"] == "pending"
|
|
assert result["params"] == {"key": "val"}
|
|
|
|
def test_build_approval_request_no_params(self):
|
|
result = build_approval_request("send", "user1")
|
|
assert result["params"] == {}
|
|
|
|
def test_resolve_approvers_empty(self):
|
|
result = resolve_approvers({})
|
|
assert result == []
|
|
|
|
def test_resolve_approvers_from_config(self):
|
|
config = {"approvers": ["user1", "zoa:user2"]}
|
|
result = resolve_approvers(config)
|
|
assert result == ["user1", "user2"]
|
|
|
|
def test_resolve_approvers_from_allowlist(self):
|
|
config = {"allowFrom": ["user1", "zalo_oa:user2"]}
|
|
result = resolve_approvers(config)
|
|
assert result == ["user1", "user2"]
|
|
|
|
def test_resolve_approvers_with_exec_name(self):
|
|
config = {"execApprovers": ["user1", "user2"]}
|
|
result = resolve_approvers(config)
|
|
assert "user1" in result
|
|
|
|
def test_normalize_approver_id_zoa(self):
|
|
assert normalize_approver_id("zoa:user1") == "user1"
|
|
|
|
def test_normalize_approver_id_zalo_oa(self):
|
|
assert normalize_approver_id("zalo_oa:user1") == "user1"
|
|
|
|
def test_normalize_approver_id_no_prefix(self):
|
|
assert normalize_approver_id("user1") == "user1"
|
|
|
|
|
|
# ==============================================================================
|
|
# Audit Logger
|
|
# ==============================================================================
|
|
|
|
|
|
class TestAuditLogger:
|
|
def setup_method(self):
|
|
self.audit = AuditLogger(enabled=True, log_level="info")
|
|
|
|
def test_record_event(self):
|
|
self.audit.record(AuditEventType.MESSAGE_SENT, {"recipient": "user1"})
|
|
assert self.audit.total_entries == 1
|
|
|
|
def test_record_string_event_type(self):
|
|
self.audit.record("custom.event", {"key": "val"})
|
|
assert self.audit.total_entries == 1
|
|
|
|
def test_disabled_logger_ignores_records(self):
|
|
audit = AuditLogger(enabled=False)
|
|
audit.record(AuditEventType.MESSAGE_SENT)
|
|
assert audit.total_entries == 0
|
|
|
|
def test_get_entries_all(self):
|
|
self.audit.record(AuditEventType.MESSAGE_SENT)
|
|
self.audit.record(AuditEventType.MESSAGE_RECEIVED)
|
|
entries = self.audit.get_entries()
|
|
assert len(entries) == 2
|
|
|
|
def test_get_entries_filtered(self):
|
|
self.audit.record(AuditEventType.MESSAGE_SENT)
|
|
self.audit.record(AuditEventType.TOKEN_REFRESHED)
|
|
entries = self.audit.get_entries(event_type=AuditEventType.MESSAGE_SENT)
|
|
assert len(entries) == 1
|
|
assert entries[0]["event_type"] == AuditEventType.MESSAGE_SENT
|
|
|
|
def test_get_summary(self):
|
|
self.audit.record(AuditEventType.MESSAGE_SENT)
|
|
self.audit.record(AuditEventType.MESSAGE_SENT)
|
|
self.audit.record(AuditEventType.AUTH_FAILURE)
|
|
summary = self.audit.get_summary()
|
|
assert summary["total_entries"] == 3
|
|
assert summary["event_counts"]["message.sent"] == 2
|
|
assert summary["event_counts"]["auth.failure"] == 1
|
|
|
|
def test_max_entries_limit(self):
|
|
audit = AuditLogger(max_entries=5)
|
|
for i in range(10):
|
|
audit.record(AuditEventType.MESSAGE_SENT, {"index": i})
|
|
assert audit.total_entries == 5
|
|
|
|
def test_entry_has_timestamp(self):
|
|
self.audit.record(AuditEventType.MESSAGE_SENT)
|
|
entries = self.audit.get_entries()
|
|
assert "timestamp" in entries[0]
|
|
|
|
def test_all_event_types_defined(self):
|
|
assert AuditEventType.MESSAGE_SENT == "message.sent"
|
|
assert AuditEventType.MESSAGE_RECEIVED == "message.received"
|
|
assert AuditEventType.MESSAGE_FAILED == "message.failed"
|
|
assert AuditEventType.MESSAGE_RECALLED == "message.recalled"
|
|
assert AuditEventType.TOKEN_REFRESHED == "token.refreshed"
|
|
assert AuditEventType.TOKEN_REFRESH_FAILED == "token.refresh_failed"
|
|
assert AuditEventType.AUTH_FAILURE == "auth.failure"
|
|
assert AuditEventType.WEBOOK_RECEIVED == "webhook.received"
|
|
assert AuditEventType.WEBOOK_RATE_LIMITED == "webhook.rate_limited"
|
|
assert AuditEventType.PAIRING_CREATED == "pairing.created"
|
|
assert AuditEventType.PAIRING_VERIFIED == "pairing.verified"
|
|
assert AuditEventType.CONFIG_CHANGED == "config.changed"
|
|
assert AuditEventType.CONNECTION_STATE == "connection.state"
|
|
|
|
|
|
# ==============================================================================
|
|
# Accounts
|
|
# ==============================================================================
|
|
|
|
|
|
class TestAccounts:
|
|
def test_list_account_ids_no_accounts(self):
|
|
result = list_account_ids({"app_id": "test"})
|
|
assert result == ["default"]
|
|
|
|
def test_list_account_ids_empty_accounts(self):
|
|
result = list_account_ids({"accounts": {}})
|
|
assert result == ["default"]
|
|
|
|
def test_list_account_ids_with_accounts(self):
|
|
config = {"accounts": {"acct1": {}, "acct2": {}}}
|
|
result = list_account_ids(config)
|
|
assert sorted(result) == ["acct1", "acct2"]
|
|
|
|
def test_list_account_ids_non_dict(self):
|
|
result = list_account_ids({"accounts": "invalid"})
|
|
assert result == ["default"]
|
|
|
|
def test_get_default_account_id_no_default(self):
|
|
result = get_default_account_id({"app_id": "test"})
|
|
assert result == "default"
|
|
|
|
def test_get_default_account_id_with_default(self):
|
|
config = {"defaultAccount": "prod", "accounts": {"prod": {"app_id": "x"}, "staging": {"app_id": "y"}}}
|
|
result = get_default_account_id(config)
|
|
assert result == "prod"
|
|
|
|
def test_get_default_account_id_default_account_key(self):
|
|
config = {"default_account": "staging", "accounts": {"prod": {}, "staging": {}}}
|
|
result = get_default_account_id(config)
|
|
assert result == "staging"
|
|
|
|
def test_resolve_account_no_accounts(self):
|
|
config = {"app_id": "test", "secret_key": "key"}
|
|
result = resolve_account(config)
|
|
assert result["app_id"] == "test"
|
|
|
|
def test_resolve_account_with_accounts(self):
|
|
config = {
|
|
"dm_policy": "open",
|
|
"accounts": {"acct1": {"app_id": "id1", "secret_key": "key1", "name": "Account 1"}},
|
|
}
|
|
result = resolve_account(config)
|
|
assert result["app_id"] == "id1"
|
|
assert result["name"] == "Account 1"
|
|
|
|
def test_resolve_account_with_default_account(self):
|
|
config = {
|
|
"defaultAccount": "prod",
|
|
"accounts": {
|
|
"prod": {"app_id": "prod_id", "name": "Production"},
|
|
"staging": {"app_id": "staging_id", "name": "Staging"},
|
|
},
|
|
}
|
|
result = resolve_account(config)
|
|
assert result["app_id"] == "prod_id"
|
|
|
|
def test_resolve_account_specific_id(self):
|
|
config = {
|
|
"accounts": {
|
|
"prod": {"app_id": "prod_id", "name": "Production"},
|
|
"staging": {"app_id": "staging_id", "name": "Staging"},
|
|
},
|
|
}
|
|
result = resolve_account(config, account_id="staging")
|
|
assert result["app_id"] == "staging_id"
|
|
|
|
def test_resolve_account_top_level_merge(self):
|
|
config = {"dm_policy": "disabled", "webhook": {"url": "https://ex.com"}, "accounts": {"acct1": {"app_id": "id1"}}}
|
|
result = resolve_account(config)
|
|
assert result["dm_policy"] == "disabled"
|
|
assert result["app_id"] == "id1"
|
|
|
|
|
|
# ==============================================================================
|
|
# Webhook Anomaly Tracker
|
|
# ==============================================================================
|
|
|
|
|
|
class TestWebhookAnomalyTracker:
|
|
def setup_method(self):
|
|
self.tracker = WebhookAnomalyTracker()
|
|
|
|
def test_record_anomaly(self):
|
|
self.tracker.record(400, "10.0.0.1", "/webhook", "Bad request")
|
|
assert self.tracker.total_anomalies == 1
|
|
|
|
def test_total_anomalies_count(self):
|
|
for i in range(5):
|
|
self.tracker.record(400, f"10.0.0.{i}")
|
|
assert self.tracker.total_anomalies == 5
|
|
|
|
def test_get_anomalies(self):
|
|
self.tracker.record(401, "10.0.0.1")
|
|
self.tracker.record(413, "10.0.0.2")
|
|
anomalies = self.tracker.get_anomalies()
|
|
assert len(anomalies) == 2
|
|
|
|
def test_get_alerts_initially_empty(self):
|
|
alerts = self.tracker.get_alerts()
|
|
assert isinstance(alerts, list)
|
|
|
|
def test_alert_thresholds_defined(self):
|
|
assert ALERT_THRESHOLD_401 == 10
|
|
assert ALERT_THRESHOLD_400 == 20
|
|
assert ALERT_THRESHOLD_413 == 5
|
|
assert ALERT_THRESHOLD_429 == 10
|
|
|
|
def test_record_with_minimal_args(self):
|
|
self.tracker.record(200)
|
|
assert self.tracker.total_anomalies == 1
|
|
|
|
|
|
# ==============================================================================
|
|
# Message Actions
|
|
# ==============================================================================
|
|
|
|
|
|
class TestMessageActions:
|
|
def test_supported_actions_list(self):
|
|
assert "send" in SUPPORTED_ACTIONS
|
|
assert "sticker" in SUPPORTED_ACTIONS
|
|
assert "list" in SUPPORTED_ACTIONS
|
|
assert "broadcast" in SUPPORTED_ACTIONS
|
|
assert "unsend" in SUPPORTED_ACTIONS
|
|
|
|
def test_all_actions_supported(self):
|
|
for action in SUPPORTED_ACTIONS:
|
|
assert is_action_supported(action) is True
|
|
|
|
def test_unsupported_action(self):
|
|
assert is_action_supported("nonexistent") is False
|
|
|
|
def test_describe_actions_returns_all(self):
|
|
actions = describe_actions()
|
|
action_names = [a["action"] for a in actions]
|
|
for supported in SUPPORTED_ACTIONS:
|
|
assert supported in action_names
|
|
|
|
def test_build_action_handler_send(self):
|
|
formatter = MagicMock()
|
|
formatter._build_text.return_value = {"text": "hello"}
|
|
result = build_action_handler("send", {"recipient_id": "user1", "content": "hello"}, formatter)
|
|
formatter._build_text.assert_called_once_with("user1", "hello")
|
|
|
|
def test_build_action_handler_sticker(self):
|
|
formatter = MagicMock()
|
|
formatter._build_sticker.return_value = {"sticker": "123"}
|
|
result = build_action_handler("sticker", {"recipient_id": "user1", "sticker_id": "123"}, formatter)
|
|
formatter._build_sticker.assert_called_once()
|
|
|
|
def test_build_action_handler_list(self):
|
|
formatter = MagicMock()
|
|
formatter.build_list.return_value = {"list": "test"}
|
|
result = build_action_handler(
|
|
"list", {"recipient_id": "user1", "elements": [{"title": "Item"}], "buttons": None}, formatter
|
|
)
|
|
formatter.build_list.assert_called_once()
|
|
|
|
def test_build_action_handler_broadcast(self):
|
|
result = build_action_handler("broadcast", {"content": "Hello all"}, MagicMock())
|
|
assert result is None
|
|
|
|
def test_build_action_handler_unsend_no_client(self):
|
|
result = build_action_handler("unsend", {"message_id": "m1", "user_id": "u1"}, None, None, client=None)
|
|
assert result is None
|
|
|
|
def test_build_action_handler_unsend_with_client(self):
|
|
result = build_action_handler("unsend", {"message_id": "m1", "user_id": "u1"}, None, None, client="mock")
|
|
assert result["_unsend_message_id"] == "m1"
|
|
assert result["_unsend_user_id"] == "u1"
|
|
|
|
def test_build_action_handler_unknown(self):
|
|
result = build_action_handler("unknown", {}, MagicMock())
|
|
assert result is None
|
|
|
|
|
|
# ==============================================================================
|
|
# Session Router (Extended)
|
|
# ==============================================================================
|
|
|
|
|
|
class TestSessionRouterExtended:
|
|
def setup_method(self):
|
|
self.router = SessionRouter(ttl_sec=3600)
|
|
|
|
def test_resolve_bind_key(self):
|
|
key = self.router.resolve_bind_key("agent1")
|
|
assert key == "agent:agent1:zalo_oa:default"
|
|
|
|
def test_resolve_bind_key_custom_account(self):
|
|
key = self.router.resolve_bind_key("agent1", "prod")
|
|
assert key == "agent:agent1:zalo_oa:prod"
|
|
|
|
def test_bind_and_check(self):
|
|
self.router.bind_agent_account("agent1", "default")
|
|
assert self.router.is_bound("agent1", "default") is True
|
|
|
|
def test_unbind(self):
|
|
self.router.bind_agent_account("agent1", "default")
|
|
self.router.unbind_agent_account("agent1", "default")
|
|
assert self.router.is_bound("agent1", "default") is False
|
|
|
|
def test_binding_ttl_expiry(self):
|
|
router = SessionRouter(ttl_sec=3600)
|
|
router.bind_agent_account("agent1", "default")
|
|
for bind_key, entry in list(router._bindings.items()):
|
|
entry["created_at"] = time.time() - 99999
|
|
assert router.is_bound("agent1", "default") is False
|
|
|
|
def test_session_set_and_get(self):
|
|
thread_key = self.router.resolve_thread_key("agent1", "user1")
|
|
self.router.set_session(thread_key, {"data": "test_value"})
|
|
session = self.router.get_session(thread_key)
|
|
assert session == {"data": "test_value"}
|
|
|
|
def test_session_clear(self):
|
|
thread_key = self.router.resolve_thread_key("agent1", "user1")
|
|
self.router.set_session(thread_key, {"data": "test"})
|
|
self.router.clear_session(thread_key)
|
|
assert self.router.get_session(thread_key) is None
|
|
|
|
def test_get_bound_sessions(self):
|
|
self.router.bind_agent_account("agent1", "default")
|
|
thread_key = self.router.resolve_thread_key("agent1", "user1")
|
|
self.router.set_session(thread_key, {"follower_id": "user1"})
|
|
sessions = self.router.get_bound_sessions("agent1", "default")
|
|
assert len(sessions) >= 1
|
|
assert sessions[0]["follower_id"] == "user1"
|
|
|
|
def test_get_bound_sessions_no_bound(self):
|
|
sessions = self.router.get_bound_sessions("agent1", "default")
|
|
assert sessions == []
|
|
|
|
def test_cleanup_expired_count(self):
|
|
thread_key = self.router.resolve_thread_key("agent1", "expired_user")
|
|
self.router.set_session(thread_key, {"data": "test"})
|
|
self.router._sessions[thread_key]["created_at"] = time.time() - 99999
|
|
removed = self.router.cleanup_expired()
|
|
assert removed >= 1
|
|
|
|
def test_binding_count(self):
|
|
assert self.router.binding_count == 0
|
|
self.router.bind_agent_account("agent1", "default")
|
|
assert self.router.binding_count == 1
|
|
|
|
def test_get_binding(self):
|
|
self.router.bind_agent_account("agent1", "default", {"extra": "info"})
|
|
binding = self.router.get_binding("agent1", "default")
|
|
assert binding["agent_id"] == "agent1"
|
|
assert binding["data"] == {"extra": "info"}
|
|
|
|
def test_get_binding_expired(self):
|
|
self.router.bind_agent_account("agent1", "default")
|
|
for bind_key, entry in list(self.router._bindings.items()):
|
|
entry["created_at"] = time.time() - 99999
|
|
assert self.router.get_binding("agent1", "default") is None
|
|
|
|
def test_resolve_session_params_default(self):
|
|
params = self.router.resolve_session_params("user1")
|
|
assert params["chat_type"] == ChatType.DIRECT
|
|
assert params["channel_chat_id"] == "user1"
|
|
|
|
def test_resolve_session_params_full(self):
|
|
params = self.router.resolve_session_params("user1", sender_name="John", avatar="url", dm_policy="open", account_id="prod")
|
|
assert params["sender_name"] == "John"
|
|
assert params["avatar"] == "url"
|
|
assert params["dm_policy"] == "open"
|
|
assert params["account_id"] == "prod"
|
|
|
|
|
|
class TestGlobalSessionFunctions:
|
|
def test_resolve_session_params_global(self):
|
|
params = resolve_session_params("user1", sender_name="Test")
|
|
assert params["chat_type"] == ChatType.DIRECT
|
|
|
|
def test_get_session_router(self):
|
|
router = get_session_router()
|
|
assert isinstance(router, SessionRouter)
|
|
|
|
|
|
# ==============================================================================
|
|
# Formatter (Extended)
|
|
# ==============================================================================
|
|
|
|
|
|
class TestFormatterExtended:
|
|
def setup_method(self):
|
|
self.formatter = ZaloOAMessageFormatter()
|
|
|
|
def _make_response(self, content, message_type=MessageType.TEXT, metadata=None):
|
|
identity = ChannelIdentity(
|
|
channel_id="zalo_oa", channel_type=ChannelType.ZALO_OA,
|
|
channel_user_id="follower_001", channel_chat_id="oa_001",
|
|
)
|
|
return ChannelResponse(identity=identity, content=content, message_type=message_type, metadata=metadata or {})
|
|
|
|
def test_build_button_url(self):
|
|
button = self.formatter.build_button_url("Open", "https://example.com")
|
|
assert button["title"] == "Open"
|
|
assert button["type"] == "oa.open.url"
|
|
assert button["payload"]["url"] == "https://example.com"
|
|
|
|
def test_build_button_query(self):
|
|
button = self.formatter.build_button_query("Query", "search_term")
|
|
assert button["title"] == "Query"
|
|
assert button["type"] == "oa.query.show"
|
|
assert button["payload"] == "search_term"
|
|
|
|
def test_build_button_hide(self):
|
|
button = self.formatter.build_button_hide()
|
|
assert button["type"] == "oa.query.hide"
|
|
|
|
def test_build_button_hide_custom_title(self):
|
|
button = self.formatter.build_button_hide("Đóng")
|
|
assert button["title"] == "Đóng"
|
|
|
|
def test_build_inline_keyboard(self):
|
|
buttons = [{"title": "Click", "type": "oa.open.url", "payload": {"url": "https://ex.com"}}]
|
|
result = self.formatter.build_inline_keyboard("user1", "Choose:", buttons)
|
|
assert result["recipient"]["user_id"] == "user1"
|
|
assert result["message"]["text"] == "Choose:"
|
|
assert result["message"]["attachment"]["payload"]["template_type"] == "button"
|
|
|
|
def test_build_image_template_basic(self):
|
|
result = self.formatter.build_image_template("user1", "att_001")
|
|
assert result["recipient"]["user_id"] == "user1"
|
|
assert result["message"]["attachment"]["payload"]["elements"][0]["attachment_id"] == "att_001"
|
|
|
|
def test_build_image_template_with_title_subtitle(self):
|
|
result = self.formatter.build_image_template("user1", "att_001", title="My Image", subtitle="Beautiful")
|
|
elem = result["message"]["attachment"]["payload"]["elements"][0]
|
|
assert elem["title"] == "My Image"
|
|
assert elem["subtitle"] == "Beautiful"
|
|
|
|
def test_build_image_template_with_buttons(self):
|
|
btn = [{"title": "View", "type": "oa.open.url", "payload": {"url": "https://ex.com"}}]
|
|
result = self.formatter.build_image_template("user1", "att_001", buttons=btn)
|
|
assert "buttons" in result["message"]["attachment"]["payload"]
|
|
|
|
def test_build_text_template_basic(self):
|
|
result = self.formatter.build_text_template("user1", "Hello")
|
|
elem = result["message"]["attachment"]["payload"]["elements"][0]
|
|
assert elem["title"] == "Hello"
|
|
|
|
def test_build_text_template_full(self):
|
|
btn = [self.formatter.build_button_url("Open", "https://ex.com")]
|
|
result = self.formatter.build_text_template("user1", "Title", subtitle="Sub", description="Desc", buttons=btn)
|
|
elem = result["message"]["attachment"]["payload"]["elements"][0]
|
|
assert elem["title"] == "Title"
|
|
assert elem["subtitle"] == "Sub"
|
|
assert elem["description"] == "Desc"
|
|
|
|
def test_format_sticker(self):
|
|
response = self._make_response("", message_type=MessageType.STICKER, metadata={"sticker_id": "stk_001"})
|
|
result = self.formatter.format(response)
|
|
elem = result["message"]["attachment"]["payload"]["elements"][0]
|
|
assert elem["media_type"] == "sticker"
|
|
assert elem["attachment_id"] == "stk_001"
|
|
|
|
def test_format_location_via_format(self):
|
|
response = self._make_response("", message_type=MessageType.LOCATION, metadata={"lat": "10.0", "lon": "106.0"})
|
|
result = self.formatter.format(response)
|
|
elem = result["message"]["attachment"]["payload"]["elements"][0]
|
|
assert elem["media_type"] == "location"
|
|
assert elem["latitude"] == 10.0
|
|
assert elem["longitude"] == 106.0
|
|
|
|
|
|
# ==============================================================================
|
|
# Outbound Media Host
|
|
# ==============================================================================
|
|
|
|
|
|
class TestOutboundMediaHost:
|
|
def setup_method(self):
|
|
self.host = OutboundMediaHost(ttl_sec=300, max_entries=10)
|
|
|
|
def test_store_and_get(self):
|
|
ref_id = self.host.store("att_001", "image")
|
|
entry = self.host.get(ref_id)
|
|
assert entry["attachment_id"] == "att_001"
|
|
assert entry["media_type"] == "image"
|
|
|
|
def test_get_nonexistent(self):
|
|
assert self.host.get("nonexistent") is None
|
|
|
|
def test_ttl_expiry(self):
|
|
host = OutboundMediaHost(ttl_sec=0)
|
|
ref_id = host.store("att_001", "image")
|
|
assert host.get(ref_id) is None
|
|
|
|
def test_resolve_attachment_id(self):
|
|
ref_id = self.host.store("att_001", "image")
|
|
assert self.host.resolve_attachment_id(ref_id) == "att_001"
|
|
|
|
def test_resolve_attachment_id_nonexistent(self):
|
|
assert self.host.resolve_attachment_id("nonexistent") is None
|
|
|
|
def test_remove(self):
|
|
ref_id = self.host.store("att_001", "image")
|
|
self.host.remove(ref_id)
|
|
assert self.host.get(ref_id) is None
|
|
|
|
def test_active_count(self):
|
|
assert self.host.active_count == 0
|
|
self.host.store("att_001", "image")
|
|
assert self.host.active_count == 1
|
|
|
|
def test_max_entries_eviction(self):
|
|
host = OutboundMediaHost(max_entries=3)
|
|
for i in range(5):
|
|
host.store(f"att_{i}", "image")
|
|
assert host.active_count <= 3
|
|
|
|
def test_cleanup_expired(self):
|
|
host = OutboundMediaHost(ttl_sec=0)
|
|
for i in range(3):
|
|
host.store(f"att_{i}", "image")
|
|
removed = host.cleanup_expired()
|
|
assert removed == 3
|
|
|
|
def test_store_with_metadata(self):
|
|
ref_id = self.host.store("att_001", "image", {"filename": "test.jpg"})
|
|
entry = self.host.get(ref_id)
|
|
assert entry["metadata"]["filename"] == "test.jpg"
|
|
|
|
|
|
class TestGlobalOutboundMediaFunctions:
|
|
def test_store_get_cycle(self):
|
|
ref_id = store_media("att_global", "image")
|
|
entry = get_media(ref_id)
|
|
assert entry["attachment_id"] == "att_global"
|
|
|
|
def test_resolve_attachment_func(self):
|
|
ref_id = store_media("att_global", "image")
|
|
assert resolve_attachment(ref_id) == "att_global"
|
|
|
|
def test_cleanup_media_cache(self):
|
|
store_media("att_temp", "image")
|
|
removed = cleanup_media_cache()
|
|
assert isinstance(removed, int)
|
|
|
|
|
|
# ==============================================================================
|
|
# Config Schema
|
|
# ==============================================================================
|
|
|
|
|
|
class TestConfigSchema:
|
|
def test_zalo_oa_config_defaults(self):
|
|
config = ZaloOAConfig()
|
|
assert config.enabled is True
|
|
assert config.dm_policy == "open"
|
|
assert config.polling_enabled is True
|
|
assert config.voice_tts_enabled is False
|
|
|
|
def test_zalo_oa_account_config_defaults(self):
|
|
config = ZaloOAAccountConfig()
|
|
assert config.app_id == ""
|
|
assert config.secret_key == ""
|
|
assert config.enabled is True
|
|
assert config.dm_policy == "open"
|
|
assert config.polling_enabled is True
|
|
assert config.voice_tts_enabled is False
|
|
assert config.heartbeat_interval_sec == 30
|
|
|
|
def test_to_flat_config_no_accounts(self):
|
|
config = ZaloOAConfig(dm_policy="allowlist", enabled=False)
|
|
flat = config.to_flat_config()
|
|
assert flat["dm_policy"] == "allowlist"
|
|
assert flat["enabled"] is False
|
|
|
|
def test_to_flat_config_with_accounts(self):
|
|
config = ZaloOAConfig(
|
|
dm_policy="allowlist",
|
|
accounts={"prod": ZaloOAAccountConfig(app_id="app_id_1", secret_key="secret_1", dm_policy="allowlist", name="Production")},
|
|
)
|
|
flat = config.to_flat_config()
|
|
assert flat["app_id"] == "app_id_1"
|
|
assert flat["dm_policy"] == "allowlist"
|
|
|
|
def test_config_field_limits(self):
|
|
config = ZaloOAConfig(
|
|
rate_limit_window_ms=60000,
|
|
rate_limit_max_requests=60,
|
|
health_check_ttl_sec=120,
|
|
dedup_window_ms=500000,
|
|
)
|
|
assert config.rate_limit_window_ms == 60000
|
|
assert config.rate_limit_max_requests == 60
|
|
assert config.health_check_ttl_sec == 120
|
|
assert config.dedup_window_ms == 500000
|
|
|
|
|
|
# ==============================================================================
|
|
# Polling
|
|
# ==============================================================================
|
|
|
|
|
|
class TestPollingExtended:
|
|
def test_poller_initial_metrics(self):
|
|
poller = ZaloOAPoller(client=None, config={"polling_interval_sec": 10, "polling_timeout_ms": 5000})
|
|
metrics = poller.get_poll_metrics()
|
|
assert metrics["running"] is False
|
|
assert metrics["poll_count"] == 0
|
|
assert metrics["error_count"] == 0
|
|
assert metrics["interval_sec"] == 10
|
|
assert metrics["timeout_ms"] == 5000
|
|
|
|
def test_poller_with_callback(self):
|
|
cb = MagicMock()
|
|
poller = ZaloOAPoller(client=None, message_callback=cb)
|
|
assert poller._message_callback == cb
|
|
|
|
def test_poller_last_poll_at_initial(self):
|
|
poller = ZaloOAPoller(client=None)
|
|
metrics = poller.get_poll_metrics()
|
|
assert metrics["last_poll_at"] is None
|
|
|
|
def test_poller_follower_delta_initial(self):
|
|
poller = ZaloOAPoller(client=None)
|
|
metrics = poller.get_poll_metrics()
|
|
assert metrics["follower_delta_total"] == 0
|
|
|
|
|
|
# ==============================================================================
|
|
# Voice
|
|
# ==============================================================================
|
|
|
|
|
|
class TestVoiceExtended:
|
|
def test_voice_defaults(self):
|
|
voice = ZaloOAVoice()
|
|
assert voice.enabled is False
|
|
assert voice.synthesis_target == "voice-note"
|
|
|
|
def test_voice_enabled_with_api_url(self):
|
|
voice = ZaloOAVoice(config={"voice_tts_enabled": True, "voice_tts_api_url": "https://api.tts.com"})
|
|
assert voice.enabled is True
|
|
|
|
def test_voice_disabled_without_api_url_but_enabled(self):
|
|
voice = ZaloOAVoice(config={"voice_tts_enabled": True})
|
|
assert voice.enabled is False
|
|
|
|
def test_voice_no_config(self):
|
|
voice = ZaloOAVoice()
|
|
assert voice.enabled is False
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_synthesize_and_send_not_enabled(self):
|
|
voice = ZaloOAVoice()
|
|
result = await voice.synthesize_and_send("hello", "user1", None)
|
|
assert result.success is False
|
|
assert "not configured" in result.error
|
|
|
|
def test_voice_custom_timeout(self):
|
|
voice = ZaloOAVoice(config={"voice_tts_enabled": True, "voice_tts_api_url": "https://api.com", "voice_tts_timeout_sec": 60})
|
|
assert voice.enabled is True
|
|
|
|
|
|
# ==============================================================================
|
|
# Normalizer - Additional Edge Cases
|
|
# ==============================================================================
|
|
|
|
|
|
class TestNormalizerEdgeCases:
|
|
def setup_method(self):
|
|
self.normalizer = ZaloOAEventNormalizer()
|
|
|
|
def test_user_submit_form(self):
|
|
payload = {
|
|
"event_name": "user_submit_form",
|
|
"sender": {"id": "u1"},
|
|
"recipient": {"id": "oa1"},
|
|
"message": {"text": "Form submitted with data"},
|
|
"timestamp": "1",
|
|
}
|
|
result = self.normalizer.normalize(payload)
|
|
assert result.event_type == EventType.CARD_ACTION
|
|
assert "Form submitted" in result.content
|
|
|
|
def test_user_click_button(self):
|
|
payload = {
|
|
"event_name": "user_click_button",
|
|
"sender": {"id": "u1"},
|
|
"recipient": {"id": "oa1"},
|
|
"message": {"payload": {"label": "Click me"}},
|
|
"timestamp": "1",
|
|
}
|
|
result = self.normalizer.normalize(payload)
|
|
assert result.event_type == EventType.CARD_ACTION
|
|
assert result.content == "Click me"
|
|
|
|
def test_user_send_gif(self):
|
|
payload = {
|
|
"event_name": "user_send_gif",
|
|
"sender": {"id": "u1"},
|
|
"recipient": {"id": "oa1"},
|
|
"message": {"msg_id": "gif1", "attachments": [{"type": "image", "payload": {"url": "https://ex.com/gif.gif", "id": "g1"}}]},
|
|
"timestamp": "1",
|
|
}
|
|
result = self.normalizer.normalize(payload)
|
|
assert result.message_type == MessageType.IMAGE
|
|
|
|
def test_text_with_mentions(self):
|
|
payload = {
|
|
"event_name": "user_send_text",
|
|
"sender": {"id": "u1"},
|
|
"recipient": {"id": "oa1"},
|
|
"message": {"msg_id": "m1", "text": "@user1 @user2 hello there"},
|
|
"timestamp": "1",
|
|
}
|
|
result = self.normalizer.normalize(payload)
|
|
assert "[mentions:" in result.content
|
|
assert "user1" in result.content
|
|
assert "user2" in result.content
|
|
|
|
def test_unknown_event_name_defaults_to_text(self):
|
|
payload = {
|
|
"event_name": "user_send_unknown",
|
|
"sender": {"id": "u1"},
|
|
"recipient": {"id": "oa1"},
|
|
"message": {"msg_id": "m1", "text": "unknown type"},
|
|
"timestamp": "1",
|
|
}
|
|
result = self.normalizer.normalize(payload)
|
|
assert result.message_type == MessageType.TEXT
|
|
assert result.event_type == EventType.MESSAGE_RECEIVED
|
|
|
|
def test_missing_fields_handled_gracefully(self):
|
|
payload = {"event_name": "follow"}
|
|
result = self.normalizer.normalize(payload)
|
|
assert result.content == "Follower followed OA"
|
|
assert result.identity.channel_user_id == "unknown"
|
|
assert result.identity.channel_chat_id == ""
|
|
|
|
|
|
class TestSkipMessageError:
|
|
def test_can_be_raised_and_caught(self):
|
|
try:
|
|
raise SkipMessageError("test skip")
|
|
except SkipMessageError as e:
|
|
assert str(e) == "test skip" |