From 555ac39aee59e5e9ce04aff9a2b0214ea4abfd2b Mon Sep 17 00:00:00 2001 From: KazusaUwU <204833675+JosephTian876@users.noreply.github.com> Date: Sat, 1 Aug 2026 01:40:19 +0800 Subject: [PATCH 1/2] fix(kb): resolve kb_id path parameter in dashboard retrieve endpoint MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `POST /api/v1/knowledge-bases/{kb_id}/retrieve` identifies the knowledge base with the `kb_id` path parameter, and the route already forwards it as `service.retrieve({"kb_id": kb_id, **body})`. However `retrieve()` only ever read `kb_names`, which that route never supplies, so every request from the WebUI retrieval panel failed with "缺少参数 kb_names 或格式错误". `kb_names` is still honoured when present, keeping the legacy `POST /api/kb/retrieve` endpoint working unchanged. --- .../services/knowledge_base_service.py | 9 ++ .../test_knowledge_base_service_contract.py | 84 +++++++++++++++++++ 2 files changed, 93 insertions(+) diff --git a/astrbot/dashboard/services/knowledge_base_service.py b/astrbot/dashboard/services/knowledge_base_service.py index 76ad91359d..9aa903c051 100644 --- a/astrbot/dashboard/services/knowledge_base_service.py +++ b/astrbot/dashboard/services/knowledge_base_service.py @@ -769,6 +769,15 @@ async def retrieve(self, data: object) -> dict[str, Any]: if not query: raise KnowledgeBaseServiceError("缺少参数 query") kb_manager = self.get_kb_manager() + if not kb_names: + # The REST route identifies the knowledge base with the kb_id path + # parameter, so resolve it into the kb_names the manager expects. + kb_id = payload.get("kb_id") + if kb_id: + kb_helper = await kb_manager.get_kb(kb_id) + if not kb_helper: + raise KnowledgeBaseServiceError("知识库不存在") + kb_names = [kb_helper.kb.kb_name] if not kb_names or not isinstance(kb_names, list): raise KnowledgeBaseServiceError("缺少参数 kb_names 或格式错误") diff --git a/tests/unit/test_knowledge_base_service_contract.py b/tests/unit/test_knowledge_base_service_contract.py index 5bcdddf2fd..80a18c6966 100644 --- a/tests/unit/test_knowledge_base_service_contract.py +++ b/tests/unit/test_knowledge_base_service_contract.py @@ -7,9 +7,11 @@ from astrbot.core.provider.provider import EmbeddingProvider from astrbot.dashboard.api.knowledge_bases import ( list_knowledge_bases, + retrieve_knowledge_base, ) from astrbot.dashboard.schemas import ( KnowledgeBaseRequest, + KnowledgeBaseRetrieveRequest, ) from astrbot.dashboard.services.knowledge_base_service import ( KnowledgeBaseService, @@ -254,6 +256,88 @@ async def test_create_kb_raises_when_embedding_provider_is_missing(): await service.create_kb({"kb_name": "Test KB"}) +@pytest.mark.asyncio +async def test_retrieve_resolves_kb_id_from_path_parameter(): + kb = make_kb("kb-1", "Docs") + kb_manager = MagicMock() + kb_manager.get_kb = AsyncMock(return_value=SimpleNamespace(kb=kb)) + kb_manager.retrieve = AsyncMock(return_value={"results": [{"chunk_id": "c-1"}]}) + service = make_service(kb_manager) + + result = await service.retrieve({"kb_id": "kb-1", "query": "hello", "top_k": 3}) + + kb_manager.get_kb.assert_awaited_once_with("kb-1") + kb_manager.retrieve.assert_awaited_once_with( + query="hello", + kb_names=["Docs"], + top_m_final=3, + ) + assert result == { + "results": [{"chunk_id": "c-1"}], + "total": 1, + "query": "hello", + } + + +@pytest.mark.asyncio +async def test_retrieve_route_passes_kb_id_to_service(): + kb = make_kb("kb-1", "Docs") + kb_manager = MagicMock() + kb_manager.get_kb = AsyncMock(return_value=SimpleNamespace(kb=kb)) + kb_manager.retrieve = AsyncMock(return_value={"results": []}) + service = make_service(kb_manager) + + response = await retrieve_knowledge_base( + "kb-1", + KnowledgeBaseRetrieveRequest(query="hello"), + _auth=object(), + service=service, + ) + + assert response["status"] == "ok" + kb_manager.retrieve.assert_awaited_once_with( + query="hello", + kb_names=["Docs"], + top_m_final=5, + ) + + +@pytest.mark.asyncio +async def test_retrieve_still_accepts_explicit_kb_names(): + kb_manager = MagicMock() + kb_manager.get_kb = AsyncMock() + kb_manager.retrieve = AsyncMock(return_value={"results": []}) + service = make_service(kb_manager) + + await service.retrieve({"query": "hello", "kb_names": ["Docs", "Notes"]}) + + kb_manager.get_kb.assert_not_awaited() + kb_manager.retrieve.assert_awaited_once_with( + query="hello", + kb_names=["Docs", "Notes"], + top_m_final=5, + ) + + +@pytest.mark.asyncio +async def test_retrieve_raises_when_kb_id_is_unknown(): + kb_manager = MagicMock() + kb_manager.get_kb = AsyncMock(return_value=None) + service = make_service(kb_manager) + + with pytest.raises(KnowledgeBaseServiceError, match="知识库不存在"): + await service.retrieve({"kb_id": "missing", "query": "hello"}) + + +@pytest.mark.asyncio +async def test_retrieve_raises_when_no_knowledge_base_is_specified(): + kb_manager = MagicMock() + service = make_service(kb_manager) + + with pytest.raises(KnowledgeBaseServiceError, match="缺少参数 kb_names 或格式错误"): + await service.retrieve({"query": "hello"}) + + @pytest.mark.asyncio async def test_create_kb_raises_when_embedding_provider_is_invalid(): kb_manager = MagicMock() From 6b9cb8aec18e24ccc6f023f71e35703e4200750a Mon Sep 17 00:00:00 2001 From: KazusaUwU <204833675+JosephTian876@users.noreply.github.com> Date: Sat, 1 Aug 2026 01:51:16 +0800 Subject: [PATCH 2/2] docs(kb): document retrieve() payload contract and kb_id precedence The method had no docstring, so the precedence between `kb_names` and the `kb_id` path parameter was only discoverable by reading the body. Spell out that an explicit `kb_names` wins, which is what keeps the legacy multi-base `POST /api/kb/retrieve` endpoint working. --- .../services/knowledge_base_service.py | 20 +++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/astrbot/dashboard/services/knowledge_base_service.py b/astrbot/dashboard/services/knowledge_base_service.py index 9aa903c051..3976903cdf 100644 --- a/astrbot/dashboard/services/knowledge_base_service.py +++ b/astrbot/dashboard/services/knowledge_base_service.py @@ -761,6 +761,26 @@ async def list_chunks_from_dashboard_query( ) async def retrieve(self, data: object) -> dict[str, Any]: + """Retrieve chunks from one or more knowledge bases. + + Args: + data: Request payload. ``query`` is required. The knowledge bases + are taken from ``kb_names`` when present, which is how the + legacy ``POST /api/kb/retrieve`` endpoint targets several bases + at once; otherwise the ``kb_id`` that the REST route reads from + its path is resolved to the matching name. ``top_k`` defaults + to 5, and ``debug`` additionally renders a t-SNE plot. + + Returns: + A dict with the ``results`` list, its ``total`` and the ``query`` + that produced it, plus ``visualization`` (or ``visualization_error``) + when ``debug`` is set. + + Raises: + KnowledgeBaseServiceError: If ``query`` is missing, the given + ``kb_id`` does not exist, or neither ``kb_names`` nor ``kb_id`` + identifies a knowledge base. + """ payload = self._payload(data) query = payload.get("query") kb_names = payload.get("kb_names")