155 lines
6.4 KiB
Python
155 lines
6.4 KiB
Python
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
|