from __future__ import annotations from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channel.application.service.session_resolver import SessionResolver from yuxi.channel.domain.model.binding.channel_binding import ChannelBinding from yuxi.channel.domain.model.session.channel_session import ChannelSession class TestSessionResolver: @pytest.fixture def mock_session_repo(self): return AsyncMock() @pytest.fixture def mock_binding_repo(self): return AsyncMock() @pytest.fixture def resolver(self, mock_session_repo, mock_binding_repo): return SessionResolver( session_repo=mock_session_repo, binding_repo=mock_binding_repo, default_agent_config_id=1, ) @pytest.fixture def sample_payload(self): return { "channel_type": "web", "sender_id": "user1", "metadata": { "account_id": "acc1", "group_id": "", "is_group": False, }, } @pytest.mark.asyncio async def test_resolve_new_session(self, resolver, mock_session_repo, mock_binding_repo, sample_payload): mock_binding_repo.find_binding.return_value = None mock_session_repo.get_or_create.return_value = ChannelSession( id=1, thread_id="thread123", user_id="user1", agent_id="1", channel_type="web", channel_session_key="user1", status="active", title=None, ) result = await resolver.resolve(sample_payload) assert result is not None assert result.thread_id == "thread123" @pytest.mark.asyncio async def test_resolve_with_binding(self, resolver, mock_session_repo, mock_binding_repo, sample_payload): mock_binding_repo.find_binding.return_value = ChannelBinding( id=1, channel_type="web", account_id="acc1", group_id="", agent_config_id=2, ) mock_session_repo.get_or_create.return_value = ChannelSession( id=1, thread_id="thread123", user_id="user1", agent_id="2", channel_type="web", channel_session_key="user1", status="active", title=None, ) result = await resolver.resolve(sample_payload) assert result is not None mock_session_repo.get_or_create.assert_awaited_once() @pytest.mark.asyncio async def test_resolve_with_agent_config_id(self, resolver, mock_session_repo, mock_binding_repo, sample_payload): sample_payload["agent_config_id"] = 3 mock_binding_repo.find_binding.return_value = None mock_session_repo.get_or_create.return_value = ChannelSession( id=1, thread_id="thread123", user_id="user1", agent_id="3", channel_type="web", channel_session_key="user1", status="active", title=None, ) result = await resolver.resolve(sample_payload) assert result is not None @pytest.mark.asyncio async def test_resolve_group_session(self, resolver, mock_session_repo, mock_binding_repo, sample_payload): sample_payload["metadata"]["is_group"] = True sample_payload["metadata"]["group_id"] = "group1" mock_binding_repo.find_binding.return_value = None mock_session_repo.get_or_create.return_value = ChannelSession( id=1, thread_id="thread123", user_id="user1", agent_id="1", channel_type="web", channel_session_key="user1:web:group1", status="active", title=None, ) result = await resolver.resolve(sample_payload) assert result is not None @pytest.mark.asyncio async def test_resolve_session_error(self, resolver, mock_session_repo, mock_binding_repo, sample_payload): mock_binding_repo.find_binding.return_value = None mock_session_repo.get_or_create.side_effect = Exception("db error") result = await resolver.resolve(sample_payload) assert result is None def test_resolve_session_key_main(self, resolver): payload = { "channel_type": "web", "sender_id": "user1", "metadata": {"is_group": False}, } key = resolver._resolve_session_key(payload, "web", "user1", None) assert key == "user1" def test_resolve_session_key_channel_group(self, resolver): payload = { "channel_type": "web", "sender_id": "user1", "metadata": {"is_group": True, "group_id": "group1"}, } key = resolver._resolve_session_key(payload, "web", "user1", None) assert key == "user1:web:group1" def test_resolve_session_key_custom(self, resolver): payload = { "channel_type": "web", "sender_id": "user1", "metadata": { "session_key_strategy": "custom", "session_key_prefix": "pre", "session_key": "custom_key", }, } key = resolver._resolve_session_key(payload, "web", "user1", None) assert key == "pre:custom_key" def test_resolve_session_key_with_binding(self, resolver): binding = ChannelBinding( id=1, channel_type="web", account_id="acc1", group_id="group1", agent_config_id=1, session_key_strategy="main", ) payload = { "channel_type": "web", "sender_id": "user1", "metadata": {}, } key = resolver._resolve_session_key(payload, "web", "user1", binding) assert key == "user1" @pytest.mark.asyncio async def test_resolve_agent_config_from_binding(self, resolver, mock_binding_repo): mock_binding_repo.find_binding.return_value = ChannelBinding( id=1, channel_type="web", account_id="acc1", group_id="", agent_config_id=5, ) result = await resolver._resolve_agent_config("web", {"account_id": "acc1", "group_id": ""}) assert result == 5 @pytest.mark.asyncio async def test_resolve_agent_config_default(self, resolver, mock_binding_repo): mock_binding_repo.find_binding.return_value = None result = await resolver._resolve_agent_config("web", {"account_id": "acc1", "group_id": ""}) assert result == 1