from __future__ import annotations from unittest.mock import AsyncMock, MagicMock, patch import pytest from sqlalchemy.exc import IntegrityError from yuxi.channel.plugins.protocol import ( BindingRoute, InboundMessage, SessionConversationRef, ) from yuxi.channel.session.manager import SessionManager from yuxi.storage.postgres.models_business import User @pytest.fixture def binding_router(): return MagicMock() @pytest.fixture def manager(binding_router): return SessionManager(binding_router=binding_router) @pytest.fixture def inbound(): return InboundMessage( channel_type="feishu", account_id="acc1", sender_id="user1", sender_name="User One", peer_id="user1", chat_type="dm", session_key="feishu:acc1:dm:user1", ) @pytest.fixture def plugin(): mock = MagicMock() ref = SessionConversationRef( session_key="feishu:acc1:dm:user1", chat_type="dm", channel_sender_id="user1", channel_metadata={}, ) mock.resolve_session_conversation = MagicMock(return_value=ref) mock.parse_session_key = MagicMock(return_value="feishu:acc1:dm:user1") mock.resolve_reply_to_mode = MagicMock(return_value=None) mock.resolve_auto_thread_id = MagicMock(return_value=None) return mock @pytest.fixture def db_session(): session = MagicMock() session.bind = None session.execute = AsyncMock() session.flush = AsyncMock() session.rollback = AsyncMock() session.refresh = AsyncMock() return session @pytest.fixture def config(): return { "channel_type": "feishu", "account_id": "acc1", "default_agent_id": "default_agent", } class TestResolveExistingSession: async def test_returns_existing_session_and_none_route(self, manager, plugin, inbound, db_session, config): existing_session = MagicMock() existing_session.session_key = "feishu:acc1:dm:user1" existing_session.channel_metadata = None with patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_repo_cls: mock_repo = MagicMock() mock_repo.get_by_session_key = AsyncMock(return_value=existing_session) mock_repo_cls.return_value = mock_repo session, route = await manager.resolve(config, plugin, inbound, db_session) assert session is existing_session assert route is None mock_repo.get_by_session_key.assert_awaited_once_with("feishu:acc1:dm:user1") async def test_updates_metadata_on_existing_session(self, manager, plugin, inbound, db_session, config): existing_session = MagicMock() existing_session.session_key = "feishu:acc1:dm:user1" existing_session.channel_metadata = {"reply_to_mode": "thread"} plugin.resolve_reply_to_mode.return_value = "message" plugin.resolve_auto_thread_id.return_value = "t1" with ( patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_repo_cls, patch("yuxi.channel.session.manager.flag_modified") as mock_flag_modified, ): mock_repo = MagicMock() mock_repo.get_by_session_key = AsyncMock(return_value=existing_session) mock_repo_cls.return_value = mock_repo await manager.resolve(config, plugin, inbound, db_session) assert existing_session.channel_metadata["reply_to_mode"] == "message" assert existing_session.channel_metadata["auto_thread_id"] == "t1" mock_flag_modified.assert_called_once_with(existing_session, "channel_metadata") db_session.flush.assert_awaited_once() async def test_no_flush_when_metadata_unchanged(self, manager, plugin, inbound, db_session, config): existing_session = MagicMock() existing_session.session_key = "feishu:acc1:dm:user1" existing_session.channel_metadata = {} plugin.resolve_reply_to_mode.return_value = None plugin.resolve_auto_thread_id.return_value = None with patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_repo_cls: mock_repo = MagicMock() mock_repo.get_by_session_key = AsyncMock(return_value=existing_session) mock_repo_cls.return_value = mock_repo await manager.resolve(config, plugin, inbound, db_session) db_session.flush.assert_not_called() class TestResolveNewSession: async def test_creates_session_and_returns_route(self, manager, plugin, inbound, db_session, config): plugin.resolve_session_conversation.return_value = None manager.router.resolve_runtime = AsyncMock( return_value=BindingRoute( agent_id="agent1", session_key="feishu:acc1:dm:user1", matched_by="default", ) ) existing_user_result = MagicMock() existing_user_result.scalar_one_or_none = MagicMock(return_value=None) db_session.execute = AsyncMock(return_value=existing_user_result) conversation = MagicMock() conversation.id = 42 conversation.thread_id = "thread-uuid" new_session = MagicMock() new_session.session_key = "feishu:acc1:dm:user1" with ( patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_session_repo_cls, patch("yuxi.channel.session.manager.ConversationRepository") as mock_conv_repo_cls, ): mock_session_repo = MagicMock() mock_session_repo.get_by_session_key = AsyncMock(return_value=None) mock_session_repo.create = AsyncMock(return_value=new_session) mock_session_repo_cls.return_value = mock_session_repo mock_conv_repo_inst = MagicMock() mock_conv_repo_inst.create_conversation = AsyncMock(return_value=conversation) mock_conv_repo_cls.return_value = mock_conv_repo_inst session, route = await manager.resolve(config, plugin, inbound, db_session) assert session is new_session assert route is not None assert route.agent_id == "agent1" manager.router.resolve_runtime.assert_awaited_once() mock_conv_repo_inst.create_conversation.assert_awaited_once() mock_session_repo.create.assert_awaited_once() added_user = db_session.add.call_args[0][0] assert isinstance(added_user, User) assert added_user.uid == "channel:feishu:acc1" async def test_raises_when_agent_id_unresolvable(self, manager, plugin, inbound, db_session, config): plugin.resolve_session_conversation.return_value = None manager.router.resolve_runtime = AsyncMock( return_value=BindingRoute( agent_id=None, session_key="feishu:acc1:dm:user1", matched_by="default", ) ) config.pop("default_agent_id") existing_user_result = MagicMock() existing_user_result.scalar_one_or_none = MagicMock(return_value=None) db_session.execute = AsyncMock(return_value=existing_user_result) with patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_session_repo_cls: mock_session_repo = MagicMock() mock_session_repo.get_by_session_key = AsyncMock(return_value=None) mock_session_repo_cls.return_value = mock_session_repo with pytest.raises(ValueError, match="Unable to resolve agent_id"): await manager.resolve(config, plugin, inbound, db_session) async def test_uses_fallback_when_plugin_returns_none(self, manager, plugin, inbound, db_session, config): plugin.resolve_session_conversation.return_value = None manager.router.resolve_runtime = AsyncMock( return_value=BindingRoute( agent_id="agent1", session_key="feishu:acc1:dm:user1", matched_by="default", ) ) existing_user_result = MagicMock() existing_user_result.scalar_one_or_none = MagicMock(return_value=None) db_session.execute = AsyncMock(return_value=existing_user_result) conversation = MagicMock() conversation.id = 42 new_session = MagicMock() with ( patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_session_repo_cls, patch("yuxi.channel.session.manager.ConversationRepository") as mock_conv_repo_cls, ): mock_session_repo = MagicMock() mock_session_repo.get_by_session_key = AsyncMock(return_value=None) mock_session_repo.create = AsyncMock(return_value=new_session) mock_session_repo_cls.return_value = mock_session_repo mock_conv_repo_inst = MagicMock() mock_conv_repo_inst.create_conversation = AsyncMock(return_value=conversation) mock_conv_repo_cls.return_value = mock_conv_repo_inst await manager.resolve(config, plugin, inbound, db_session) plugin.parse_session_key.assert_called_once_with(inbound) async def test_handles_integrity_error_on_session_create(self, manager, plugin, inbound, db_session, config): plugin.resolve_session_conversation.return_value = None manager.router.resolve_runtime = AsyncMock( return_value=BindingRoute( agent_id="agent1", session_key="feishu:acc1:dm:user1", matched_by="default", ) ) existing_user_result = MagicMock() existing_user_result.scalar_one_or_none = MagicMock(return_value=None) db_session.execute = AsyncMock(return_value=existing_user_result) recovered_session = MagicMock() recovered_session.session_key = "feishu:acc1:dm:user1" with ( patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_session_repo_cls, patch("yuxi.channel.session.manager.ConversationRepository") as mock_conv_repo_cls, ): mock_session_repo = MagicMock() mock_session_repo.get_by_session_key = AsyncMock(side_effect=[None, recovered_session]) mock_session_repo.create = AsyncMock(side_effect=IntegrityError("", {}, None)) mock_session_repo_cls.return_value = mock_session_repo mock_conv_repo_inst = MagicMock() mock_conv_repo_inst.create_conversation = AsyncMock(return_value=MagicMock(id=42)) mock_conv_repo_cls.return_value = mock_conv_repo_inst session, route = await manager.resolve(config, plugin, inbound, db_session) assert session is recovered_session assert route is None db_session.rollback.assert_awaited_once() db_session.refresh.assert_awaited_once_with(recovered_session) async def test_reraises_integrity_error_when_session_not_recovered( self, manager, plugin, inbound, db_session, config ): plugin.resolve_session_conversation.return_value = None manager.router.resolve_runtime = AsyncMock( return_value=BindingRoute( agent_id="agent1", session_key="feishu:acc1:dm:user1", matched_by="default", ) ) existing_user_result = MagicMock() existing_user_result.scalar_one_or_none = MagicMock(return_value=None) db_session.execute = AsyncMock(return_value=existing_user_result) with ( patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_session_repo_cls, patch("yuxi.channel.session.manager.ConversationRepository") as mock_conv_repo_cls, ): mock_session_repo = MagicMock() mock_session_repo.get_by_session_key = AsyncMock(side_effect=[None, None]) mock_session_repo.create = AsyncMock(side_effect=IntegrityError("", {}, None)) mock_session_repo_cls.return_value = mock_session_repo mock_conv_repo_inst = MagicMock() mock_conv_repo_inst.create_conversation = AsyncMock(return_value=MagicMock(id=42)) mock_conv_repo_cls.return_value = mock_conv_repo_inst with pytest.raises(IntegrityError): await manager.resolve(config, plugin, inbound, db_session) class TestGetOrCreateChannelUser: async def test_returns_existing_user(self, manager, db_session): existing_user = MagicMock() existing_user.uid = "channel:feishu:acc1" result = MagicMock() result.scalar_one_or_none = MagicMock(return_value=existing_user) db_session.execute = AsyncMock(return_value=result) user = await manager._get_or_create_channel_user("feishu", "acc1", db_session) assert user is existing_user async def test_creates_user_when_not_exists(self, manager, db_session): result = MagicMock() result.scalar_one_or_none = MagicMock(return_value=None) db_session.execute = AsyncMock(return_value=result) user = await manager._get_or_create_channel_user("feishu", "acc1", db_session) assert isinstance(user, User) assert user.uid == "channel:feishu:acc1" db_session.add.assert_called_once_with(user) db_session.flush.assert_awaited_once() async def test_recovers_from_integrity_error(self, manager, db_session): result_first = MagicMock() result_first.scalar_one_or_none = MagicMock(return_value=None) result_second = MagicMock() recovered_user = MagicMock() result_second.scalar_one_or_none = MagicMock(return_value=recovered_user) db_session.execute = AsyncMock(side_effect=[result_first, result_second]) db_session.flush = AsyncMock(side_effect=IntegrityError("", {}, None)) user = await manager._get_or_create_channel_user("feishu", "acc1", db_session) assert user is recovered_user db_session.rollback.assert_awaited_once() async def test_reraises_integrity_error_when_user_not_recovered(self, manager, db_session): result = MagicMock() result.scalar_one_or_none = MagicMock(return_value=None) db_session.execute = AsyncMock(return_value=result) db_session.flush = AsyncMock(side_effect=IntegrityError("", {}, None)) with pytest.raises(IntegrityError): await manager._get_or_create_channel_user("feishu", "acc1", db_session) class TestRecordMessage: async def test_delegates_to_repository(self, manager, inbound, db_session): session = MagicMock() with patch("yuxi.channel.session.manager.ChannelSessionRepository") as mock_repo_cls: mock_repo = MagicMock() mock_repo.update_last_message_at = AsyncMock() mock_repo_cls.return_value = mock_repo await manager.record_message(session, inbound, db_session) mock_repo.update_last_message_at.assert_awaited_once_with( session, channel_message_id=inbound.channel_message_id ) class TestSessionKeyLock: async def test_acquire_lock_no_op_without_bind(self, manager, db_session): db_session.bind = None await manager._acquire_session_key_lock(db_session, "key") db_session.execute.assert_not_called() async def test_acquire_lock_no_op_for_non_postgres(self, manager, db_session): bind = MagicMock() bind.dialect.name = "sqlite" db_session.bind = bind await manager._acquire_session_key_lock(db_session, "key") db_session.execute.assert_not_called() async def test_acquire_lock_for_postgres(self, manager, db_session): bind = MagicMock() bind.dialect.name = "postgresql" db_session.bind = bind await manager._acquire_session_key_lock(db_session, "key") db_session.execute.assert_awaited_once() def test_lock_id_is_deterministic(self, manager): assert manager._session_key_lock_id("key") == manager._session_key_lock_id("key") assert manager._session_key_lock_id("key1") != manager._session_key_lock_id("key2") def test_thread_id_is_deterministic(self, manager): assert manager._session_key_to_thread_id("feishu:acc1:dm:user1") == manager._session_key_to_thread_id( "feishu:acc1:dm:user1" )