mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
test: dashboard routes unit tests for chat/conversation/open_api
This commit is contained in:
+278
-72
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user