From d582a8518cc38a74c02f2b2a66fb8b79198daf82 Mon Sep 17 00:00:00 2001 From: LIghtJUNction Date: Wed, 29 Apr 2026 08:40:45 +0800 Subject: [PATCH] test: dashboard routes unit tests for chat/conversation/open_api --- tests/dashboard/test_chat.py | 350 +++++++++++++++---- tests/dashboard/test_conversation.py | 483 +++++++++++++++++++++++++-- tests/dashboard/test_open_api.py | 391 ++++++++++++++++++++-- 3 files changed, 1087 insertions(+), 137 deletions(-) diff --git a/tests/dashboard/test_chat.py b/tests/dashboard/test_chat.py index 12eba9f19..ba2651790 100644 --- a/tests/dashboard/test_chat.py +++ b/tests/dashboard/test_chat.py @@ -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 diff --git a/tests/dashboard/test_conversation.py b/tests/dashboard/test_conversation.py index d93a0e161..ac808b15d 100644 --- a/tests/dashboard/test_conversation.py +++ b/tests/dashboard/test_conversation.py @@ -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 diff --git a/tests/dashboard/test_open_api.py b/tests/dashboard/test_open_api.py index 3d847fcdd..27ed88e0f 100644 --- a/tests/dashboard/test_open_api.py +++ b/tests/dashboard/test_open_api.py @@ -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()