403 lines
16 KiB
Python
403 lines
16 KiB
Python
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"
|
|
)
|