本次提交修复了多个测试文件中的问题: 1. 将 ChannelType 枚举调用改为字符串实例化方式 2. 修正了日志断言、异步mock使用、配置参数等多处测试细节 3. 新增了会话聚合根、跨渠道关联策略等单元测试用例 4. 修复了路由测试中的路径方法错误与断言逻辑 5. 调整了依赖导入与测试夹具的兼容性 6. 统一了重试回退调度的列表/元组使用规范
1711 lines
52 KiB
Python
1711 lines
52 KiB
Python
"""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.dashboard import MessageStats
|
||
from yuxi.channels.contract.dtos.conversation import (
|
||
AppendOperationHistoryCmd,
|
||
AssociateConversationCmd,
|
||
MergeConversationCmd,
|
||
MergeResult,
|
||
ResolveConversationCmd,
|
||
SaveMessageCmd,
|
||
UpdateMessageChannelStatusCmd,
|
||
)
|
||
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]:
|
||
"""构造 adapter,patch 三个 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")
|