ForcePilot/backend/test/unit/channel/session/test_manager.py
Kris bab30f2715
Some checks failed
Deploy VitePress site to Pages / build (push) Has been cancelled
Ruff Format Check / Ruff Format & Lint (push) Has been cancelled
Deploy VitePress site to Pages / Deploy (push) Has been cancelled
feat:0715
2026-07-15 12:30:58 +08:00

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"
)