ForcePilot/backend/test/unit/channels/adapters/test_conversation_adapter.py

1864 lines
63 KiB
Python
Raw Normal View History

"""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",
role: str = "assistant",
content: str = "hello",
) -> SaveMessageCmd:
"""构造保存消息命令。"""
return SaveMessageCmd(
conversation_id=conversation_id,
role=role,
content=content,
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()
logger = AsyncMock()
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, logger)
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
def _make_session_row(*, session_id: str = "sess-1", peer_id: str = "peer-1", channel_account_id: str = "acc-1"):
"""构造 LATERAL 子查询返回的主渠道会话行桩。"""
row = MagicMock()
row.session_id = session_id
row.peer_id = peer_id
row.channel_account_id = channel_account_id
return row
def _make_message_row(message_orm, conversation_orm=None, session=None):
"""构造 ``(MessageORM, ConversationORM, session)`` 三元组行桩。
``listMessages`` / ``searchMessages`` / ``getMessageById`` 通过 join
``ConversationORM`` LATERAL 子查询返回此类行测试中需模拟行的属性
访问与下标访问``row[2]`` session
"""
row = MagicMock()
row.MessageORM = message_orm
row.ConversationORM = conversation_orm
row.__getitem__.return_value = session
return row
def _make_one_or_none_result(value):
"""构造 ``one_or_none`` 返回 ``value`` 的结果桩。"""
result = MagicMock()
result.one_or_none.return_value = value
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
# init_channel_fields 的 commit 会过期 conversation 关系,
# 需 refresh 重新加载以避免 lazy-load 抛 MissingGreenlet
db.refresh.assert_awaited_once_with(message_orm, ["conversation"])
@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_orm = _make_message_orm()
message_repo.init_channel_fields = AsyncMock(return_value=message_orm)
tx = MagicMock()
# Act
await adapter.saveMessage(_make_save_cmd(), tx=tx)
# Assert
db.commit.assert_not_awaited()
# 即使 flush 也需 refresh conversation统一防御 lazy-load
db.refresh.assert_awaited_once_with(message_orm, ["conversation"])
@pytest.mark.asyncio
async def test_save_message_raises_not_found_when_init_returns_none(self):
"""入站回复路径init_channel_fields 返回 None消息被并发删除时抛 NotFoundError。"""
db = _make_db()
found_orm = _make_message_orm(message_id=5)
db.execute = AsyncMock(return_value=_make_scalar_result(found_orm))
adapter, message_repo, _, _ = _build_adapter(db)
message_repo.init_channel_fields = AsyncMock(return_value=None)
with pytest.raises(NotFoundError) as exc_info:
await adapter.saveMessage(_make_save_cmd())
assert exc_info.value.details["id"] == "5"
# 不应调用 refreshlatest is None
db.refresh.assert_not_awaited()
@pytest.mark.asyncio
async def test_save_message_creates_new_when_no_latest_message(self):
"""管理员直发路径:会话无消息时创建新 Message 记录。
验证 ``saveMessage`` ``_findLatestMessage`` 返回 None
1. 调用 ``db.add`` 添加新 MessageORM
2. 调用 ``db.commit````commit=True``
3. 返回的 Message 携带 cmd 中的 ``content`` ``channel_status``
"""
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(None))
adapter, message_repo, _, _ = _build_adapter(db)
cmd = _make_save_cmd(role="admin", content="亲亲")
# Act
message = await adapter.saveMessage(cmd)
# Assert
assert isinstance(message, Message)
# 新增 MessageORM 至 session
db.add.assert_called_once()
added_orm = db.add.call_args.args[0]
assert added_orm.role == "admin"
assert added_orm.content == "亲亲"
assert added_orm.channel_status == "sent"
# commit=True 自主提交
db.commit.assert_awaited_once()
# 不应调用 init_channel_fields新消息在构造时已设置渠道字段
message_repo.init_channel_fields.assert_not_called()
@pytest.mark.asyncio
async def test_save_message_creates_new_flushes_when_tx_provided(self):
"""管理员直发路径且事务由应用层控制时仅 flush。"""
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_scalar_result(None))
adapter, _, _, _ = _build_adapter(db)
tx = MagicMock()
# Act
await adapter.saveMessage(_make_save_cmd(), tx=tx)
# Assert
db.add.assert_called_once()
db.flush.assert_awaited_once()
db.commit.assert_not_awaited()
@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) as exc_info:
await adapter.getMessages("abc")
assert exc_info.value.field == "conversation_id"
assert exc_info.value.message == "conversation_id format invalid"
assert exc_info.value.to_dict()["id"] == "abc"
assert "abc" not in exc_info.value.message
@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()
# H-10: mergeConversations 需要全部会话,不传 limit
session_repo.list_by_conversation.assert_awaited_once_with(100)
@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.asyncio
async def test_merge_uses_batch_update_for_session_soft_delete(self):
"""H-9: N 条会话合并仅执行 1 条批量 UPDATE 软删除,不再逐条 delete_by_id。"""
# Arrange —— 5 条源会话
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_id=f"sess-{i}") for i in range(5)]
)
session_repo.delete_by_id = AsyncMock()
# Act
await adapter.mergeConversations(_make_merge_cmd())
# Assert —— 不再逐条软删除
session_repo.delete_by_id.assert_not_called()
# 恰好 1 条针对 channel_sessions 的批量软删除 UPDATE与 N 无关)
session_soft_delete_calls = [
call
for call in db.execute.call_args_list
if "UPDATE channel_sessions" in str(call.args[0])
and "is_deleted" in str(call.args[0])
]
assert len(session_soft_delete_calls) == 1
@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"
# H-10: 仅取首条会话,传 limit=1 避免加载全部会话
session_repo.list_by_conversation.assert_awaited_once_with(100, limit=1)
@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()
# H-10: 仅取首条会话,传 limit=1 避免加载全部会话
session_repo.list_by_conversation.assert_awaited_once_with(100, limit=1)
@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()
conv = _make_conversation_orm()
session = _make_session_row()
db.execute = AsyncMock(
return_value=_make_rows_result(
[
_make_message_row(_make_message_orm(message_id=200), conv, session),
_make_message_row(_make_message_orm(message_id=201), conv, session),
]
)
)
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"
assert items[0].channel_session_id == session.session_id
assert items[0].channel_account_id == session.channel_account_id
assert items[0].peer_id == session.peer_id
assert items[0].conversation_title == conv.title
@pytest.mark.asyncio
async def test_search_returns_empty_list_when_no_matches(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_rows_result([]))
adapter, _, _, _ = _build_adapter(db)
# Act
items = await adapter.searchMessages("nomatch")
# Assert
assert items == ()
@pytest.mark.asyncio
async def test_search_fills_session_context_from_row(self):
# Arrange
db = _make_db()
conv = _make_conversation_orm()
session = _make_session_row(session_id="sess-1", peer_id="peer-1", channel_account_id="acc-1")
db.execute = AsyncMock(
return_value=_make_rows_result([_make_message_row(_make_message_orm(message_id=200), conv, session)])
)
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"
assert items[0].channel_account_id == "acc-1"
assert items[0].peer_id == "peer-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) as exc_info:
await adapter.countMessages("abc")
assert exc_info.value.field == "conversation_id"
assert exc_info.value.message == "conversation_id format invalid"
assert exc_info.value.to_dict()["id"] == "abc"
assert "abc" not in exc_info.value.message
@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
# H-10: 仅取首条会话,传 limit=1 避免加载全部会话
session_repo.list_by_conversation.assert_awaited_once_with(100, limit=1)
@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()
message_orm = _make_message_orm()
conv = _make_conversation_orm()
session = _make_session_row()
db.execute = AsyncMock(return_value=_make_one_or_none_result(_make_message_row(message_orm, conv, session)))
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.getMessageById("200")
# Assert
assert isinstance(result, Message)
assert result.message_id == "200"
assert result.channel_session_id == session.session_id
assert result.channel_account_id == session.channel_account_id
assert result.peer_id == session.peer_id
assert result.conversation_title == conv.title
@pytest.mark.asyncio
async def test_returns_none_when_not_found(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_one_or_none_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_tuple_when_found(self):
# Arrange
db = _make_db()
conv = _make_conversation_orm()
session = _make_session_row()
db.execute = AsyncMock(
return_value=_make_rows_result(
[
_make_message_row(_make_message_orm(message_id=200), conv, session),
_make_message_row(_make_message_orm(message_id=201), conv, session),
]
)
)
adapter, _, _, _ = _build_adapter(db)
# Act
result = await adapter.listMessages(role="assistant")
# Assert
assert isinstance(result, tuple)
assert len(result) == 2
assert all(isinstance(m, Message) for m in result)
assert result[0].channel_session_id == session.session_id
assert result[0].channel_account_id == session.channel_account_id
assert result[0].peer_id == session.peer_id
assert result[0].conversation_title == conv.title
@pytest.mark.asyncio
async def test_returns_empty_tuple_when_no_matches(self):
# Arrange
db = _make_db()
db.execute = AsyncMock(return_value=_make_rows_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_rows_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")
@pytest.mark.unit
class TestConversationAdapterIntegrityError:
"""补充覆盖各查询方法的 IntegrityError→ConflictError 翻译路径。
现有测试仅覆盖了 SQLAlchemyErrorDependencyError 翻译本类补充
IntegrityError并发冲突翻译为 ConflictError 的覆盖确保
``_translate_db_error`` 在以下查询方法上均被验证
getMessages / getMessageById / getMessageByChannelMsgId /
countMessages / listMessages / searchMessages / countSearchMessages
"""
@pytest.mark.asyncio
async def test_getMessages_translates_integrity_error_to_conflict_error(self):
# getMessages 在 session.execute 抛 IntegrityError 时应翻译为 ConflictError
db = _make_db()
db.execute = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter, _, _, _ = _build_adapter(db)
with pytest.raises(ConflictError):
await adapter.getMessages("100")
@pytest.mark.asyncio
async def test_getMessageById_translates_integrity_error_to_conflict_error(self):
# getMessageById 在 session.execute 抛 IntegrityError 时应翻译为 ConflictError
db = _make_db()
db.execute = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter, _, _, _ = _build_adapter(db)
with pytest.raises(ConflictError):
await adapter.getMessageById("200")
@pytest.mark.asyncio
async def test_getMessageByChannelMsgId_translates_integrity_error_to_conflict_error(self):
# getMessageByChannelMsgId 在 session.execute 抛 IntegrityError 时应翻译为 ConflictError
db = _make_db()
db.execute = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter, _, _, _ = _build_adapter(db)
with pytest.raises(ConflictError):
await adapter.getMessageByChannelMsgId("cmid-1")
@pytest.mark.asyncio
async def test_countMessages_translates_integrity_error_to_conflict_error(self):
# countMessages 在 session.execute 抛 IntegrityError 时应翻译为 ConflictError
db = _make_db()
db.execute = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter, _, _, _ = _build_adapter(db)
with pytest.raises(ConflictError):
await adapter.countMessages("100")
@pytest.mark.asyncio
async def test_listMessages_translates_integrity_error_to_conflict_error(self):
# listMessages 在 session.execute 抛 IntegrityError 时应翻译为 ConflictError
db = _make_db()
db.execute = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter, _, _, _ = _build_adapter(db)
with pytest.raises(ConflictError):
await adapter.listMessages(role="assistant")
@pytest.mark.asyncio
async def test_searchMessages_translates_integrity_error_to_conflict_error(self):
# searchMessages 在 session.execute 抛 IntegrityError 时应翻译为 ConflictError
db = _make_db()
db.execute = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter, _, _, _ = _build_adapter(db)
with pytest.raises(ConflictError):
await adapter.searchMessages("hello")
@pytest.mark.asyncio
async def test_countSearchMessages_translates_integrity_error_to_conflict_error(self):
# countSearchMessages 在 session.execute 抛 IntegrityError 时应翻译为 ConflictError
db = _make_db()
db.execute = AsyncMock(side_effect=IntegrityError("stmt", {}, Exception("orig")))
adapter, _, _, _ = _build_adapter(db)
with pytest.raises(ConflictError):
await adapter.countSearchMessages("hello")