Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions astrbot/dashboard/services/knowledge_base_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -769,6 +789,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 或格式错误")

Expand Down
84 changes: 84 additions & 0 deletions tests/unit/test_knowledge_base_service_contract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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()
Expand Down