新增了Twitch、Telegram、Discord、Slack、Mattermost、WeChat、Zalo等多渠道的单元测试用例,覆盖了令牌处理、速率限制、消息去重、会话解析、格式转换、安全策略等模块 同时在测试配置中添加了测试用的OpenAI API密钥环境变量
136 lines
4.3 KiB
Python
136 lines
4.3 KiB
Python
from __future__ import annotations
|
|
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
from sqlalchemy.exc import IntegrityError
|
|
|
|
from yuxi.channels.models import ChannelIdentity, ChannelMessage, ChannelType
|
|
from yuxi.channels.session_mapper import SessionMapper, VIRTUAL_DEPARTMENT_ID, USER_SOURCE_PREFIX
|
|
|
|
|
|
def _make_message(channel_user_id: str = "u1", channel_chat_id: str = "c1") -> ChannelMessage:
|
|
identity = ChannelIdentity(
|
|
channel_id="test",
|
|
channel_type=ChannelType.WEBCHAT,
|
|
channel_user_id=channel_user_id,
|
|
channel_chat_id=channel_chat_id,
|
|
)
|
|
return ChannelMessage(identity=identity, content="hello")
|
|
|
|
|
|
class TestSessionMapperInit:
|
|
def test_default_department_id(self):
|
|
db = MagicMock()
|
|
mapper = SessionMapper(db)
|
|
assert mapper.department_id == VIRTUAL_DEPARTMENT_ID
|
|
assert mapper.db is db
|
|
|
|
def test_custom_department_id(self):
|
|
db = MagicMock()
|
|
mapper = SessionMapper(db, department_id=42)
|
|
assert mapper.department_id == 42
|
|
|
|
|
|
class TestResolveUser:
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_user_existing_mapping(self):
|
|
db = AsyncMock()
|
|
message = _make_message()
|
|
mapper = SessionMapper(db)
|
|
mapper._get_user_mapping = AsyncMock(return_value=MagicMock(internal_user_id="existing_user"))
|
|
|
|
user_id = await mapper.resolve_user(message)
|
|
assert user_id == "existing_user"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_user_creates_new_mapping(self):
|
|
db = AsyncMock()
|
|
message = _make_message()
|
|
mapper = SessionMapper(db)
|
|
mapper._get_user_mapping = AsyncMock(return_value=None)
|
|
|
|
user_id = await mapper.resolve_user(message)
|
|
assert user_id.startswith("ch_")
|
|
db.commit.assert_called()
|
|
|
|
|
|
class TestResolveThread:
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_thread_existing(self):
|
|
db = AsyncMock()
|
|
message = _make_message()
|
|
mapper = SessionMapper(db)
|
|
|
|
existing = MagicMock()
|
|
existing.thread_id = "thread_existing"
|
|
mapper._get_thread_mapping = AsyncMock(return_value=existing)
|
|
|
|
thread_id = await mapper.resolve_thread(message, "internal_user")
|
|
assert thread_id == "thread_existing"
|
|
db.flush.assert_called()
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_thread_creates_new(self):
|
|
db = AsyncMock()
|
|
message = _make_message()
|
|
mapper = SessionMapper(db)
|
|
mapper._get_thread_mapping = AsyncMock(return_value=None)
|
|
|
|
thread_id = await mapper.resolve_thread(message, "internal_user")
|
|
assert thread_id is not None
|
|
db.commit.assert_called()
|
|
|
|
|
|
class TestResetThread:
|
|
@pytest.mark.asyncio
|
|
async def test_reset_thread_updates_existing(self):
|
|
db = AsyncMock()
|
|
message = _make_message()
|
|
mapper = SessionMapper(db)
|
|
|
|
existing = MagicMock()
|
|
existing.thread_id = "old_thread"
|
|
mapper._get_thread_mapping = AsyncMock(return_value=existing)
|
|
|
|
thread_id = await mapper.reset_thread(message, "internal_user")
|
|
assert thread_id != "old_thread"
|
|
db.commit.assert_called()
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reset_thread_creates_if_missing(self):
|
|
db = AsyncMock()
|
|
message = _make_message()
|
|
mapper = SessionMapper(db)
|
|
mapper._get_thread_mapping = AsyncMock(return_value=None)
|
|
|
|
thread_id = await mapper.reset_thread(message, "internal_user")
|
|
assert thread_id is not None
|
|
|
|
|
|
class TestIntegrityErrorHandling:
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_user_retry_on_integrity_error(self):
|
|
db = AsyncMock()
|
|
message = _make_message()
|
|
mapper = SessionMapper(db)
|
|
|
|
call_count = [0]
|
|
|
|
async def mock_get_user_mapping(_channel_id, _channel_user_id):
|
|
call_count[0] += 1
|
|
if call_count[0] == 2:
|
|
return MagicMock(internal_user_id="concurrent_user")
|
|
return None
|
|
|
|
mapper._get_user_mapping = mock_get_user_mapping
|
|
db.commit.side_effect = [IntegrityError("fake", {}, None), None]
|
|
|
|
user_id = await mapper.resolve_user(message)
|
|
assert user_id == "concurrent_user"
|
|
db.rollback.assert_called()
|
|
|
|
|
|
class TestUserSource:
|
|
def test_user_source_prefix(self):
|
|
assert USER_SOURCE_PREFIX == "channel:" |