from __future__ import annotations from datetime import datetime, timedelta from unittest.mock import AsyncMock, MagicMock, patch import pytest from yuxi.channel.security.pairing import PairingManager from yuxi.storage.postgres.model_channel import ChannelPairingRecord from yuxi.utils.datetime_utils import utc_now_naive @pytest.fixture def pairing_manager() -> PairingManager: return PairingManager() @pytest.fixture def mock_session_ctx(monkeypatch): session = AsyncMock() ctx = MagicMock() ctx.__aenter__ = AsyncMock(return_value=session) ctx.__aexit__ = AsyncMock(return_value=None) manager = MagicMock() manager.get_async_session_context = MagicMock(return_value=ctx) monkeypatch.setattr("yuxi.channel.security.pairing.pg_manager", manager) return session @pytest.mark.unit class TestPairingManagerGenerate: def test_generate_code_string_returns_six_digit_code(self, pairing_manager: PairingManager) -> None: code = pairing_manager.generate_code_string() assert len(code) == PairingManager.CODE_LENGTH assert code.isdigit() async def test_generate_returns_six_digit_code( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: code = await pairing_manager.generate_pairing_code("feishu", "acc_1", "peer_1") assert len(code) == PairingManager.CODE_LENGTH assert code.isdigit() async def test_generate_creates_record_with_expected_fields( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock, monkeypatch ) -> None: with patch("yuxi.channel.security.pairing.secrets") as mock_secrets: mock_secrets.choice = MagicMock(side_effect=["1", "2", "3", "4", "5", "6"]) mock_secrets.token_urlsafe = MagicMock(return_value="token_xxx") code = await pairing_manager.generate_pairing_code("feishu", "acc_1", "peer_1") assert code == "123456" mock_session_ctx.execute.assert_awaited_once() added_record = mock_session_ctx.add.call_args[0][0] assert isinstance(added_record, ChannelPairingRecord) assert added_record.channel_type == "feishu" assert added_record.account_id == "acc_1" assert added_record.peer_id == "peer_1" assert added_record.pairing_code == "123456" assert added_record.pairing_token == "token_xxx" assert added_record.status == "pending" assert isinstance(added_record.expires_at, datetime) mock_session_ctx.commit.assert_awaited_once() async def test_generate_with_explicit_code_uses_it( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: code = await pairing_manager.generate_pairing_code("feishu", "acc_1", "peer_1", code="987654") assert code == "987654" added_record = mock_session_ctx.add.call_args[0][0] assert added_record.pairing_code == "987654" async def test_generate_with_qr_content_and_pairing_mode( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: await pairing_manager.generate_pairing_code( "feishu", "acc_1", "peer_1", platform_user_id="user_1", qr_content="https://example.com/bind?code=123456", pairing_mode="qr", ) added_record = mock_session_ctx.add.call_args[0][0] assert added_record.platform_user_id == "user_1" assert added_record.qr_content == "https://example.com/bind?code=123456" assert added_record.pairing_mode == "qr" async def test_generate_removes_existing_pending_records( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: await pairing_manager.generate_pairing_code("feishu", "acc_1", "peer_1") executed = mock_session_ctx.execute.call_args[0][0] # The first call must be a delete against pending records for the same channel/peer. assert "DELETE" in str(executed).upper() params = executed.compile().params assert params.get("status_1") == "pending" @pytest.mark.unit class TestPairingManagerVerify: async def test_verify_success_updates_record_and_returns_true( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: record = MagicMock() record.pairing_code = "123456" result = MagicMock() result.scalar_one_or_none = MagicMock(return_value=record) mock_session_ctx.execute = AsyncMock(return_value=result) assert await pairing_manager.verify_pairing_code("feishu", "acc_1", "peer_1", "123456") is True assert record.status == "paired" assert isinstance(record.paired_at, datetime) mock_session_ctx.commit.assert_awaited_once() async def test_verify_wrong_code_returns_false( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: record = MagicMock() record.pairing_code = "654321" result = MagicMock() result.scalar_one_or_none = MagicMock(return_value=record) mock_session_ctx.execute = AsyncMock(return_value=result) assert await pairing_manager.verify_pairing_code("feishu", "acc_1", "peer_1", "123456") is False mock_session_ctx.commit.assert_not_awaited() async def test_verify_no_record_returns_false( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: result = MagicMock() result.scalar_one_or_none = MagicMock(return_value=None) mock_session_ctx.execute = AsyncMock(return_value=result) assert await pairing_manager.verify_pairing_code("feishu", "acc_1", "peer_1", "123456") is False async def test_verify_expired_record_returns_false( self, pairing_manager: PairingManager, mock_session_ctx: AsyncMock ) -> None: record = MagicMock() record.pairing_code = "123456" record.expires_at = utc_now_naive() - timedelta(minutes=PairingManager.CODE_TTL_MINUTES + 1) result = MagicMock() result.scalar_one_or_none = MagicMock(return_value=None) mock_session_ctx.execute = AsyncMock(return_value=result) assert await pairing_manager.verify_pairing_code("feishu", "acc_1", "peer_1", "123456") is False @pytest.mark.unit class TestPairingManagerConstants: def test_code_length_and_ttl(self) -> None: assert PairingManager.CODE_LENGTH == 6 assert PairingManager.CODE_TTL_MINUTES == 10