from __future__ import annotations from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channels.adapters.wechat.adapter import WeChatAdapter from yuxi.channels.adapters.wechat.approval import ( WeChatApprovalAdapter, ApprovalRequest, ) from yuxi.channels.models import ( ChannelIdentity, ChannelResponse, ChannelType, DeliveryResult, ) class TestApprovalAdapter: @pytest.fixture def adapter(self): adapter = WeChatApprovalAdapter() adapter.configure({"approval": {"enabled": True, "timeout_minutes": 5}}) return adapter def test_configure_enabled(self, adapter): assert adapter._enabled is True assert adapter._timeout_minutes == 5 def test_configure_disabled(self): adapter = WeChatApprovalAdapter() adapter.configure({"approval": {"enabled": False}}) assert adapter._enabled is False def test_create_request(self, adapter): req = adapter.create_request("req_1", "exec", "user123", "chat_1", {"cmd": "ls"}) assert req is not None assert req.request_id == "req_1" assert req.action == "exec" assert req.channel_user_id == "user123" assert req.status == "pending" assert req.chat_type == "direct" assert req.expires_at > req.created_at def test_create_request_disabled(self): adapter = WeChatApprovalAdapter() adapter._enabled = False req = adapter.create_request("req_1", "exec", "user123", "chat_1") assert req is None def test_create_request_dm_disabled(self, adapter): adapter._dm_enabled = False req = adapter.create_request("req_1", "exec", "user123", "chat_1", chat_type="direct") assert req is None def test_create_request_group_enabled(self, adapter): req = adapter.create_request("req_1", "exec", "user123", "chat_1", chat_type="group") assert req is not None assert req.chat_type == "group" def test_approve_request(self, adapter): adapter.create_request("req_1", "exec", "user123", "chat_1") result = adapter.approve("req_1", "admin_user") assert result is True assert adapter.is_approved("req_1") is False def test_deny_request(self, adapter): adapter.create_request("req_2", "exec", "user456", "chat_2") result = adapter.deny("req_2", "unauthorized command", "admin_user") assert result is True assert adapter.is_pending("req_2") is False def test_approve_nonexistent(self, adapter): result = adapter.approve("nonexistent") assert result is False def test_approve_already_handled(self, adapter): adapter.create_request("req_3", "exec", "user", "chat") adapter.deny("req_3", "denied") result = adapter.approve("req_3") assert result is False def test_list_pending_by_chat_type(self, adapter): adapter.create_request("req_dm", "exec", "u1", "c1", chat_type="direct") adapter.create_request("req_group", "exec", "u2", "c2", chat_type="group") dm_pending = adapter.list_pending(chat_type="direct") assert len(dm_pending) == 1 assert dm_pending[0]["request_id"] == "req_dm" group_pending = adapter.list_pending(chat_type="group") assert len(group_pending) == 1 assert group_pending[0]["request_id"] == "req_group" def test_resolve_timeout(self, adapter): import time adapter.create_request("req_timeout", "exec", "user", "chat") req = adapter._pending.get("req_timeout") if req: req.expires_at = time.time() - 1 expired = adapter.resolve_timeout_requests() assert "req_timeout" in expired assert adapter.is_pending("req_timeout") is False def test_get_history(self, adapter): adapter.create_request("req_h", "exec", "user", "chat") adapter.approve("req_h") history = adapter.get_history() assert len(history) >= 1 assert history[0]["outcome"] == "approved" assert history[0]["request_id"] == "req_h" def test_callback_on_approve(self, adapter): callbacks = [] def on_approve(req_id, outcome, payload): callbacks.append((req_id, outcome)) adapter.on("approved", on_approve) adapter.create_request("req_cb", "exec", "user", "chat") adapter.approve("req_cb") assert len(callbacks) == 1 assert callbacks[0] == ("req_cb", "approved") def test_callback_on_deny(self, adapter): callbacks = [] def on_deny(req_id, outcome, payload): callbacks.append((req_id, outcome, payload.get("reason"))) adapter.on("denied", on_deny) adapter.create_request("req_deny", "exec", "user", "chat") adapter.deny("req_deny", "blocked command") assert len(callbacks) == 1 assert callbacks[0][2] == "blocked command" def test_callback_on_timeout(self, adapter): import time callbacks = [] def on_timeout(req_id, outcome, payload): callbacks.append(req_id) adapter.on("timeout", on_timeout) adapter.create_request("req_to", "exec", "user", "chat") req = adapter._pending.get("req_to") if req: req.expires_at = time.time() - 1 adapter.resolve_timeout_requests() assert len(callbacks) == 1 assert callbacks[0] == "req_to" def test_get_pending_count(self, adapter): adapter.create_request("r1", "exec", "u1", "c1", chat_type="direct") adapter.create_request("r2", "exec", "u2", "c2", chat_type="group") adapter.create_request("r3", "exec", "u3", "c3", chat_type="direct") assert adapter.get_pending_count() == 3 assert adapter.get_pending_count("direct") == 2 assert adapter.get_pending_count("group") == 1 def test_max_pending_limit(self, adapter): for i in range(60): adapter.create_request(f"req_{i}", "exec", f"user_{i}", f"chat_{i}") pending = adapter.list_pending() assert len(pending) <= 50 def test_history_max_size(self, adapter): for i in range(250): adapter.create_request(f"req_h_{i}", "exec", "user", "chat") adapter.approve(f"req_h_{i}") history = adapter.get_history(limit=500) assert len(history) <= 200 class TestSecurityIntegration: @pytest.fixture def adapter(self): config = { "dm_policy": "open", "group_policy": "open", "allow_from": ["wx:trusted_user"], } return WeChatAdapter(config) def test_dm_policy_allow_all(self, adapter): assert adapter._check_dm_policy("any_user") is True def test_dm_policy_deny_all(self): adapter = WeChatAdapter({"dm_policy": "disabled"}) assert adapter._check_dm_policy("any_user") is False def test_dm_policy_allowlist_match(self): adapter = WeChatAdapter({"dm_policy": "allowlist", "allow_from": ["wx:user_a"]}) assert adapter._check_dm_policy("user_a") is True def test_dm_policy_allowlist_mismatch(self): adapter = WeChatAdapter({"dm_policy": "allowlist", "allow_from": ["wx:user_a"]}) assert adapter._check_dm_policy("user_b") is False def test_group_policy_per_group_allow(self): adapter = WeChatAdapter({ "group_policy": "allowlist", "groups": {"room_1": {"allow_from": ["wx:member_a"]}}, }) assert adapter._check_group_policy("room_1", "member_a") is True def test_group_policy_disabled(self): adapter = WeChatAdapter({"group_policy": "disabled"}) assert adapter._check_group_policy("any_room", "any_user") is False def test_at_bot_wecom_mention(self): adapter = WeChatAdapter({"dm_policy": "open"}) adapter._mode = "wecom" adapter.config["agent_name"] = "TestBot" payload = {"Content": "@TestBot what's up"} assert adapter._is_at_bot(payload) is True def test_at_bot_wecom_no_mention(self): adapter = WeChatAdapter({"dm_policy": "open"}) adapter._mode = "wecom" payload = {"Content": "just a message"} assert adapter._is_at_bot(payload) is False def test_at_bot_bridge_no_at(self): adapter = WeChatAdapter({"bridge_url": "http://b", "dm_policy": "open"}) adapter._mode = "personal" payload = {"at_list": []} assert adapter._is_at_bot(payload) is False @pytest.mark.asyncio async def test_banned_send_wecom(self): adapter = WeChatAdapter({"corp_id": "test", "corp_secret": "test", "agent_id": "1", "dm_policy": "open"}) adapter._mode = "wecom" adapter._banned = True adapter._ban_permanent = True adapter._banned_reason = "48001" response = ChannelResponse( identity=ChannelIdentity( channel_id="wechat", channel_type=ChannelType.WECHAT, channel_user_id="u1", channel_chat_id="u1", ), content="test", ) result = await adapter.send(response) assert result.success is False assert "permanently banned" in result.error.lower() @pytest.mark.asyncio async def test_banned_health_check_degraded(self): adapter = WeChatAdapter({"corp_id": "test", "corp_secret": "test", "agent_id": "1", "dm_policy": "open"}) adapter._mode = "wecom" adapter._banned = True adapter._ban_permanent = False adapter._banned_reason = "48001" adapter._http_client = AsyncMock() result = await adapter.health_check() assert result.status == "degraded" @pytest.mark.asyncio async def test_banned_health_check_permanent(self): adapter = WeChatAdapter({"corp_id": "test", "corp_secret": "test", "agent_id": "1", "dm_policy": "open"}) adapter._mode = "wecom" adapter._ban_permanent = True adapter._banned_reason = "48001" adapter._http_client = AsyncMock() result = await adapter.health_check() assert result.status == "unhealthy" assert "permanently" in (result.last_error or "").lower() def test_validate_config_new_options_valid(self): from yuxi.channels.adapters.wechat.config_schema import validate_wechat_config config = { "corp_id": "test", "corp_secret": "test", "agent_id": "1", "retry_attempts": 5, "retry_min_delay": 1.0, "retry_max_delay": 60.0, "outbound_priority_enabled": True, "exec_approval_timeout": 10, "ban_backoff_intervals": [60, 300, 900], } errors = validate_wechat_config(config) assert errors == [] def test_validate_config_invalid_retry(self): from yuxi.channels.adapters.wechat.config_schema import validate_wechat_config config = { "corp_id": "test", "corp_secret": "test", "agent_id": "1", "retry_attempts": -1, } errors = validate_wechat_config(config) assert len(errors) >= 1 def test_validate_config_invalid_ban_backoff(self): from yuxi.channels.adapters.wechat.config_schema import validate_wechat_config config = { "corp_id": "test", "corp_secret": "test", "agent_id": "1", "ban_backoff_intervals": [0, -1], } errors = validate_wechat_config(config) assert len(errors) >= 1