ForcePilot/backend/test/unit/channels/adapters/test_conversation_adapter.py
Kris b323922cdc style: 清理测试文件中的多余空行和导入顺序
本次提交清理了多个单元测试文件中的多余空行,调整了部分导入的顺序,同时修复了一处feishu错误翻译的导入顺序问题,移除了application/conftest.py中未使用的fake_report_repo fixture,优化代码格式提升可读性。
2026-07-06 20:52:20 +08:00

1604 lines
51 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""yuxi.channels.adapters.conversation_adapter 单元测试。
覆盖 ``ConversationAdapter`` 的 ``saveMessage`` /
``updateMessageChannelStatus`` / ``getMessages`` /
``resolveConversation`` / ``associateConversation`` /
``mergeConversations`` / ``getSessionOwner`` /
``transferSessionOwner`` / ``appendOperationHistory`` /
``searchMessages`` / ``markMessageRecalled`` 方法patch
``ChannelConversationRepository`` / ``ChannelMessageRepository`` /
``ChannelSessionRepository`` 返回 mock 桩,使用 ``MagicMock`` 模拟
``AsyncSession``,不连接真实 DB。
"""
from __future__ import annotations
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlalchemy.exc import IntegrityError, SQLAlchemyError
from yuxi.channels.adapters.conversation_adapter import ConversationAdapter
from yuxi.channels.contract.dtos.channel import (
ChannelType,
Message,
MessageSearchItem,
)
from yuxi.channels.contract.dtos.common import Operator, OperatorRole
from yuxi.channels.contract.dtos.conversation import (
AppendOperationHistoryCmd,
AssociateConversationCmd,
MergeConversationCmd,
MergeResult,
ResolveConversationCmd,
SaveMessageCmd,
UpdateMessageChannelStatusCmd,
)
from yuxi.channels.contract.dtos.dashboard import MessageStats
from yuxi.channels.contract.dtos.session import (
ChatType,
OwnerTransferCmd,
SessionOwner,
)
from yuxi.channels.contract.errors import (
ConflictError,
DependencyError,
NotFoundError,
ValidationError,
)
pytestmark = pytest.mark.unit
def _make_db() -> MagicMock:
"""构造 AsyncSession 桩。"""
db = MagicMock()
db.add = MagicMock()
db.flush = AsyncMock()
db.commit = AsyncMock()
db.rollback = AsyncMock()
db.refresh = AsyncMock()
db.execute = AsyncMock()
db.get = AsyncMock()
return db
def _make_operator() -> Operator:
"""构造操作人 DTO。"""
return Operator(
user_id="admin-1",
role=OperatorRole.SUPERADMIN_USER,
ip="127.0.0.1",
request_id="req-1",
)
def _make_save_cmd(*, conversation_id: str = "100") -> SaveMessageCmd:
"""构造保存消息命令。"""
return SaveMessageCmd(
conversation_id=conversation_id,
role="assistant",
content="hello",
channel_status="sent",
channel_msg_id="cmid-1",
)
def _make_update_status_cmd(*, message_id: str = "200") -> UpdateMessageChannelStatusCmd:
"""构造状态回写命令。"""
return UpdateMessageChannelStatusCmd(
message_id=message_id,
channel_status="read",
event={"status": "read", "at": datetime(2026, 1, 1).isoformat()},
read_at=datetime(2026, 1, 1),
)
def _make_resolve_cmd(
*,
create_if_not_found: bool = False,
unified_identity_id: str | None = None,
) -> ResolveConversationCmd:
"""构造解析会话命令。"""
return ResolveConversationCmd(
peer_id="peer-1",
channel_type=ChannelType("feishu"),
account_id="acc-1",
chat_type=ChatType.P2P,
unified_identity_id=unified_identity_id,
create_if_not_found=create_if_not_found,
)
def _make_associate_cmd(*, channel_session_id: str | None = "sess-1") -> AssociateConversationCmd:
"""构造关联会话命令。"""
return AssociateConversationCmd(
unified_identity_id="uid-1",
route_binding={},
channel_session_id=channel_session_id,
)
def _make_merge_cmd() -> MergeConversationCmd:
"""构造合并会话命令。"""
return MergeConversationCmd(
source_conversation_id="100",
target_conversation_id="200",
operator=_make_operator(),
reason="duplicate",
)
def _make_owner_transfer_cmd() -> OwnerTransferCmd:
"""构造所有者转移命令。"""
return OwnerTransferCmd(
conversation_id="100",
new_owner_id="peer-2",
operator=_make_operator(),
)
def _make_append_op_cmd(*, message_id: str = "200") -> AppendOperationHistoryCmd:
"""构造追加操作历史命令。"""
return AppendOperationHistoryCmd(
message_id=message_id,
entry={
"operation": "recall",
"at": datetime(2026, 1, 1).isoformat(),
"success": True,
},
)
def _make_message_orm(*, message_id: int = 200, conversation_id: int = 100) -> MagicMock:
"""构造 Message ORM 桩。"""
orm = MagicMock()
orm.id = message_id
orm.conversation_id = conversation_id
orm.role = "assistant"
orm.content = "hello"
orm.channel_status = "sent"
orm.channel_msg_id = "cmid-1"
orm.ref_channel_msg_id = None
orm.channel_status_history = []
orm.operations_history = []
orm.channel_read_at = None
orm.channel_recalled_at = None
orm.channel_edited_at = None
orm.created_at = datetime(2026, 1, 1)
orm.conversation = MagicMock()
orm.conversation.channel_type = "feishu"
return orm
def _make_session_orm(
*,
session_id: str = "sess-1",
conversation_id: int | None = 100,
owner_peer_id: str | None = None,
is_temporary: bool = False,
) -> MagicMock:
"""构造 ChannelSession ORM 桩。
默认 ``owner_peer_id`` / ``deleted_at`` / ``closed_at`` 均为 None
使 ``_orm_to_session_aggregate`` 重构的聚合根可通过 ``canBeMerged()```
校验,供合并测试使用。
"""
orm = MagicMock()
orm.id = 1
orm.session_id = session_id
orm.channel_type = "feishu"
orm.account_id = 1
orm.peer_id = "peer-1"
orm.chat_type = "p2p"
orm.conversation_id = conversation_id
orm.unified_identity_id = None
orm.owner_peer_id = owner_peer_id
orm.is_temporary = is_temporary
orm.created_at = datetime(2026, 1, 1)
orm.updated_at = datetime(2026, 1, 1)
orm.deleted_at = None
orm.closed_at = None
orm.version = 1
return orm
def _make_conversation_orm(*, conv_id: int = 100, status: str = "active") -> MagicMock:
"""构造 Conversation ORM 桩。"""
orm = MagicMock()
orm.id = conv_id
orm.thread_id = "thread-1"
orm.uid = "peer-1"
orm.agent_id = "agent-1"
orm.title = "title"
orm.status = status
orm.extra_metadata = {}
orm.unified_identity_id = None
orm.channel_type = "feishu"
orm.channel_account_id = "acc-1"
orm.channel_session_id = ""
orm.created_at = datetime(2026, 1, 1)
orm.updated_at = datetime(2026, 1, 1)
return orm
def _build_adapter(
db: MagicMock,
) -> tuple[ConversationAdapter, MagicMock, MagicMock, MagicMock]:
"""构造 adapterpatch 三个 repo 类返回桩。
返回 ``(adapter, message_repo, conversation_repo, session_repo)``
调用方可按需在各 repo 桩上挂载 ``AsyncMock`` 方法。
"""
message_repo = MagicMock()
conversation_repo = MagicMock()
session_repo = MagicMock()
with (
patch(
"yuxi.channels.adapters.conversation_adapter.ChannelMessageRepository",
return_value=message_repo,
),
patch(
"yuxi.channels.adapters.conversation_adapter.ChannelConversationRepository",
return_value=conversation_repo,
),
patch(
"yuxi.channels.adapters.conversation_adapter.ChannelSessionRepository",
return_value=session_repo,
),
):
adapter = ConversationAdapter(db, MagicMock())
return adapter, message_repo, conversation_repo, session_repo
def _make_scalar_result(value):
"""构造 ``scalar_one_or_none`` 返回 ``value`` 的结果桩。"""
result = MagicMock()
result.scalar_one_or_none.return_value = value
return result
def _make_scalars_result(values):
"""构造 ``scalars().all()`` 返回 ``values`` 的结果桩。"""
result = MagicMock()
result.scalars.return_value.all.return_value = values
return result
@pytest.mark.unit
class TestConversationSaveMessage:
@pytest.mark.asyncio
async def test_save_message_returns_message_when_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(_make_message_orm()))
adapter, message_repo, _, _ = _build_adapter(db)
message_orm = _make_message_orm()
message_repo.init_channel_fields = AsyncMock(return_value=message_orm)
# Act
message = await adapter.saveMessage(_make_save_cmd())
# Assert
# 源码 saveMessage 委托 message_repo.init_channel_fields(commit=True)
# 自主提交,不直接调用 db.commit
assert isinstance(message, Message)
assert message.message_id == "200"
message_repo.init_channel_fields.assert_awaited_once()
assert message_repo.init_channel_fields.call_args.kwargs["commit"] is True
@pytest.mark.asyncio
async def test_save_message_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(_make_message_orm()))
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.init_channel_fields = AsyncMock(return_value=_make_message_orm())
tx = MagicMock()
# Act
await adapter.saveMessage(_make_save_cmd(), tx=tx)
# Assert
db.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_save_message_raises_not_found_when_no_latest_message(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(None))
adapter, message_repo, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.saveMessage(_make_save_cmd())
@pytest.mark.asyncio
async def test_save_message_raises_not_found_when_invalid_conversation_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.saveMessage(_make_save_cmd(conversation_id="abc"))
@pytest.mark.asyncio
async def test_save_message_translates_integrity_error_to_conflict(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(_make_message_orm()))
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.init_channel_fields = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
# Act / Assert
with pytest.raises(ConflictError):
await adapter.saveMessage(_make_save_cmd())
@pytest.mark.asyncio
async def test_save_message_translates_sqlalchemy_error_to_dependency(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.saveMessage(_make_save_cmd())
@pytest.mark.unit
class TestConversationUpdateMessageChannelStatus:
@pytest.mark.asyncio
async def test_update_status_returns_message_when_found(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(return_value=_make_message_orm())
# Act
message = await adapter.updateMessageChannelStatus(_make_update_status_cmd())
# Assert
# 源码 updateMessageChannelStatus 委托 message_repo.update_channel_status(commit=True)
# 自主提交,不直接调用 db.commit
assert isinstance(message, Message)
assert message.message_id == "200"
message_repo.update_channel_status.assert_awaited_once()
assert message_repo.update_channel_status.call_args.kwargs["commit"] is True
@pytest.mark.asyncio
async def test_update_status_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(return_value=_make_message_orm())
tx = MagicMock()
# Act
await adapter.updateMessageChannelStatus(_make_update_status_cmd(), tx=tx)
# Assert
db.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_status_raises_not_found_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.updateMessageChannelStatus(_make_update_status_cmd(message_id="abc"))
@pytest.mark.asyncio
async def test_update_status_raises_not_found_when_message_missing(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(return_value=None)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.updateMessageChannelStatus(_make_update_status_cmd())
@pytest.mark.asyncio
async def test_update_status_translates_integrity_error_to_conflict(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
# Act / Assert
with pytest.raises(ConflictError):
await adapter.updateMessageChannelStatus(_make_update_status_cmd())
@pytest.mark.asyncio
async def test_update_status_translates_sqlalchemy_error_to_dependency(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(side_effect=SQLAlchemyError("db failure"))
# Act / Assert
with pytest.raises(DependencyError):
await adapter.updateMessageChannelStatus(_make_update_status_cmd())
@pytest.mark.unit
class TestConversationGetMessages:
@pytest.mark.asyncio
async def test_get_messages_returns_tuple(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(
return_value=_make_scalars_result([_make_message_orm(message_id=200), _make_message_orm(message_id=201)])
)
adapter, _, _, _ = _build_adapter(db)
# Act
messages = await adapter.getMessages("100")
# Assert
assert isinstance(messages, tuple)
assert len(messages) == 2
assert all(isinstance(m, Message) for m in messages)
@pytest.mark.asyncio
async def test_get_messages_returns_empty_tuple_when_no_records(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalars_result([]))
adapter, _, _, _ = _build_adapter(db)
# Act
messages = await adapter.getMessages("100")
# Assert
assert messages == ()
@pytest.mark.asyncio
async def test_get_messages_raises_validation_error_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(ValidationError):
await adapter.getMessages("abc")
@pytest.mark.asyncio
async def test_get_messages_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.getMessages("100")
@pytest.mark.unit
class TestConversationResolveConversation:
@pytest.mark.asyncio
async def test_resolve_returns_existing_session_conversation_id(self):
# Arrange
db = _make_db()
session_orm = _make_session_orm(conversation_id=100)
db.execute = AsyncMock(
side_effect=[
_make_scalar_result(1),
_make_scalar_result(session_orm),
]
)
adapter, _, _, _ = _build_adapter(db)
# Act
cid = await adapter.resolveConversation(_make_resolve_cmd())
# Assert
assert cid == "100"
@pytest.mark.asyncio
async def test_resolve_raises_not_found_when_account_missing(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(None))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.resolveConversation(_make_resolve_cmd())
@pytest.mark.asyncio
async def test_resolve_raises_not_found_when_create_false_and_no_session(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(
side_effect=[
_make_scalar_result(1),
_make_scalar_result(None),
]
)
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.resolveConversation(_make_resolve_cmd(create_if_not_found=False))
@pytest.mark.asyncio
async def test_resolve_creates_new_conversation_when_create_true(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(
side_effect=[
_make_scalar_result(1),
_make_scalar_result(None),
]
)
adapter, _, _, _ = _build_adapter(db)
new_conv = _make_conversation_orm(conv_id=999)
# Act
with (
patch(
"yuxi.channels.adapters.conversation_adapter.ConversationORM",
return_value=new_conv,
),
patch("yuxi.channels.adapters.conversation_adapter.ConversationStatsORM"),
):
cid = await adapter.resolveConversation(_make_resolve_cmd(create_if_not_found=True))
# Assert
assert cid == "999"
db.add.assert_called()
db.flush.assert_awaited_once()
db.commit.assert_awaited_once()
db.refresh.assert_awaited_once()
@pytest.mark.asyncio
async def test_resolve_returns_existing_conversation_by_identity(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(
side_effect=[
_make_scalar_result(1),
_make_scalar_result(None),
MagicMock(),
]
)
adapter, _, conversation_repo, _ = _build_adapter(db)
existing_conv = _make_conversation_orm(conv_id=888)
conversation_repo.find_by_unified_identity = AsyncMock(return_value=[existing_conv])
# Act
cid = await adapter.resolveConversation(_make_resolve_cmd(unified_identity_id="uid-1"))
# Assert
assert cid == "888"
@pytest.mark.asyncio
async def test_resolve_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.resolveConversation(_make_resolve_cmd())
@pytest.mark.unit
class TestConversationAssociateConversation:
@pytest.mark.asyncio
async def test_associate_links_conversation(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=MagicMock())
adapter, _, conversation_repo, session_repo = _build_adapter(db)
conv_orm = _make_conversation_orm()
conversation_repo.find_by_unified_identity = AsyncMock(return_value=[conv_orm])
session_orm = _make_session_orm()
session_repo.get_by_session_id = AsyncMock(return_value=session_orm)
conversation_repo.link_channel_session = AsyncMock(return_value=conv_orm)
# Act
await adapter.associateConversation(_make_associate_cmd())
# Assert
# 源码 associateConversation 委托 conversation_repo.link_channel_session(commit=True)
# 自主提交,不直接调用 db.commit
conversation_repo.link_channel_session.assert_awaited_once()
assert conversation_repo.link_channel_session.call_args.kwargs["commit"] is True
@pytest.mark.asyncio
async def test_associate_raises_not_found_when_no_conversation(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=MagicMock())
adapter, _, conversation_repo, _ = _build_adapter(db)
conversation_repo.find_by_unified_identity = AsyncMock(return_value=[])
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.associateConversation(_make_associate_cmd())
@pytest.mark.asyncio
async def test_associate_raises_not_found_when_session_missing(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=MagicMock())
adapter, _, conversation_repo, session_repo = _build_adapter(db)
conv_orm = _make_conversation_orm()
conversation_repo.find_by_unified_identity = AsyncMock(return_value=[conv_orm])
session_repo.get_by_session_id = AsyncMock(return_value=None)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.associateConversation(_make_associate_cmd())
@pytest.mark.asyncio
async def test_associate_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.associateConversation(_make_associate_cmd())
@pytest.mark.unit
class TestConversationMergeConversations:
@pytest.mark.asyncio
async def test_merge_returns_result_when_successful(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=MagicMock(rowcount=3))
db.get = AsyncMock(
side_effect=[
_make_conversation_orm(conv_id=100, status="active"),
_make_conversation_orm(conv_id=200, status="active"),
]
)
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[_make_session_orm()])
session_repo.delete_by_id = AsyncMock()
# Act
result = await adapter.mergeConversations(_make_merge_cmd())
# Assert
assert isinstance(result, MergeResult)
assert result.migrated_message_count == 3
assert result.source_soft_deleted is True
db.commit.assert_awaited_once()
@pytest.mark.asyncio
async def test_merge_returns_idempotent_when_source_deleted(self):
# Arrange
db = _make_db()
db.get = AsyncMock(
side_effect=[
_make_conversation_orm(conv_id=100, status="deleted"),
_make_conversation_orm(conv_id=200, status="active"),
]
)
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.mergeConversations(_make_merge_cmd())
# Assert
assert result.migrated_message_count == 0
assert result.source_soft_deleted is True
db.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_merge_raises_not_found_when_source_missing(self):
# Arrange
db = _make_db()
db.get = AsyncMock(
side_effect=[
None,
_make_conversation_orm(conv_id=200),
]
)
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.mergeConversations(_make_merge_cmd())
@pytest.mark.asyncio
async def test_merge_raises_not_found_when_target_missing(self):
# Arrange
db = _make_db()
db.get = AsyncMock(
side_effect=[
_make_conversation_orm(conv_id=100),
None,
]
)
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.mergeConversations(_make_merge_cmd())
@pytest.mark.asyncio
async def test_merge_raises_not_found_when_invalid_source_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
cmd = MergeConversationCmd(
source_conversation_id="abc",
target_conversation_id="200",
operator=_make_operator(),
reason="duplicate",
)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.mergeConversations(cmd)
@pytest.mark.asyncio
async def test_merge_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=MagicMock(rowcount=1))
db.get = AsyncMock(
side_effect=[
_make_conversation_orm(conv_id=100, status="active"),
_make_conversation_orm(conv_id=200, status="active"),
]
)
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[_make_session_orm()])
session_repo.delete_by_id = AsyncMock()
tx = MagicMock()
# Act
await adapter.mergeConversations(_make_merge_cmd(), tx=tx)
# Assert
db.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_merge_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.mergeConversations(_make_merge_cmd())
@pytest.mark.unit
class TestConversationGetSessionOwner:
@pytest.mark.asyncio
async def test_get_session_owner_returns_owner(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[_make_session_orm(owner_peer_id="peer-1")])
# Act
owner = await adapter.getSessionOwner("100")
# Assert
assert isinstance(owner, SessionOwner)
assert owner.owner_peer_id == "peer-1"
assert owner.conversation_id == "100"
@pytest.mark.asyncio
async def test_get_session_owner_returns_none_when_no_sessions(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[])
# Act
owner = await adapter.getSessionOwner("100")
# Assert
assert owner is None
@pytest.mark.asyncio
async def test_get_session_owner_returns_none_when_no_owner(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[_make_session_orm(owner_peer_id=None)])
# Act
owner = await adapter.getSessionOwner("100")
# Assert
assert owner is None
@pytest.mark.asyncio
async def test_get_session_owner_raises_not_found_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.getSessionOwner("abc")
@pytest.mark.asyncio
async def test_get_session_owner_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(side_effect=SQLAlchemyError("db failure"))
# Act / Assert
with pytest.raises(DependencyError):
await adapter.getSessionOwner("100")
@pytest.mark.unit
class TestConversationTransferSessionOwner:
@pytest.mark.asyncio
async def test_transfer_returns_owner(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_orm = _make_session_orm(owner_peer_id="peer-1")
updated_orm = _make_session_orm(owner_peer_id="peer-2")
session_repo.list_by_conversation = AsyncMock(return_value=[session_orm])
session_repo.update = AsyncMock(return_value=updated_orm)
# Act
owner = await adapter.transferSessionOwner(_make_owner_transfer_cmd())
# Assert
assert isinstance(owner, SessionOwner)
assert owner.owner_peer_id == "peer-2"
db.commit.assert_awaited_once()
@pytest.mark.asyncio
async def test_transfer_raises_not_found_when_no_sessions(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[])
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.transferSessionOwner(_make_owner_transfer_cmd())
@pytest.mark.asyncio
async def test_transfer_raises_not_found_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
cmd = OwnerTransferCmd(
conversation_id="abc",
new_owner_id="peer-2",
operator=_make_operator(),
)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.transferSessionOwner(cmd)
@pytest.mark.asyncio
async def test_transfer_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_orm = _make_session_orm()
session_repo.list_by_conversation = AsyncMock(return_value=[session_orm])
session_repo.update = AsyncMock(return_value=session_orm)
tx = MagicMock()
# Act
await adapter.transferSessionOwner(_make_owner_transfer_cmd(), tx=tx)
# Assert
db.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_transfer_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(side_effect=SQLAlchemyError("db failure"))
# Act / Assert
with pytest.raises(DependencyError):
await adapter.transferSessionOwner(_make_owner_transfer_cmd())
@pytest.mark.unit
class TestConversationAppendOperationHistory:
@pytest.mark.asyncio
async def test_append_op_returns_message(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.append_operation_history = AsyncMock(return_value=_make_message_orm())
# Act
message = await adapter.appendOperationHistory(_make_append_op_cmd())
# Assert
# 源码 appendOperationHistory 委托 message_repo.append_operation_history(commit=True)
# 自主提交,不直接调用 db.commit
assert isinstance(message, Message)
assert message.message_id == "200"
message_repo.append_operation_history.assert_awaited_once()
assert message_repo.append_operation_history.call_args.kwargs["commit"] is True
@pytest.mark.asyncio
async def test_append_op_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.append_operation_history = AsyncMock(return_value=_make_message_orm())
tx = MagicMock()
# Act
await adapter.appendOperationHistory(_make_append_op_cmd(), tx=tx)
# Assert
db.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_append_op_raises_not_found_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.appendOperationHistory(_make_append_op_cmd(message_id="abc"))
@pytest.mark.asyncio
async def test_append_op_raises_not_found_when_message_missing(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.append_operation_history = AsyncMock(return_value=None)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.appendOperationHistory(_make_append_op_cmd())
@pytest.mark.asyncio
async def test_append_op_translates_integrity_error_to_conflict(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.append_operation_history = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
# Act / Assert
with pytest.raises(ConflictError):
await adapter.appendOperationHistory(_make_append_op_cmd())
@pytest.mark.asyncio
async def test_append_op_translates_sqlalchemy_error_to_dependency(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.append_operation_history = AsyncMock(side_effect=SQLAlchemyError("db failure"))
# Act / Assert
with pytest.raises(DependencyError):
await adapter.appendOperationHistory(_make_append_op_cmd())
@pytest.mark.unit
class TestConversationSearchMessages:
@pytest.mark.asyncio
async def test_search_returns_items(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(
return_value=_make_scalars_result([_make_message_orm(message_id=200), _make_message_orm(message_id=201)])
)
adapter, _, _, _ = _build_adapter(db)
# Act
items = await adapter.searchMessages("hello")
# Assert
assert len(items) == 2
assert all(isinstance(i, MessageSearchItem) for i in items)
assert items[0].message_id == "200"
@pytest.mark.asyncio
async def test_search_returns_empty_list_when_no_matches(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalars_result([]))
adapter, _, _, _ = _build_adapter(db)
# Act
items = await adapter.searchMessages("nomatch")
# Assert
assert items == []
@pytest.mark.asyncio
async def test_search_with_channel_session_id_fills_session_id(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalars_result([_make_message_orm(message_id=200)]))
adapter, _, _, _ = _build_adapter(db)
# Act
items = await adapter.searchMessages("hello", channel_session_id="sess-1")
# Assert
assert len(items) == 1
assert items[0].channel_session_id == "sess-1"
@pytest.mark.asyncio
async def test_search_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.searchMessages("hello")
@pytest.mark.unit
class TestConversationMarkMessageRecalled:
@pytest.mark.asyncio
async def test_mark_recalled_succeeds(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(return_value=_make_message_orm())
# Act
await adapter.markMessageRecalled("200", datetime(2026, 1, 1))
# Assert
# 源码 markMessageRecalled 委托 message_repo.update_channel_status(commit=True)
# 自主提交,不直接调用 db.commit
message_repo.update_channel_status.assert_awaited_once()
assert message_repo.update_channel_status.call_args.kwargs["commit"] is True
@pytest.mark.asyncio
async def test_mark_recalled_skips_commit_when_tx_provided(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(return_value=_make_message_orm())
tx = MagicMock()
# Act
await adapter.markMessageRecalled("200", datetime(2026, 1, 1), tx=tx)
# Assert
db.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_mark_recalled_raises_not_found_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.markMessageRecalled("abc", datetime(2026, 1, 1))
@pytest.mark.asyncio
async def test_mark_recalled_raises_not_found_when_message_missing(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(return_value=None)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.markMessageRecalled("200", datetime(2026, 1, 1))
@pytest.mark.asyncio
async def test_mark_recalled_translates_integrity_error_to_conflict(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
# Act / Assert
with pytest.raises(ConflictError):
await adapter.markMessageRecalled("200", datetime(2026, 1, 1))
@pytest.mark.asyncio
async def test_mark_recalled_translates_sqlalchemy_error_to_dependency(self):
# Arrange
db = _make_db()
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.update_channel_status = AsyncMock(side_effect=SQLAlchemyError("db failure"))
# Act / Assert
with pytest.raises(DependencyError):
await adapter.markMessageRecalled("200", datetime(2026, 1, 1))
def _make_scalar_one_result(value):
"""构造 ``scalar_one`` 返回 ``value`` 的结果桩(用于 count 查询)。"""
result = MagicMock()
result.scalar_one.return_value = value
return result
def _make_rows_result(rows):
"""构造 ``all()`` 返回 ``rows`` 的结果桩(用于 getMessageStats 聚合)。"""
result = MagicMock()
result.all.return_value = rows
return result
def _make_stats_row(channel_type, role, channel_status, count):
"""构造聚合查询行桩。"""
row = MagicMock()
row.channel_type = channel_type
row.role = role
row.channel_status = channel_status
row.count = count
return row
@pytest.mark.unit
class TestConversationGetMessageByChannelMsgId:
"""``getMessageByChannelMsgId`` 按渠道消息 ID 定位消息。"""
@pytest.mark.asyncio
async def test_returns_none_when_channel_msg_id_empty(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageByChannelMsgId("")
# Assert
assert result is None
db.execute.assert_not_awaited()
@pytest.mark.asyncio
async def test_returns_message_when_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(_make_message_orm()))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageByChannelMsgId("cmid-1")
# Assert
assert isinstance(result, Message)
assert result.channel_msg_id == "cmid-1"
@pytest.mark.asyncio
async def test_returns_none_when_not_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(None))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageByChannelMsgId("missing")
# Assert
assert result is None
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.getMessageByChannelMsgId("cmid-1")
@pytest.mark.unit
class TestConversationCountMessages:
"""``countMessages`` 统计会话消息总数(分页用例)。"""
@pytest.mark.asyncio
async def test_returns_count(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_one_result(42))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.countMessages("100")
# Assert
assert result == 42
@pytest.mark.asyncio
async def test_raises_validation_error_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(ValidationError):
await adapter.countMessages("abc")
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.countMessages("100")
@pytest.mark.unit
class TestConversationIsTemporarySession:
"""``isTemporarySession`` 检测会话是否为临时会话。"""
@pytest.mark.asyncio
async def test_returns_true_when_session_temporary(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[_make_session_orm(is_temporary=True)])
# Act
result = await adapter.isTemporarySession("100")
# Assert
assert result is True
@pytest.mark.asyncio
async def test_returns_false_when_session_not_temporary(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[_make_session_orm(is_temporary=False)])
# Act
result = await adapter.isTemporarySession("100")
# Assert
assert result is False
@pytest.mark.asyncio
async def test_returns_false_when_no_sessions(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(return_value=[])
# Act
result = await adapter.isTemporarySession("100")
# Assert
assert result is False
@pytest.mark.asyncio
async def test_raises_validation_error_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(ValidationError):
await adapter.isTemporarySession("abc")
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
adapter, _, _, session_repo = _build_adapter(db)
session_repo.list_by_conversation = AsyncMock(side_effect=SQLAlchemyError("db failure"))
# Act / Assert
with pytest.raises(DependencyError):
await adapter.isTemporarySession("100")
@pytest.mark.unit
class TestConversationGetMessageById:
"""``getMessageById`` 按内部消息 ID 查询消息。"""
@pytest.mark.asyncio
async def test_returns_message_when_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(_make_message_orm()))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageById("200")
# Assert
assert isinstance(result, Message)
assert result.message_id == "200"
@pytest.mark.asyncio
async def test_returns_none_when_not_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(None))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageById("999")
# Assert
assert result is None
@pytest.mark.asyncio
async def test_raises_not_found_error_when_invalid_id(self):
# Arrange
db = _make_db()
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(NotFoundError):
await adapter.getMessageById("abc")
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.getMessageById("200")
@pytest.mark.unit
class TestConversationListMessages:
"""``listMessages`` 按条件跨会话查询消息列表。"""
@pytest.mark.asyncio
async def test_returns_list_when_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(
return_value=_make_scalars_result([_make_message_orm(message_id=200), _make_message_orm(message_id=201)])
)
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.listMessages(role="assistant")
# Assert
assert isinstance(result, list)
assert len(result) == 2
assert all(isinstance(m, Message) for m in result)
@pytest.mark.asyncio
async def test_returns_empty_list_when_no_matches(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalars_result([]))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.listMessages(role="assistant")
# Assert
assert result == []
@pytest.mark.asyncio
async def test_uses_default_limit_and_offset(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalars_result([]))
adapter, _, _, _ = _build_adapter(db)
# Act
await adapter.listMessages()
# Assert
# 默认 limit=50, offset=0
db.execute.assert_awaited_once()
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.listMessages(role="assistant")
@pytest.mark.unit
class TestConversationCountMessagesByRole:
"""``countMessagesByRole`` 按角色统计消息总数。"""
@pytest.mark.asyncio
async def test_returns_count(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_one_result(10))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.countMessagesByRole(role="assistant")
# Assert
assert result == 10
@pytest.mark.asyncio
async def test_returns_zero_when_no_matches(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_one_result(0))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.countMessagesByRole(role="assistant")
# Assert
assert result == 0
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.countMessagesByRole(role="assistant")
@pytest.mark.unit
class TestConversationGetMessageStats:
"""``getMessageStats`` 聚合统计消息(按渠道/角色/投递状态分组)。"""
@pytest.mark.asyncio
async def test_aggregates_rows_into_stats(self):
# Arrange
db = _make_db()
rows = [
_make_stats_row("feishu", "user", "sent", 5),
_make_stats_row("feishu", "assistant", "delivered", 3),
_make_stats_row("wecom", "user", "read", 2),
]
db.execute = AsyncMock(return_value=_make_rows_result(rows))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageStats()
# Assert
assert isinstance(result, MessageStats)
assert result.total == 10
assert result.by_channel == {"feishu": 8, "wecom": 2}
assert result.by_role == {"user": 7, "assistant": 3}
assert result.by_delivery_status == {
"sent": 5,
"delivered": 3,
"read": 2,
"recalled": 0,
"edited": 0,
}
@pytest.mark.asyncio
async def test_pre_initializes_delivery_status_dict(self):
# Arrange
db = _make_db()
# 空结果集:验证 by_delivery_status 仍包含全部已知状态键
db.execute = AsyncMock(return_value=_make_rows_result([]))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageStats()
# Assert
assert result.total == 0
assert result.by_channel == {}
assert result.by_role == {}
assert result.by_delivery_status == {
"sent": 0,
"delivered": 0,
"read": 0,
"recalled": 0,
"edited": 0,
}
@pytest.mark.asyncio
async def test_skips_unknown_delivery_status(self):
# Arrange
db = _make_db()
# NULL/未知 channel_status 不计入 by_delivery_status 但计入 total
rows = [
_make_stats_row("feishu", "user", None, 4),
_make_stats_row("feishu", "assistant", "unknown_status", 1),
]
db.execute = AsyncMock(return_value=_make_rows_result(rows))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageStats()
# Assert
assert result.total == 5
assert result.by_delivery_status == {
"sent": 0,
"delivered": 0,
"read": 0,
"recalled": 0,
"edited": 0,
}
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.getMessageStats()
@pytest.mark.unit
class TestConversationCountSearchMessages:
"""``countSearchMessages`` 统计搜索匹配的消息总数。"""
@pytest.mark.asyncio
async def test_returns_count(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_one_result(7))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.countSearchMessages("hello")
# Assert
assert result == 7
@pytest.mark.asyncio
async def test_returns_zero_when_no_matches(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_one_result(0))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.countSearchMessages("nomatch")
# Assert
assert result == 0
@pytest.mark.asyncio
async def test_translates_sqlalchemy_error(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(side_effect=SQLAlchemyError("db failure"))
adapter, _, _, _ = _build_adapter(db)
# Act / Assert
with pytest.raises(DependencyError):
await adapter.countSearchMessages("hello")