mirror of
https://github.com/AstrBotDevs/AstrBot
synced 2026-07-20 02:55:08 +08:00
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:
@@ -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="检索失败",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user