from __future__ import annotations import time from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channels.adapters.wechat.pairing import ( PAIRING_CODE_CHARS, PAIRING_CODE_LENGTH, PAIRING_MAX_PENDING, PAIRING_RATE_LIMIT_S, PAIRING_VALIDITY_S, PairingRequest, WeChatPairingAdapter, ) class TestPairingCodeGeneration: @pytest.fixture def adapter(self): return WeChatPairingAdapter() def test_generate_code_length(self, adapter): code = adapter.generate_code() assert len(code) == PAIRING_CODE_LENGTH def test_generate_code_charset(self, adapter): code = adapter.generate_code() for ch in code: assert ch in PAIRING_CODE_CHARS def test_generate_code_uniqueness(self, adapter): codes = {adapter.generate_code() for _ in range(100)} assert len(codes) > 90 class TestPairingFlow: @pytest.fixture def adapter(self): return WeChatPairingAdapter() def test_request_pairing_returns_code(self, adapter): code = adapter.request_pairing("sender1", "wxid_abc") assert code is not None assert len(code) == PAIRING_CODE_LENGTH assert adapter.is_approved("wxid_abc") is False def test_request_pairing_already_approved_returns_none(self, adapter): adapter._approved.add("wxid_abc") code = adapter.request_pairing("sender1", "wxid_abc") assert code is None def test_approve_pairing_success(self, adapter): code = adapter.request_pairing("sender1", "wxid_abc") assert code is not None result = adapter.approve(code) assert result is True assert adapter.is_approved("wxid_abc") is True def test_approve_invalid_code(self, adapter): result = adapter.approve("INVALID1") assert result is False def test_reject_pairing_success(self, adapter): code = adapter.request_pairing("sender1", "wxid_abc") assert code is not None result = adapter.reject(code) assert result is True assert adapter.is_approved("wxid_abc") is False def test_list_pending_requests(self, adapter): code1 = adapter.request_pairing("sender1", "wxid_abc") code2 = adapter.request_pairing("sender2", "wxid_def") pending = adapter.list_pending() assert len(pending) == 2 codes = [p["code"] for p in pending] assert code1 in codes assert code2 in codes def test_list_pending_includes_fields(self, adapter): adapter.request_pairing("sender1", "wxid_abc") pending = adapter.list_pending() assert len(pending) == 1 assert "code" in pending[0] assert "sender_id" in pending[0] assert "channel_user_id" in pending[0] assert "requested_at" in pending[0] assert "expires_at" in pending[0] class TestPairingRateLimit: @pytest.fixture def adapter(self): adapter = WeChatPairingAdapter() adapter._rate_limit["wxid_abc"] = time.monotonic() return adapter def test_rate_limit_blocks_immediate_retry(self, adapter): code = adapter.request_pairing("sender1", "wxid_abc") assert code is None def test_rate_limit_allows_after_window(self, adapter): adapter._rate_limit["wxid_abc"] = time.monotonic() - PAIRING_RATE_LIMIT_S - 1 code = adapter.request_pairing("sender1", "wxid_abc") assert code is not None class TestPairingMaxPending: @pytest.fixture def adapter(self): adapter = WeChatPairingAdapter() for i in range(PAIRING_MAX_PENDING): adapter.request_pairing(f"sender{i}", f"wxid_{i}") return adapter def test_max_pending_limit(self, adapter): code = adapter.request_pairing("sender_extra", "wxid_extra") assert code is None def test_max_pending_after_expiry(self, adapter): for code, req in list(adapter._pending.items()): req.expires_at = time.time() - 1 code = adapter.request_pairing("sender_new", "wxid_new") assert code is not None class TestPairingPersistence: @pytest.fixture def adapter(self, tmp_path): return WeChatPairingAdapter(storage_dir=str(tmp_path)) @pytest.mark.asyncio async def test_save_and_load_allowlist(self, adapter): adapter._approved.add("wxid_saved") await adapter.save_allowlist("test_account") adapter2 = WeChatPairingAdapter(storage_dir=adapter._storage_dir) loaded = await adapter2.load_allowlist("test_account") assert "wxid_saved" in loaded assert adapter2.is_approved("wxid_saved") @pytest.mark.asyncio async def test_save_and_load_pending(self, adapter): adapter.request_pairing("sender1", "wxid_pending") await adapter.save_pending("test_account") adapter2 = WeChatPairingAdapter(storage_dir=adapter._storage_dir) await adapter2.load_pending("test_account") pending = adapter2.list_pending() assert len(pending) == 1 assert pending[0]["sender_id"] == "sender1" @pytest.mark.asyncio async def test_load_nonexistent_allowlist(self, adapter): loaded = await adapter.load_allowlist("nonexistent") assert loaded == [] class TestPairingNormalizeAllowEntry: def test_normalize_with_prefix(self): result = WeChatPairingAdapter.normalize_allow_entry("wx:user123") assert result == "wx:user123" def test_normalize_without_prefix(self): result = WeChatPairingAdapter.normalize_allow_entry("user123") assert result == "wx:user123" def test_normalize_whitespace(self): result = WeChatPairingAdapter.normalize_allow_entry(" user123 ") assert result == "wx:user123" class TestPairingCleanup: @pytest.fixture def adapter(self): return WeChatPairingAdapter() def test_cleanup_expired_requests(self, adapter): code = adapter.request_pairing("sender1", "wxid_abc") req = adapter._pending[code] req.expires_at = time.time() - 1 adapter._cleanup_expired() assert code not in adapter._pending assert adapter.list_pending() == [] class TestPairingNotify: @pytest.mark.asyncio async def test_notify_approval_calls_send(self): adapter = WeChatPairingAdapter() mock_send = AsyncMock() await adapter.notify_approval(mock_send, "wxid_abc") mock_send.assert_called_once() call_args = mock_send.call_args[0] assert call_args[0] == "wxid_abc" assert "批准" in call_args[1] @pytest.mark.asyncio async def test_notify_approval_handles_error(self): adapter = WeChatPairingAdapter() mock_send = AsyncMock(side_effect=RuntimeError("send failed")) await adapter.notify_approval(mock_send, "wxid_abc") mock_send.assert_called_once() class TestPairingRequestDataclass: def test_pairing_request_creation(self): now = time.time() req = PairingRequest( code="ABCD1234", sender_id="sender1", channel_user_id="wxid_abc", requested_at=now, expires_at=now + 3600, ) assert req.code == "ABCD1234" assert req.sender_id == "sender1" assert req.channel_user_id == "wxid_abc" assert req.requested_at == now assert req.expires_at == now + 3600