test: dashboard routes unit tests for chat/conversation/open_api

This commit is contained in:
LIghtJUNction
2026-04-29 08:40:45 +08:00
parent 169a856a3a
commit d582a8518c
3 changed files with 1087 additions and 137 deletions
+278 -72
View File
@@ -1,94 +1,300 @@
"""Import smoke tests for the dashboard chat route module.
"""Unit tests for standalone functions/classes in astrbot/dashboard/routes/chat.py.
Verifies that the main class and its key method signatures from
``chat.py`` can be imported without errors.
Tests cover BotMessageAccumulator, extract_reasoning_from_message_parts,
collect_plain_text_from_message_parts, _sanitize_upload_filename, and
track_conversation. No Quart app fixture required.
"""
import inspect
import pytest
from astrbot.dashboard.routes.chat import (
BotMessageAccumulator,
ChatRoute,
SSE_HEARTBEAT,
_sanitize_upload_filename,
collect_plain_text_from_message_parts,
extract_reasoning_from_message_parts,
track_conversation,
)
class TestChatRouteClass:
def test_class_exists(self):
assert ChatRoute is not None
def test_has_sse_heartbeat(self):
assert SSE_HEARTBEAT == ": heartbeat\n\n"
def test_init_method_signature(self):
sig = inspect.signature(ChatRoute.__init__)
params = list(sig.parameters.keys())
assert "self" in params
assert "context" in params
assert "db" in params
assert "core_lifecycle" in params
def test_chat_method_is_async(self):
assert inspect.iscoroutinefunction(ChatRoute.chat)
def test_chat_method_signature(self):
sig = inspect.signature(ChatRoute.chat)
params = list(sig.parameters.keys())
assert "self" in params
assert "post_data" in params
def test_new_session_method_is_async(self):
assert inspect.iscoroutinefunction(ChatRoute.new_session)
def test_new_session_method_signature(self):
sig = inspect.signature(ChatRoute.new_session)
params = list(sig.parameters.keys())
assert "self" in params
def test_get_session_method_is_async(self):
assert inspect.iscoroutinefunction(ChatRoute.get_session)
def test_get_session_method_signature(self):
sig = inspect.signature(ChatRoute.get_session)
params = list(sig.parameters.keys())
assert "self" in params
# ---------------------------------------------------------------------------
# extract_reasoning_from_message_parts
# ---------------------------------------------------------------------------
class TestBotMessageAccumulatorClass:
def test_class_exists(self):
assert BotMessageAccumulator is not None
class TestExtractReasoningFromMessageParts:
def test_empty_list_returns_empty(self):
assert extract_reasoning_from_message_parts([]) == ""
def test_has_content_method(self):
assert callable(BotMessageAccumulator.has_content)
def test_no_think_parts_returns_empty(self):
parts = [{"type": "plain", "text": "hello"}]
assert extract_reasoning_from_message_parts(parts) == ""
def test_add_plain_method(self):
sig = inspect.signature(BotMessageAccumulator.add_plain)
params = list(sig.parameters.keys())
assert "self" in params
assert "result_text" in params
assert "chain_type" in params
assert "streaming" in params
def test_single_think_part(self):
parts = [{"type": "think", "think": "deep reasoning"}]
assert extract_reasoning_from_message_parts(parts) == "deep reasoning"
def test_build_message_parts_method(self):
sig = inspect.signature(BotMessageAccumulator.build_message_parts)
params = list(sig.parameters.keys())
assert "self" in params
assert "include_pending_tool_calls" in params
def test_multiple_think_parts_concatenated(self):
parts = [
{"type": "think", "think": "first "},
{"type": "plain", "text": "skip"},
{"type": "think", "think": "second"},
]
assert extract_reasoning_from_message_parts(parts) == "first second"
def test_non_string_think_value_skipped(self):
parts = [
{"type": "think", "think": 42},
{"type": "think", "think": "valid"},
]
assert extract_reasoning_from_message_parts(parts) == "valid"
class TestStandaloneFunctions:
def test_track_conversation_is_async_gen(self):
assert inspect.isasyncgenfunction(track_conversation)
# ---------------------------------------------------------------------------
# collect_plain_text_from_message_parts
# ---------------------------------------------------------------------------
def test_collect_plain_text_from_message_parts_is_callable(self):
assert callable(collect_plain_text_from_message_parts)
def test_extract_reasoning_from_message_parts_is_callable(self):
assert callable(extract_reasoning_from_message_parts)
class TestCollectPlainTextFromMessageParts:
def test_empty_list_returns_empty(self):
assert collect_plain_text_from_message_parts([]) == ""
def test_sanitize_upload_filename_is_callable(self):
assert callable(_sanitize_upload_filename)
def test_no_plain_parts_returns_empty(self):
parts = [{"type": "think", "think": "hidden"}]
assert collect_plain_text_from_message_parts(parts) == ""
def test_single_plain_part(self):
parts = [{"type": "plain", "text": "hello world"}]
assert collect_plain_text_from_message_parts(parts) == "hello world"
def test_multiple_plain_parts_concatenated(self):
parts = [
{"type": "plain", "text": "hello "},
{"type": "think", "think": "hidden"},
{"type": "plain", "text": "world"},
]
assert collect_plain_text_from_message_parts(parts) == "hello world"
def test_non_string_text_field_skipped(self):
parts = [{"type": "plain", "text": 99}]
assert collect_plain_text_from_message_parts(parts) == ""
# ---------------------------------------------------------------------------
# _sanitize_upload_filename
# ---------------------------------------------------------------------------
class TestSanitizeUploadFilename:
def test_empty_returns_random_hex(self):
result = _sanitize_upload_filename("")
assert isinstance(result, str)
assert len(result) == 16
assert all(c in "0123456789abcdef" for c in result)
def test_null_bytes_removed(self):
assert _sanitize_upload_filename("file\x00name.txt") == "filename.txt"
def test_path_traversal_stripped(self):
assert _sanitize_upload_filename("../../etc/passwd") == "passwd"
def test_windows_fakepath_stripped(self):
assert _sanitize_upload_filename("C:\\fakepath\\doc.pdf") == "doc.pdf"
def test_windows_fakepath_lowercase_drive(self):
assert _sanitize_upload_filename("c:\\fakepath\\photo.png") == "photo.png"
def test_backslash_converted(self):
assert _sanitize_upload_filename("folder\\sub\\file.txt") == "file.txt"
def test_normal_filename_preserved(self):
assert _sanitize_upload_filename("report.pdf") == "report.pdf"
def test_dot_returns_random(self):
result = _sanitize_upload_filename(".")
assert isinstance(result, str) and len(result) == 16
def test_trailing_slash_stripped(self):
assert _sanitize_upload_filename("dir/") == "dir"
# ---------------------------------------------------------------------------
# track_conversation (async context manager)
# ---------------------------------------------------------------------------
class TestTrackConversation:
@pytest.mark.asyncio
async def test_adds_and_removes_key(self):
convs: dict = {}
async with track_conversation(convs, "test-id"):
assert convs.get("test-id") is True
assert "test-id" not in convs
@pytest.mark.asyncio
async def test_cleans_up_on_exception(self):
convs: dict = {}
with pytest.raises(RuntimeError):
async with track_conversation(convs, "test-id"):
raise RuntimeError("boom")
assert "test-id" not in convs
# ---------------------------------------------------------------------------
# BotMessageAccumulator
# ---------------------------------------------------------------------------
class TestBotMessageAccumulatorInit:
def test_initial_state(self):
acc = BotMessageAccumulator()
assert acc.parts == []
assert acc.pending_text == ""
assert acc.pending_tool_calls == {}
assert acc.has_content() is False
def test_has_content_true_with_pending_text(self):
acc = BotMessageAccumulator()
acc.pending_text = "x"
assert acc.has_content() is True
def test_has_content_true_with_parts(self):
acc = BotMessageAccumulator()
acc.parts.append({"type": "plain", "text": "x"})
assert acc.has_content() is True
def test_has_content_true_with_pending_tool_calls(self):
acc = BotMessageAccumulator()
acc.pending_tool_calls["c1"] = {"id": "c1"}
assert acc.has_content() is True
class TestBotMessageAccumulatorAddPlain:
def test_streaming_appends_to_pending(self):
acc = BotMessageAccumulator()
acc.add_plain("Hello ", chain_type=None, streaming=True)
acc.add_plain("World", chain_type=None, streaming=True)
assert acc.pending_text == "Hello World"
def test_non_streaming_replaces_pending(self):
acc = BotMessageAccumulator()
acc.add_plain("Hello", chain_type=None, streaming=True)
acc.add_plain("World", chain_type=None, streaming=False)
assert acc.pending_text == "World"
def test_reasoning_chain_flushes_and_stores_think(self):
acc = BotMessageAccumulator()
acc.add_plain("visible", chain_type=None, streaming=True)
acc.add_plain("hidden", chain_type="reasoning", streaming=False)
assert acc.pending_text == ""
assert acc.reasoning_text() == "hidden"
assert acc.plain_text() == "visible"
def test_tool_call_stores_pending_call(self):
acc = BotMessageAccumulator()
acc.add_plain(
'{"id": "c1", "name": "search"}',
chain_type="tool_call",
streaming=False,
)
assert "c1" in acc.pending_tool_calls
def test_tool_call_result_creates_part(self):
acc = BotMessageAccumulator()
acc._store_tool_call('{"id": "c1", "name": "search"}')
acc._store_tool_call_result(
'{"id": "c1", "result": "data", "ts": 1}'
)
assert "c1" not in acc.pending_tool_calls
assert len(acc.parts) == 1
assert acc.parts[0]["tool_calls"][0]["result"] == "data"
def test_tool_call_invalid_json_ignored(self):
acc = BotMessageAccumulator()
acc.add_plain("not-json", chain_type="tool_call", streaming=False)
assert acc.pending_tool_calls == {}
class TestBotMessageAccumulatorAddAttachment:
def test_none_ignored(self):
acc = BotMessageAccumulator()
acc.add_attachment(None)
assert acc.parts == []
def test_valid_part_appended(self):
acc = BotMessageAccumulator()
acc.add_attachment({"type": "image", "url": "test.jpg"})
assert len(acc.parts) == 1
assert acc.parts[0]["url"] == "test.jpg"
def test_flushes_pending_text_before_append(self):
acc = BotMessageAccumulator()
acc.add_plain("text", chain_type=None, streaming=True)
acc.add_attachment({"type": "image"})
assert acc.pending_text == ""
assert acc.parts[0]["type"] == "plain"
class TestBotMessageAccumulatorBuildMessageParts:
def test_flushes_pending_text(self):
acc = BotMessageAccumulator()
acc.add_plain("hello", chain_type=None, streaming=True)
parts = acc.build_message_parts()
assert len(parts) == 1
assert parts[0] == {"type": "plain", "text": "hello"}
def test_includes_pending_tool_calls_when_requested(self):
acc = BotMessageAccumulator()
acc.pending_tool_calls["c1"] = {"id": "c1", "name": "search"}
parts = acc.build_message_parts(include_pending_tool_calls=True)
assert len(parts) == 1
assert parts[0]["type"] == "tool_call"
def test_skips_pending_tool_calls_by_default(self):
acc = BotMessageAccumulator()
acc.pending_tool_calls["c1"] = {"id": "c1"}
parts = acc.build_message_parts(include_pending_tool_calls=False)
assert parts == []
class TestBotMessageAccumulatorInternalFlushAndThink:
def test_flush_creates_new_plain_part(self):
acc = BotMessageAccumulator()
acc.pending_text = "hello"
acc._flush_pending_text()
assert acc.parts == [{"type": "plain", "text": "hello"}]
assert acc.pending_text == ""
def test_flush_appends_to_existing_plain_part(self):
acc = BotMessageAccumulator()
acc.parts.append({"type": "plain", "text": "Hello"})
acc.pending_text = " World"
acc._flush_pending_text()
assert acc.parts == [{"type": "plain", "text": "Hello World"}]
def test_append_think_creates_new_part(self):
acc = BotMessageAccumulator()
acc._append_think_part("step 1")
assert acc.parts == [{"type": "think", "think": "step 1"}]
def test_append_think_appends_to_existing(self):
acc = BotMessageAccumulator()
acc._append_think_part("step 1")
acc._append_think_part(" step 2")
assert acc.parts == [{"type": "think", "think": "step 1 step 2"}]
def test_append_think_empty_ignored(self):
acc = BotMessageAccumulator()
acc._append_think_part("")
assert acc.parts == []
class TestBotMessageAccumulatorParseJsonObject:
def test_valid_dict(self):
assert BotMessageAccumulator._parse_json_object(
'{"a": 1}'
) == {"a": 1}
def test_valid_list_returns_none(self):
assert BotMessageAccumulator._parse_json_object("[1,2]") is None
def test_invalid_json_returns_none(self):
assert BotMessageAccumulator._parse_json_object("not json") is None
+451 -32
View File
@@ -1,46 +1,465 @@
"""Import smoke tests for the dashboard conversation route module.
"""Mock-based unit tests for ConversationRoute in conversation.py.
Verifies that the main class and its key method signatures from
``conversation.py`` can be imported without errors.
All tests mock ``request`` and ``g`` at the module level so no Quart app
fixture is required. The Route base class requires a mock ``app`` on the
context object; register_routes() calls add_url_rule through the mock.
"""
import inspect
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from astrbot.dashboard.routes.conversation import ConversationRoute
from astrbot.dashboard.routes.route import RouteContext
# Sentinel to distinguish "no conv_mgr argument" from "conv_mgr=None".
_UNSET = object()
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class TestConversationRouteClass:
def test_class_exists(self):
assert ConversationRoute is not None
def _make_route_context():
"""Return a RouteContext with a mock app and config."""
return RouteContext(
app=MagicMock(),
config=MagicMock(),
)
def test_init_method_signature(self):
sig = inspect.signature(ConversationRoute.__init__)
params = list(sig.parameters.keys())
assert "self" in params
assert "context" in params
assert "db_helper" in params
assert "core_lifecycle" in params
def test_list_conversations_method_is_async(self):
assert inspect.iscoroutinefunction(ConversationRoute.list_conversations)
def _make_mock_args_getter(values: dict):
"""Build a callable that mimics request.args.get with optional type
coercion (the pattern used by Quart's MultiDict)."""
def test_list_conversations_method_signature(self):
sig = inspect.signature(ConversationRoute.list_conversations)
params = list(sig.parameters.keys())
assert "self" in params
def get(key, default=None, type=None): # noqa: A002
val = values.get(key)
if val is not None and type is not None:
return type(val)
return val if val is not None else default
def test_del_conv_method_is_async(self):
assert inspect.iscoroutinefunction(ConversationRoute.del_conv)
return get
def test_del_conv_method_signature(self):
sig = inspect.signature(ConversationRoute.del_conv)
params = list(sig.parameters.keys())
assert "self" in params
def test_export_conversations_method_is_async(self):
assert inspect.iscoroutinefunction(ConversationRoute.export_conversations)
def _build_route(conv_mgr=_UNSET):
"""Return a (ConversationRoute, mock_db, mock_core_lifecycle) tuple.
def test_export_conversations_method_signature(self):
sig = inspect.signature(ConversationRoute.export_conversations)
params = list(sig.parameters.keys())
assert "self" in params
All dependencies are MagicMock based. Pass ``conv_mgr=None`` to simulate
an unavailable conversation manager; the default creates a fresh mock.
"""
ctx = _make_route_context()
mock_db = MagicMock()
mock_core = MagicMock()
if conv_mgr is _UNSET:
mock_core.conversation_manager = MagicMock()
else:
mock_core.conversation_manager = conv_mgr
route = ConversationRoute(
context=ctx,
db_helper=mock_db,
core_lifecycle=mock_core,
)
return route, mock_db, mock_core
# ---------------------------------------------------------------------------
# list_conversations
# ---------------------------------------------------------------------------
class TestListConversations:
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.g")
@pytest.mark.asyncio
async def test_default_pagination(self, mock_g, mock_request):
"""Default page/page_size when no query params are supplied."""
mock_request.args.get.side_effect = _make_mock_args_getter({})
mock_g.get.return_value = "testuser"
route, _db, _core = _build_route()
result = await route.list_conversations()
assert result["status"] == "ok"
pag = result["data"]["pagination"]
assert pag["page"] == 1
assert pag["page_size"] == 20
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.g")
@pytest.mark.asyncio
async def test_custom_pagination(self, mock_g, mock_request):
"""Custom page/page_size passed as query strings."""
mock_request.args.get.side_effect = _make_mock_args_getter(
{"page": "3", "page_size": "50"}
)
mock_g.get.return_value = "testuser"
route, _db, _core = _build_route()
result = await route.list_conversations()
assert result["status"] == "ok"
pag = result["data"]["pagination"]
assert pag["page"] == 3
assert pag["page_size"] == 50
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.g")
@pytest.mark.asyncio
async def test_page_size_clamped_to_max_100(self, mock_g, mock_request):
"""page_size beyond 100 is clamped."""
mock_request.args.get.side_effect = _make_mock_args_getter(
{"page_size": "999"}
)
mock_g.get.return_value = "testuser"
route, _db, _core = _build_route()
result = await route.list_conversations()
assert result["status"] == "ok"
assert result["data"]["pagination"]["page_size"] == 100
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.g")
@pytest.mark.asyncio
async def test_page_clamped_to_minimum_1(self, mock_g, mock_request):
"""page below 1 is raised to 1."""
mock_request.args.get.side_effect = _make_mock_args_getter(
{"page": "0"}
)
mock_g.get.return_value = "testuser"
route, _db, _core = _build_route()
result = await route.list_conversations()
assert result["status"] == "ok"
assert result["data"]["pagination"]["page"] == 1
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.g")
@pytest.mark.asyncio
async def test_filter_params_passed_to_manager(self, mock_g, mock_request):
"""Filter query params are split and forwarded."""
mock_request.args.get.side_effect = _make_mock_args_getter(
{
"platforms": "webchat,discord",
"message_types": "FriendMessage",
"search": "hello",
}
)
mock_g.get.return_value = "testuser"
route, _db, core = _build_route()
await route.list_conversations()
core.conversation_manager.get_filtered_conversations.assert_awaited_once_with(
page=1,
page_size=20,
platforms=["webchat", "discord"],
message_types=["FriendMessage"],
search_query="hello",
exclude_ids=[],
exclude_platforms=[],
)
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.g")
@pytest.mark.asyncio
async def test_error_when_conv_mgr_unavailable(self, mock_g, mock_request):
"""Returns error when conversation_manager is falsy."""
mock_g.get.return_value = "testuser"
route, _db, _core = _build_route(conv_mgr=None)
result = await route.list_conversations()
assert result["status"] == "error"
assert "not available" in (result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.g")
@pytest.mark.asyncio
async def test_db_query_exception_wrapped(self, mock_g, mock_request):
"""Exception from get_filtered_conversations is caught and returned as error."""
mock_g.get.return_value = "testuser"
route, _db, core = _build_route()
core.conversation_manager.get_filtered_conversations.side_effect = (
RuntimeError("db connection lost")
)
result = await route.list_conversations()
assert result["status"] == "error"
assert "数据库查询出错" in str(result.get("message") or "")
# ---------------------------------------------------------------------------
# get_conv_detail
# ---------------------------------------------------------------------------
class TestGetConvDetail:
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_missing_params_returns_error(self, mock_request):
mock_request.get_json = AsyncMock(return_value={})
route, _db, _core = _build_route()
result = await route.get_conv_detail()
assert result["status"] == "error"
assert "缺少必要参数" in str(result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_conversation_not_found(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={"user_id": "user1", "cid": "cid1"}
)
route, _db, core = _build_route()
core.conversation_manager.get_conversation = AsyncMock(return_value=None)
result = await route.get_conv_detail()
assert result["status"] == "error"
assert "不存在" in str(result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_success_with_valid_params(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={"user_id": "user1", "cid": "cid1"}
)
route, _db, core = _build_route()
mock_conv = MagicMock()
mock_conv.title = "test-title"
mock_conv.persona_id = "p1"
mock_conv.history = "[]"
mock_conv.created_at = "2024-01-01"
mock_conv.updated_at = "2024-01-02"
core.conversation_manager.get_conversation = AsyncMock(return_value=mock_conv)
result = await route.get_conv_detail()
assert result["status"] == "ok"
assert result["data"]["title"] == "test-title"
assert result["data"]["persona_id"] == "p1"
# ---------------------------------------------------------------------------
# upd_conv
# ---------------------------------------------------------------------------
class TestUpdConv:
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_missing_params_returns_error(self, mock_request):
mock_request.get_json = AsyncMock(return_value={})
route, _db, _core = _build_route()
result = await route.upd_conv()
assert result["status"] == "error"
assert "缺少必要参数" in str(result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_not_found(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={"user_id": "user1", "cid": "cid1"}
)
route, _db, core = _build_route()
core.conversation_manager.get_conversation = AsyncMock(return_value=None)
result = await route.upd_conv()
assert result["status"] == "error"
assert "不存在" in str(result.get("message") or "")
# ---------------------------------------------------------------------------
# del_conv (single + batch)
# ---------------------------------------------------------------------------
class TestDelConv:
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_single_missing_params(self, mock_request):
mock_request.get_json = AsyncMock(return_value={})
route, _db, _core = _build_route()
result = await route.del_conv()
assert result["status"] == "error"
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_batch_delete(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={
"conversations": [
{"user_id": "u1", "cid": "c1"},
{"user_id": "u2", "cid": "c2"},
]
}
)
route, _db, core = _build_route()
core.conversation_manager.delete_conversation = AsyncMock()
result = await route.del_conv()
assert result["status"] == "ok"
assert result["data"]["deleted_count"] == 2
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_batch_with_failures(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={
"conversations": [
{"user_id": "u1", "cid": "c1"},
{"user_id": "", "cid": ""}, # missing params
]
}
)
route, _db, core = _build_route()
core.conversation_manager.delete_conversation = AsyncMock()
result = await route.del_conv()
assert result["status"] == "ok"
assert result["data"]["deleted_count"] == 1
assert result["data"]["failed_count"] == 1
# ---------------------------------------------------------------------------
# update_history
# ---------------------------------------------------------------------------
class TestUpdateHistory:
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_missing_user_id_or_cid(self, mock_request):
mock_request.get_json = AsyncMock(return_value={"history": []})
route, _db, _core = _build_route()
result = await route.update_history()
assert result["status"] == "error"
assert "缺少必要参数" in str(result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_missing_history(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={"user_id": "u1", "cid": "c1"}
)
route, _db, _core = _build_route()
result = await route.update_history()
assert result["status"] == "error"
assert "缺少必要参数" in str(result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_invalid_json_string_returns_error(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={
"user_id": "u1",
"cid": "c1",
"history": "not-json",
}
)
route, _db, _core = _build_route()
result = await route.update_history()
assert result["status"] == "error"
assert "有效的 JSON" in str(result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_valid_list_history(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={
"user_id": "u1",
"cid": "c1",
"history": [{"role": "user", "content": "hello"}],
}
)
route, _db, core = _build_route()
core.conversation_manager.get_conversation = AsyncMock(
return_value=MagicMock()
)
core.conversation_manager.update_conversation = AsyncMock()
result = await route.update_history()
assert result["status"] == "ok"
# ---------------------------------------------------------------------------
# export_conversations
# ---------------------------------------------------------------------------
class TestExportConversations:
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_empty_list_returns_error(self, mock_request):
mock_request.get_json = AsyncMock(
return_value={"conversations": []}
)
route, _db, _core = _build_route()
result = await route.export_conversations()
assert result["status"] == "error"
assert "不能为空" in str(result.get("message") or "")
@patch("astrbot.dashboard.routes.conversation.request")
@patch("astrbot.dashboard.routes.conversation.send_file")
@pytest.mark.asyncio
async def test_export_success(self, mock_send_file, mock_request):
"""A valid export request returns a file response via send_file."""
mock_send_file.return_value = {"status": "ok", "_mock_file": True}
mock_request.get_json = AsyncMock(
return_value={
"conversations": [
{"user_id": "u1", "cid": "c1"},
]
}
)
route, _db, core = _build_route()
mock_conv = MagicMock()
mock_conv.history = "[]"
mock_conv.title = "t1"
mock_conv.persona_id = None
mock_conv.platform_id = "webchat"
mock_conv.created_at = "2024-01-01"
mock_conv.updated_at = "2024-01-02"
core.conversation_manager.get_conversation = AsyncMock(return_value=mock_conv)
result = await route.export_conversations()
assert result["_mock_file"] is True
mock_send_file.assert_awaited_once()
@patch("astrbot.dashboard.routes.conversation.request")
@pytest.mark.asyncio
async def test_export_skips_items_missing_params(self, mock_request):
"""Items without user_id/cid are reported as failures."""
mock_request.get_json = AsyncMock(
return_value={
"conversations": [
{"user_id": "u1", "cid": "c1"},
{"user_id": "", "cid": ""}, # missing
]
}
)
route, _db, core = _build_route()
core.conversation_manager.get_conversation = AsyncMock(
side_effect=[
MagicMock(
history="[]",
title="t",
persona_id=None,
platform_id="w",
created_at="",
updated_at="",
),
]
)
with patch("astrbot.dashboard.routes.conversation.send_file") as mock_sf:
mock_sf.return_value = {"_mock": True}
result = await route.export_conversations()
assert result["_mock"] is True
+358 -33
View File
@@ -1,47 +1,372 @@
"""Import smoke tests for the dashboard OpenAPI route module.
"""Mock-based unit tests for standalone methods in OpenApiRoute (open_api.py).
Verifies that the main class and its key method signatures from
``open_api.py`` can be imported without errors.
Tests cover _resolve_open_username, _get_chat_config_list,
_resolve_chat_config_id, _extract_ws_api_key, _ensure_chat_session, and
_ensure_runtime_ready. No Quart app fixture required; all external
dependencies (request, websocket, db, core_lifecycle) are mocked.
"""
import inspect
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from astrbot.dashboard.routes.open_api import OpenApiRoute
from astrbot.dashboard.routes.route import RouteContext
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class TestOpenApiRouteClass:
def test_class_exists(self):
assert OpenApiRoute is not None
def _make_route_context():
"""Return a RouteContext with a mock app and config."""
return RouteContext(app=MagicMock(), config=MagicMock())
def test_init_method_signature(self):
sig = inspect.signature(OpenApiRoute.__init__)
params = list(sig.parameters.keys())
assert "self" in params
assert "context" in params
assert "db" in params
assert "core_lifecycle" in params
assert "chat_route" in params
def test_chat_send_method_is_async(self):
assert inspect.iscoroutinefunction(OpenApiRoute.chat_send)
def _build_openapi_route(core_lifecycle=None, db=None, chat_route=None):
"""Return an (OpenApiRoute, mock_db, mock_core_lifecycle, mock_chat_route)
tuple, filling in MagicMock defaults for omitted dependencies.
"""
ctx = _make_route_context()
mock_core = core_lifecycle or MagicMock()
mock_db = db or MagicMock()
mock_chat = chat_route or MagicMock()
route = OpenApiRoute(
context=ctx,
db=mock_db,
core_lifecycle=mock_core,
chat_route=mock_chat,
)
return route, mock_db, mock_core, mock_chat
def test_chat_send_method_signature(self):
sig = inspect.signature(OpenApiRoute.chat_send)
params = list(sig.parameters.keys())
assert "self" in params
def test_get_chat_configs_method_is_async(self):
assert inspect.iscoroutinefunction(OpenApiRoute.get_chat_configs)
# ---------------------------------------------------------------------------
# _resolve_open_username (static)
# ---------------------------------------------------------------------------
def test_get_chat_configs_method_signature(self):
sig = inspect.signature(OpenApiRoute.get_chat_configs)
params = list(sig.parameters.keys())
assert "self" in params
def test_send_message_method_is_async(self):
assert inspect.iscoroutinefunction(OpenApiRoute.send_message)
class TestResolveOpenUsername:
def test_none_returns_error(self):
username, err = OpenApiRoute._resolve_open_username(None)
assert username is None
assert err == "Missing key: username"
def test_send_message_method_signature(self):
sig = inspect.signature(OpenApiRoute.send_message)
params = list(sig.parameters.keys())
assert "self" in params
def test_empty_string_returns_error(self):
username, err = OpenApiRoute._resolve_open_username("")
assert username is None
assert err == "username is empty"
def test_whitespace_only_returns_error(self):
username, err = OpenApiRoute._resolve_open_username(" ")
assert username is None
assert err == "username is empty"
def test_valid_username(self):
username, err = OpenApiRoute._resolve_open_username("alice")
assert username == "alice"
assert err is None
def test_valid_username_trimmed(self):
username, err = OpenApiRoute._resolve_open_username(" bob ")
assert username == "bob"
assert err is None
# ---------------------------------------------------------------------------
# _get_chat_config_list
# ---------------------------------------------------------------------------
class TestGetChatConfigList:
def test_returns_empty_when_mgr_is_none(self):
core = MagicMock()
core.astrbot_config_mgr = None
route, *_ = _build_openapi_route(core_lifecycle=core)
assert route._get_chat_config_list() == []
def test_returns_config_list(self):
core = MagicMock()
core.astrbot_config_mgr.get_conf_list.return_value = [
{"id": "default", "name": "Default Config", "path": "/cfg1"},
{"id": "cfg2", "name": "My Config", "path": "/cfg2"},
]
route, *_ = _build_openapi_route(core_lifecycle=core)
result = route._get_chat_config_list()
assert len(result) == 2
assert result[0] == {
"id": "default",
"name": "Default Config",
"path": "/cfg1",
"is_default": True,
}
assert result[1] == {
"id": "cfg2",
"name": "My Config",
"path": "/cfg2",
"is_default": False,
}
# ---------------------------------------------------------------------------
# _resolve_chat_config_id
# ---------------------------------------------------------------------------
class TestResolveChatConfigId:
def test_no_config_id_or_name_returns_none_none(self):
route, *_ = _build_openapi_route()
config_id, err = route._resolve_chat_config_id({})
assert config_id is None
assert err is None
def test_config_id_found(self):
core = MagicMock()
core.astrbot_config_mgr.get_conf_list.return_value = [
{"id": "cfg1", "name": "Cfg 1", "path": "/p1"},
]
route, *_ = _build_openapi_route(core_lifecycle=core)
config_id, err = route._resolve_chat_config_id({"config_id": "cfg1"})
assert config_id == "cfg1"
assert err is None
def test_config_id_not_found(self):
core = MagicMock()
core.astrbot_config_mgr.get_conf_list.return_value = []
route, *_ = _build_openapi_route(core_lifecycle=core)
config_id, err = route._resolve_chat_config_id({"config_id": "missing"})
assert config_id is None
assert "not found" in (err or "")
def test_config_name_found(self):
core = MagicMock()
core.astrbot_config_mgr.get_conf_list.return_value = [
{"id": "c1", "name": "My Config", "path": "/p1"},
]
route, *_ = _build_openapi_route(core_lifecycle=core)
config_id, err = route._resolve_chat_config_id(
{"config_name": "My Config"}
)
assert config_id == "c1"
assert err is None
def test_config_name_not_found(self):
core = MagicMock()
core.astrbot_config_mgr.get_conf_list.return_value = []
route, *_ = _build_openapi_route(core_lifecycle=core)
config_id, err = route._resolve_chat_config_id(
{"config_name": "Nope"}
)
assert config_id is None
assert "not found" in (err or "")
def test_config_name_ambiguous(self):
core = MagicMock()
core.astrbot_config_mgr.get_conf_list.return_value = [
{"id": "c1", "name": "Same Name", "path": "/p1"},
{"id": "c2", "name": "Same Name", "path": "/p2"},
]
route, *_ = _build_openapi_route(core_lifecycle=core)
config_id, err = route._resolve_chat_config_id(
{"config_name": "Same Name"}
)
assert config_id is None
assert "ambiguous" in (err or "")
def test_config_name_empty_after_strip_returns_none(self):
"""When config_name is present but only whitespace."""
core = MagicMock()
core.astrbot_config_mgr.get_conf_list.return_value = []
route, *_ = _build_openapi_route(core_lifecycle=core)
config_id, err = route._resolve_chat_config_id(
{"config_name": " "}
)
# config_name is stripped to "" -> the method returns (None, "config_name is empty")
assert config_id is None
assert err == "config_name is empty"
# ---------------------------------------------------------------------------
# _extract_ws_api_key (static, uses ``websocket`` from quart)
# ---------------------------------------------------------------------------
class TestExtractWsApiKey:
@patch("astrbot.dashboard.routes.open_api.websocket")
def test_from_api_key_arg(self, mock_ws):
mock_ws.args.get.side_effect = lambda k, default=None: (
"my-api-key" if k == "api_key" else default
)
result = OpenApiRoute._extract_ws_api_key()
assert result == "my-api-key"
@patch("astrbot.dashboard.routes.open_api.websocket")
def test_from_key_arg(self, mock_ws):
mock_ws.args.get.side_effect = lambda k, default=None: (
"key-from-arg" if k == "key" else default
)
result = OpenApiRoute._extract_ws_api_key()
assert result == "key-from-arg"
@patch("astrbot.dashboard.routes.open_api.websocket")
def test_from_x_api_key_header(self, mock_ws):
mock_ws.args.get.return_value = None
mock_ws.headers.get.side_effect = lambda k, default=None: (
"header-key" if k == "X-API-Key" else default
)
result = OpenApiRoute._extract_ws_api_key()
assert result == "header-key"
@patch("astrbot.dashboard.routes.open_api.websocket")
def test_from_bearer_auth(self, mock_ws):
mock_ws.args.get.return_value = None
mock_ws.headers.get.side_effect = lambda k, default=None: (
"Bearer token123" if k == "Authorization" else default
)
result = OpenApiRoute._extract_ws_api_key()
assert result == "token123"
@patch("astrbot.dashboard.routes.open_api.websocket")
def test_from_apikey_auth(self, mock_ws):
mock_ws.args.get.return_value = None
mock_ws.headers.get.side_effect = lambda k, default=None: (
"ApiKey api-key-value" if k == "Authorization" else default
)
result = OpenApiRoute._extract_ws_api_key()
assert result == "api-key-value"
@patch("astrbot.dashboard.routes.open_api.websocket")
def test_no_key_found_returns_none(self, mock_ws):
mock_ws.args.get.return_value = None
mock_ws.headers.get.return_value = ""
result = OpenApiRoute._extract_ws_api_key()
assert result is None
# ---------------------------------------------------------------------------
# _ensure_chat_session
# ---------------------------------------------------------------------------
class TestEnsureChatSession:
@pytest.mark.asyncio
async def test_session_exists_and_belongs_to_user(self):
db = MagicMock()
db.get_platform_session_by_id = AsyncMock(return_value=MagicMock(creator="alice"))
route, *_ = _build_openapi_route(db=db)
err = await route._ensure_chat_session("alice", "sid1")
assert err is None
@pytest.mark.asyncio
async def test_session_exists_but_wrong_user(self):
db = MagicMock()
db.get_platform_session_by_id = AsyncMock(
return_value=MagicMock(creator="bob")
)
route, *_ = _build_openapi_route(db=db)
err = await route._ensure_chat_session("alice", "sid1")
assert err is not None
assert "belongs to another" in err
@pytest.mark.asyncio
async def test_session_does_not_exist_created(self):
db = MagicMock()
db.get_platform_session_by_id = AsyncMock(return_value=None)
db.create_platform_session = AsyncMock()
route, *_ = _build_openapi_route(db=db)
err = await route._ensure_chat_session("alice", "new-sid")
assert err is None
db.create_platform_session.assert_awaited_once()
@pytest.mark.asyncio
async def test_create_race_recovered(self):
"""If create_platform_session raises but a concurrent creation
succeeded for the same user, we recover gracefully."""
db = MagicMock()
# First get returns None, create raises
call_count = 0
async def get_session(*args, **kwargs): # noqa: ARG001
nonlocal call_count
call_count += 1
if call_count == 1:
return None # first call: not found
return MagicMock(creator="alice") # second call: found (race)
db.get_platform_session_by_id.side_effect = get_session
db.create_platform_session = AsyncMock(side_effect=Exception("duplicate"))
route, *_ = _build_openapi_route(db=db)
err = await route._ensure_chat_session("alice", "race-sid")
assert err is None # recovered
# ---------------------------------------------------------------------------
# _send_chat_ws_error (uses ``websocket`` from quart)
# ---------------------------------------------------------------------------
class TestSendChatWsError:
@patch("astrbot.dashboard.routes.open_api.websocket")
@pytest.mark.asyncio
async def test_sends_error_json(self, mock_ws):
route, *_ = _build_openapi_route()
await route._send_chat_ws_error("Something broke", "ERR_CODE")
mock_ws.send_json.assert_awaited_once_with(
{"type": "error", "code": "ERR_CODE", "data": "Something broke"}
)
# ---------------------------------------------------------------------------
# _update_session_config_route
# ---------------------------------------------------------------------------
class TestUpdateSessionConfigRoute:
@pytest.mark.asyncio
async def test_no_config_id_returns_none(self):
route, *_ = _build_openapi_route()
err = await route._update_session_config_route(
username="alice",
session_id="sid1",
config_id=None,
)
assert err is None
@pytest.mark.asyncio
async def test_router_not_available(self):
core = MagicMock()
core.umop_config_router = None
route, *_ = _build_openapi_route(core_lifecycle=core)
err = await route._update_session_config_route(
username="alice",
session_id="sid1",
config_id="cfg1",
)
assert err is not None
assert "not available" in err
@pytest.mark.asyncio
async def test_delete_route_for_default(self):
core = MagicMock()
core.umop_config_router.delete_route = AsyncMock()
route, *_ = _build_openapi_route(core_lifecycle=core)
err = await route._update_session_config_route(
username="alice",
session_id="sid1",
config_id="default",
)
assert err is None
core.umop_config_router.delete_route.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_route_for_custom_config(self):
core = MagicMock()
core.umop_config_router.update_route = AsyncMock()
route, *_ = _build_openapi_route(core_lifecycle=core)
err = await route._update_session_config_route(
username="alice",
session_id="sid1",
config_id="mycfg",
)
assert err is None
core.umop_config_router.update_route.assert_awaited_once()