fix: align knowledge base CRUD API contract with kb_name, canonical payload fields, and matching OpenAPI types (#9000)

Warning: Potential breaking changes in /kb related openapi.

* fix(kb): align CRUD API contract end-to-end

Unify knowledge base create/update/list with kb_name, canonical_payload,
optional list pagination with total, and matching OpenAPI plus dashboard types.

* fix: restore _model_dict references, add KnowledgeBaseCreateRequest, fix OpenAPI allOf validation

- Replace 3 remaining _model_dict() calls with payload.model_dump(exclude_none=True)
- Add KnowledgeBaseCreateRequest schema with required kb_name + embedding_provider_id
- Remove additionalProperties: false from OpenAPI allOf sub-schema to avoid validation failures
- Use Pydantic Field(alias="name") + populate_by_name=True for legacy name migration
- Delegate _canonical_kb_payload to KnowledgeBaseRequest to eliminate duplicated logic
- Add getattr default None for update_kb to prevent AttributeError
- Add validation error tests: missing kb_name, missing embedding_provider_id

* fix: address knowledge base contract follow-ups

* fix: preserve knowledge base create error shape
This commit is contained in:
lxfight
2026-07-03 22:18:56 +08:00
committed by GitHub
parent e2f3b00088
commit ea9e3421d9
9 changed files with 419 additions and 60 deletions

View File

@@ -9,6 +9,7 @@ from astrbot.core import logger
from astrbot.dashboard.async_utils import run_maybe_async
from astrbot.dashboard.responses import error, ok
from astrbot.dashboard.schemas import (
KnowledgeBaseCreateRequest,
KnowledgeBaseImportRequest,
KnowledgeBaseRequest,
KnowledgeBaseRetrieveRequest,
@@ -53,14 +54,6 @@ def _to_int(value: Any, default: int) -> int:
return default
def _model_dict(payload) -> dict[str, Any]:
if payload is None:
return {}
if hasattr(payload, "model_dump"):
return payload.model_dump(exclude_none=True)
return payload if isinstance(payload, dict) else {}
async def _run(operation, *, prefix: str):
try:
result = await run_maybe_async(operation)
@@ -102,12 +95,12 @@ async def list_knowledge_bases(
@router.post("/knowledge-bases")
async def create_knowledge_base(
payload: KnowledgeBaseRequest,
payload: KnowledgeBaseCreateRequest,
_auth: AuthContext = Depends(require_kb_scope),
service: KnowledgeBaseService = Depends(get_service),
):
return await _run(
lambda: service.create_kb(_model_dict(payload)),
lambda: service.create_kb(payload.canonical_payload()),
prefix="创建知识库失败",
)
@@ -140,9 +133,8 @@ async def update_knowledge_base(
_auth: AuthContext = Depends(require_kb_scope),
service: KnowledgeBaseService = Depends(get_service),
):
body = _model_dict(payload)
return await _run(
lambda: service.update_kb({"kb_id": kb_id, **body}),
lambda: service.update_kb({**payload.canonical_payload(), "kb_id": kb_id}),
prefix="更新知识库失败",
)
@@ -213,7 +205,7 @@ async def import_knowledge_base_documents(
_auth: AuthContext = Depends(require_kb_scope),
service: KnowledgeBaseService = Depends(get_service),
):
body = _model_dict(payload)
body = payload.model_dump(exclude_none=True)
return await _run(
lambda: service.import_documents({"kb_id": kb_id, **body}),
prefix="导入文档失败",
@@ -227,7 +219,7 @@ async def import_knowledge_base_document_url(
_auth: AuthContext = Depends(require_kb_scope),
service: KnowledgeBaseService = Depends(get_service),
):
body = _model_dict(payload)
body = payload.model_dump(exclude_none=True)
return await _run(
lambda: service.upload_document_from_url({"kb_id": kb_id, **body}),
prefix="从URL上传文档失败",
@@ -307,7 +299,7 @@ async def retrieve_knowledge_base(
_auth: AuthContext = Depends(require_kb_scope),
service: KnowledgeBaseService = Depends(get_service),
):
body = _model_dict(payload)
body = payload.model_dump(exclude_none=True)
return await _run(
lambda: service.retrieve({"kb_id": kb_id, **body}),
prefix="检索失败",

View File

@@ -207,13 +207,49 @@ class ImMessageRequest(OpenModel):
class KnowledgeBaseRequest(OpenModel):
kb_id: str | None = None
name: str | None = None
kb_name: str | None = Field(None, alias="name")
description: str | None = None
emoji: str | None = None
embedding_provider_id: str | None = None
rerank_provider_id: str | None = None
chunk_size: int | None = None
chunk_overlap: int | None = None
top_k_dense: int | None = None
top_k_sparse: int | None = None
top_m_final: int | None = None
model_config = ConfigDict(populate_by_name=True, extra="allow")
def canonical_payload(self) -> dict[str, Any]:
"""Return the service-facing knowledge base payload.
Returns:
Dictionary accepted by KnowledgeBaseService.
"""
return self.model_dump(
exclude_unset=True,
include={
"kb_name",
"description",
"emoji",
"embedding_provider_id",
"rerank_provider_id",
"chunk_size",
"chunk_overlap",
"top_k_dense",
"top_k_sparse",
"top_m_final",
},
by_alias=False,
)
class KnowledgeBaseCreateRequest(KnowledgeBaseRequest):
model_config = ConfigDict(
populate_by_name=True,
extra="allow",
json_schema_extra={"required": ["name", "embedding_provider_id"]},
)
class KnowledgeBaseImportRequest(OpenModel):

View File

@@ -12,6 +12,7 @@ from astrbot.core import logger
from astrbot.core.core_lifecycle import AstrBotCoreLifecycle
from astrbot.core.provider.provider import EmbeddingProvider, RerankProvider
from astrbot.core.utils.astrbot_path import get_astrbot_temp_path
from astrbot.dashboard.schemas import KnowledgeBaseRequest
from astrbot.dashboard.utils import generate_tsne_visualization
@@ -29,6 +30,19 @@ class KnowledgeBaseService:
def _payload(data: object) -> dict[str, Any]:
return data if isinstance(data, dict) else {}
@staticmethod
def _canonical_kb_payload(data: object) -> dict[str, Any]:
"""Normalize knowledge base create/update payloads.
Uses KnowledgeBaseRequest to handle the legacy ``name`` →
``kb_name`` migration while preserving operational fields
like ``kb_id``.
"""
raw = KnowledgeBaseService._payload(data)
canonical = KnowledgeBaseRequest(**raw).canonical_payload()
raw.update(canonical)
return raw
def get_kb_manager(self):
return self.core_lifecycle.kb_manager
@@ -293,7 +307,7 @@ class KnowledgeBaseService:
async def create_kb(self, data: object) -> tuple[dict[str, Any], str]:
kb_manager = self.get_kb_manager()
payload = self._payload(data)
payload = self._canonical_kb_payload(data)
kb_name = payload.get("kb_name")
if not kb_name:
raise KnowledgeBaseServiceError("知识库名称不能为空")
@@ -363,7 +377,7 @@ class KnowledgeBaseService:
return await self.get_kb(kb_id)
async def update_kb(self, data: object) -> tuple[dict[str, Any], str]:
payload = self._payload(data)
payload = self._canonical_kb_payload(data)
kb_id = payload.get("kb_id")
if not kb_id:
raise KnowledgeBaseServiceError("缺少参数 kb_id")
@@ -380,28 +394,20 @@ class KnowledgeBaseService:
"top_k_sparse",
"top_m_final",
]
if all(payload.get(key) is None for key in update_keys):
provided_updates = {key: payload[key] for key in update_keys if key in payload}
if not provided_updates:
raise KnowledgeBaseServiceError("至少需要提供一个更新字段")
current_kb = await self.get_kb_manager().get_kb(kb_id)
kb_name = payload.get("kb_name")
if kb_name is None:
if not current_kb:
raise KnowledgeBaseServiceError("知识库不存在")
kb_name = current_kb.kb.kb_name
if not current_kb:
raise KnowledgeBaseServiceError("知识库不存在")
current = current_kb.kb
update_data = {key: getattr(current, key, None) for key in update_keys}
update_data.update(provided_updates)
kb_helper = await self.get_kb_manager().update_kb(
kb_id=kb_id,
kb_name=kb_name,
description=payload.get("description"),
emoji=payload.get("emoji"),
embedding_provider_id=payload.get("embedding_provider_id"),
rerank_provider_id=payload.get("rerank_provider_id"),
chunk_size=payload.get("chunk_size"),
chunk_overlap=payload.get("chunk_overlap"),
top_k_dense=payload.get("top_k_dense"),
top_k_sparse=payload.get("top_k_sparse"),
top_m_final=payload.get("top_m_final"),
**update_data,
)
if not kb_helper:
raise KnowledgeBaseServiceError("知识库不存在")
@@ -762,11 +768,11 @@ class KnowledgeBaseService:
if not query:
raise KnowledgeBaseServiceError("缺少参数 query")
kb_manager = self.get_kb_manager()
if not kb_names or not isinstance(kb_names, list):
raise KnowledgeBaseServiceError("缺少参数 kb_names 或格式错误")
top_k = payload.get("top_k", 5)
kb_manager = self.get_kb_manager()
results = await kb_manager.retrieve(
query=query,
kb_names=kb_names,