本次提交修复了多个测试文件中的问题: 1. 为各管道测试类添加tracer_port属性初始化 2. 修复QQBot客户端关闭异常的错误类型匹配 3. 修正配对接口查询的状态值小写格式 4. 移除ChannelType枚举值的显式.value调用 5. 新增QQ常量的导出项与Instagram测试桩函数 6. 修复API接口返回值解析,正确访问data字段 7. 调整企业微信富文本消息的断言结构 8. 修正微信iLink的消息类型测试用例 9. 替换废弃的datetime.utc相关导入为UTC常量 10. 修复QQBot白名单适配器的返回值结构 11. 调整outbox工具的channel_msg_id处理逻辑 12. 新增多个适配器与服务的测试用例,覆盖异常降级、缓存处理等场景 13. 重构部分QQBot适配器的辅助函数测试,清理冗余代码 14. 修复微信iLink客户端的上传接口调用参数与缓存清理逻辑 15. 修正配置导出接口的敏感字段过滤与返回值结构
1864 lines
63 KiB
Python
1864 lines
63 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.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]:
|
||
"""构造 adapter,patch 三个 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"
|
||
# 不应调用 refresh(latest 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 翻译路径。
|
||
|
||
现有测试仅覆盖了 SQLAlchemyError→DependencyError 翻译,本类补充
|
||
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")
|