ForcePilot/backend/package/yuxi/knowledge/base.py

1705 lines
64 KiB
Python
Raw Normal View History

import asyncio
import mimetypes
import os
import re
from abc import ABC, abstractmethod
from typing import Any
from yuxi.knowledge.chunking.ragflow_like.presets import ensure_chunk_defaults_in_additional_params
from yuxi.knowledge.schemas import FindOutputSchema, FindWindowSchema, SearchOutputSchema, SearchResultSchema
from yuxi.knowledge.utils import resolve_processing_params, sanitize_processing_params
from yuxi.utils import logger
from yuxi.utils.datetime_utils import coerce_any_to_utc_datetime, utc_isoformat
class FileStatus:
UPLOADED = "uploaded"
PARSING = "parsing"
PARSED = "parsed"
ERROR_PARSING = "error_parsing"
INDEXING = "indexing"
INDEXED = "indexed"
ERROR_INDEXING = "error_indexing"
class KnowledgeBaseException(Exception):
"""知识库统一异常基类"""
pass
class KBNotFoundError(KnowledgeBaseException):
"""知识库不存在错误"""
pass
class KBOperationError(KnowledgeBaseException):
"""知识库操作错误"""
pass
class KnowledgeBase(ABC):
"""知识库抽象基类,定义统一接口"""
requires_embedding_model = True
supports_documents = True
apply_chunk_defaults = True
# 类级别的处理队列,跟踪所有正在处理的文件
_processing_files = set()
_processing_lock = None
def __init__(self, work_dir: str):
"""
初始化知识库
Args:
work_dir: 工作目录
"""
import threading
self.work_dir = work_dir
self.databases_meta: dict[str, dict] = {}
self.files_meta: dict[str, dict] = {}
self.benchmarks_meta: dict[str, dict] = {}
self._metadata_loaded = False # 标记元数据是否已加载
# 初始化类级别的锁
if KnowledgeBase._processing_lock is None:
KnowledgeBase._processing_lock = threading.Lock()
os.makedirs(work_dir, exist_ok=True)
# 注意:不在 __init__ 中加载元数据,由 KnowledgeBaseManager 统一管理加载
def load_metadata(
self, global_databases_meta: dict[str, dict], files_meta: dict[str, dict], benchmarks_meta: dict[str, dict]
):
"""由 KnowledgeBaseManager 调用,同步加载元数据"""
# 过滤出当前 kb_type 的知识库
self.databases_meta = {}
for db_id, meta in global_databases_meta.items():
if meta.get("kb_type") == self.kb_type:
normalized_additional_params = self.normalize_additional_params(meta.get("additional_params"))
self.databases_meta[db_id] = {
"name": meta.get("name"),
"description": meta.get("description"),
"kb_type": meta.get("kb_type"),
"embedding_model_spec": meta.get("embedding_model_spec"),
"llm_model_spec": meta.get("llm_model_spec"),
"query_params": meta.get("query_params"),
"metadata": normalized_additional_params,
"created_at": meta.get("created_at"),
}
# 过滤文件
self.files_meta = {}
for file_id, meta in files_meta.items():
if meta.get("database_id") in self.databases_meta:
db_id = meta.get("database_id")
kb_additional_params = self.databases_meta.get(db_id, {}).get("metadata") or {}
normalized_meta = dict(meta)
normalized_meta["processing_params"] = resolve_processing_params(
kb_additional_params=kb_additional_params,
file_processing_params=meta.get("processing_params"),
)
self.files_meta[file_id] = normalized_meta
# 过滤评估基准
self.benchmarks_meta = {}
for kb_id, benchmarks in benchmarks_meta.items():
if kb_id in self.databases_meta:
self.benchmarks_meta[kb_id] = benchmarks
self._normalize_metadata_state()
self._metadata_loaded = True
logger.info(f"{self.kb_type}: 加载了 {len(self.databases_meta)} 个数据库的元数据")
def _ensure_metadata_loaded(self):
"""确保元数据已加载(延迟加载)"""
if not self._metadata_loaded:
logger.warning(f"{self.kb_type}: 元数据尚未加载,请确保 KnowledgeBaseManager 已调用 load_metadata()")
@staticmethod
def _normalize_timestamp(value: Any) -> str | None:
"""Convert persisted timestamps to a normalized UTC ISO string."""
try:
dt_value = coerce_any_to_utc_datetime(value)
except (TypeError, ValueError) as exc: # noqa: BLE001
logger.warning(f"Invalid timestamp encountered: {value!r} ({exc})")
return None
if not dt_value:
return None
return utc_isoformat(dt_value)
def _normalize_metadata_state(self) -> None:
"""Ensure in-memory metadata uses normalized timestamp formats."""
for meta in self.databases_meta.values():
if "created_at" in meta:
normalized = self._normalize_timestamp(meta.get("created_at"))
if normalized:
meta["created_at"] = normalized
for file_info in self.files_meta.values():
if "created_at" in file_info:
normalized = self._normalize_timestamp(file_info.get("created_at"))
if normalized:
file_info["created_at"] = normalized
for db_benchmarks in self.benchmarks_meta.values():
for b in db_benchmarks.values():
if "created_at" in b:
normalized = self._normalize_timestamp(b.get("created_at"))
if normalized:
b["created_at"] = normalized
if "updated_at" in b:
normalized = self._normalize_timestamp(b.get("updated_at"))
if normalized:
b["updated_at"] = normalized
@property
@abstractmethod
def kb_type(self) -> str:
"""知识库类型标识"""
pass
@classmethod
def get_create_params_config(cls) -> dict[str, Any]:
"""获取创建知识库时的类型特定参数配置。"""
return {"options": []}
@classmethod
def validate_additional_params(cls, additional_params: dict | None) -> dict:
"""校验并规范化类型特定配置。"""
return dict(additional_params or {})
@classmethod
def normalize_additional_params(cls, additional_params: dict | None) -> dict:
"""规范化 additional_params仅文档型知识库补充分块默认值。"""
params = cls.validate_additional_params(additional_params)
if cls.apply_chunk_defaults:
return ensure_chunk_defaults_in_additional_params(params)
return params
@abstractmethod
async def _create_kb_instance(self, db_id: str, config: dict) -> Any:
"""
创建底层知识库实例
Args:
db_id: 数据库ID
config: 配置信息
Returns:
底层知识库实例
"""
pass
@abstractmethod
async def _initialize_kb_instance(self, instance: Any) -> None:
"""
初始化底层知识库实例
Args:
instance: 底层知识库实例
"""
pass
async def add_file_record(
self, db_id: str, item: str, params: dict | None = None, operator_id: str | None = None
) -> dict:
"""
Add a file record to metadata (Status: UPLOADED)
Args:
db_id: Database ID
item: File path or URL
params: Parameters
operator_id: Operator ID who created the file
Returns:
File metadata record
"""
from yuxi.knowledge.utils.kb_utils import prepare_item_metadata
params = params or {}
content_type = params.get("content_type", "file")
# Prepare metadata
metadata = await prepare_item_metadata(item, content_type, db_id, params=params)
file_id = metadata["file_id"]
kb_additional_params = self.databases_meta.get(db_id, {}).get("metadata") or {}
metadata["processing_params"] = resolve_processing_params(
kb_additional_params=kb_additional_params,
file_processing_params=metadata.get("processing_params"),
)
# Fallback: fetch file size from MinIO if not provided
if metadata.get("size") is None and content_type == "file":
try:
from yuxi.knowledge.utils.kb_utils import is_minio_url, parse_minio_url
from yuxi.storage.minio import get_minio_client
file_path = metadata.get("path") or item
if is_minio_url(file_path):
bucket_name, obj_name = parse_minio_url(file_path)
minio_client = get_minio_client()
file_size = await minio_client.astat_file(bucket_name, obj_name)
if file_size is not None:
metadata["size"] = file_size
except Exception as exc:
logger.warning(f"Failed to stat file size from MinIO for {item}: {exc}")
# Initial status
metadata["status"] = FileStatus.UPLOADED
metadata["created_at"] = utc_isoformat()
if operator_id:
metadata["created_by"] = operator_id
# Save to metadata
self.files_meta[file_id] = metadata
await self._persist_file(file_id)
return metadata
async def parse_file(self, db_id: str, file_id: str, operator_id: str | None = None) -> dict:
"""
Parse file to Markdown and save to MinIO (Status: PARSING -> PARSED/ERROR_PARSING)
Args:
db_id: Database ID
file_id: File ID
operator_id: ID of the user performing the operation
Returns:
Updated file metadata
"""
if file_id not in self.files_meta:
raise ValueError(f"File {file_id} not found")
file_meta = self.files_meta[file_id]
current_status = file_meta.get("status")
# Validate current status - only allow parsing from these states
allowed_statuses = {
FileStatus.UPLOADED,
FileStatus.ERROR_PARSING,
"failed", # Legacy status
}
if current_status not in allowed_statuses:
raise ValueError(
f"Cannot parse file with status '{current_status}'. "
f"File must be in one of these states: {', '.join(allowed_statuses)}"
)
file_path = file_meta.get("path")
if not file_path:
raise ValueError(f"File {file_id} has no valid path in metadata")
# Clear previous error if any
if "error" in file_meta:
self.files_meta[file_id].pop("error", None)
# Update status to PARSING and add to processing queue
self.files_meta[file_id]["status"] = FileStatus.PARSING
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
await self._persist_file(file_id)
# Add to processing queue
self._add_to_processing_queue(file_id)
try:
from yuxi.plugins.parser.unified import Parser
# Prepare params
params = file_meta.get("processing_params", {}) or {}
params["image_bucket"] = "public"
params["image_prefix"] = f"{db_id}/kb-images"
markdown_content = await Parser.aparse(
source=file_path,
params=params,
)
# Save Markdown to MinIO
markdown_file_path = await self._save_markdown_to_minio(db_id, file_id, markdown_content)
# Update metadata
self.files_meta[file_id]["status"] = FileStatus.PARSED
self.files_meta[file_id]["markdown_file"] = markdown_file_path
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
await self._persist_file(file_id)
return self.files_meta[file_id]
except Exception as e:
error_msg = str(e)
logger.error(f"Failed to parse file {file_id}: {error_msg}")
self.files_meta[file_id]["status"] = FileStatus.ERROR_PARSING
self.files_meta[file_id]["error"] = error_msg
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
await self._persist_file(file_id)
raise
finally:
# Remove from processing queue
self._remove_from_processing_queue(file_id)
async def update_file_params(self, db_id: str, file_id: str, params: dict, operator_id: str | None = None) -> None:
"""Update file processing params"""
if file_id not in self.files_meta:
raise ValueError(f"File {file_id} not found")
# Skip if no params to update
if not params:
return
current_params = self.files_meta[file_id].get("processing_params", {}) or {}
kb_additional_params = self.databases_meta.get(db_id, {}).get("metadata") or {}
logger.debug(f"[update_file_params] file_id={file_id}, current_params={current_params}, new_params={params}")
current_params = resolve_processing_params(
kb_additional_params=kb_additional_params,
file_processing_params=current_params,
request_params=params,
)
self.files_meta[file_id]["processing_params"] = current_params
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
logger.debug(f"[update_file_params] file_id={file_id}, updated_params={current_params}")
await self._persist_file(file_id)
async def _mark_file_unparsed(self, file_id: str, operator_id: str | None = None) -> None:
if file_id not in self.files_meta:
return
self.files_meta[file_id]["status"] = FileStatus.UPLOADED
self.files_meta[file_id].pop("markdown_file", None)
self.files_meta[file_id].pop("error", None)
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
await self._persist_file(file_id)
async def _save_markdown_to_minio(self, db_id: str, file_id: str, content: str) -> str:
"""Save markdown content to MinIO and return HTTP URL"""
from yuxi.storage.minio import get_minio_client
minio_client = get_minio_client()
bucket_name = minio_client.KB_BUCKETS["parsed"]
await asyncio.to_thread(minio_client.ensure_bucket_exists, bucket_name)
object_name = f"{db_id}/parsed/{file_id}.md"
data = content.encode("utf-8")
# Return standard HTTP URL from UploadResult
upload_result = await minio_client.aupload_file(
bucket_name=bucket_name,
object_name=object_name,
data=data,
)
return upload_result.url
async def _read_minio_bytes(self, file_path: str) -> bytes:
from yuxi.knowledge.utils.kb_utils import is_minio_url, parse_minio_url
from yuxi.storage.minio import get_minio_client
if not file_path or not is_minio_url(file_path):
raise ValueError(f"Invalid MinIO path format: {file_path}")
bucket_name, object_name = parse_minio_url(file_path)
minio_client = get_minio_client()
return await minio_client.adownload_file(bucket_name, object_name)
async def _read_markdown_from_minio(self, file_path: str) -> str:
"""Read markdown content from MinIO"""
content_bytes = await self._read_minio_bytes(file_path)
return content_bytes.decode("utf-8")
def _get_file_meta(self, db_id: str, file_id: str) -> dict:
file_meta = self.files_meta.get(file_id)
if not file_meta or file_meta.get("database_id") != db_id:
raise ValueError(f"File {file_id} not found")
return file_meta
@staticmethod
def _original_file_path(file_meta: dict) -> str | None:
return file_meta.get("minio_url") or file_meta.get("path")
def _knowledge_preview_variants(self, file_meta: dict) -> list[dict]:
variants = []
original_path = self._original_file_path(file_meta)
if original_path:
variants.append({"key": "original", "label": "Source", "supported": True})
if file_meta.get("markdown_file"):
variants.append({"key": "parsed", "label": "MD", "supported": True})
return variants
def _knowledge_file_entry(self, db_id: str, file_id: str, file_meta: dict) -> dict:
is_dir = bool(file_meta.get("is_folder"))
variants = [] if is_dir else self._knowledge_preview_variants(file_meta)
preview_modes = [item["key"] for item in variants]
default_preview_mode = None
if "parsed" in preview_modes:
default_preview_mode = "parsed"
elif preview_modes:
default_preview_mode = preview_modes[0]
path = f"/{file_id}"
if is_dir:
path = f"{path}/"
return {
"source": "knowledge",
"db_id": db_id,
"file_id": file_id,
"parent_id": file_meta.get("parent_id"),
"path": path,
"virtual_path": f"/knowledge/{db_id}/{file_id}",
"name": file_meta.get("filename") or file_meta.get("original_filename") or file_id,
"is_dir": is_dir,
"size": 0 if is_dir else file_meta.get("size") or 0,
"modified_at": file_meta.get("updated_at") or file_meta.get("created_at") or "",
"readonly": True,
"status": file_meta.get("status", "done"),
"preview_modes": preview_modes,
"default_preview_mode": default_preview_mode,
}
def _sort_file_entries(self, entries: list[dict]) -> list[dict]:
return sorted(
entries,
key=lambda item: (not bool(item.get("is_dir")), str(item.get("name") or "").lower()),
)
def _list_knowledge_children(
self,
db_id: str,
parent_id: str | None,
*,
recursive: bool,
files_only: bool,
) -> list[dict]:
children = [
(file_id, meta)
for file_id, meta in self.files_meta.items()
if meta.get("database_id") == db_id and meta.get("parent_id") == parent_id
]
entries = []
for file_id, meta in children:
if not files_only or not meta.get("is_folder"):
entries.append(self._knowledge_file_entry(db_id, file_id, meta))
if recursive and meta.get("is_folder"):
entries.extend(
self._list_knowledge_children(
db_id,
file_id,
recursive=True,
files_only=files_only,
)
)
return self._sort_file_entries(entries)
async def list_file_tree(
self,
db_id: str,
parent_id: str | None = None,
recursive: bool = False,
files_only: bool = False,
) -> dict:
if db_id not in self.databases_meta:
raise ValueError(f"Database {db_id} not found")
if parent_id:
parent_meta = self._get_file_meta(db_id, parent_id)
if not parent_meta.get("is_folder"):
raise ValueError("Parent is not a folder")
return {
"entries": self._list_knowledge_children(
db_id,
parent_id,
recursive=recursive,
files_only=files_only,
),
"readonly": True,
}
async def read_file_preview(self, db_id: str, file_id: str, variant: str = "parsed") -> dict:
from yuxi.services.viewer_filesystem_service import _detect_preview_type
file_meta = self._get_file_meta(db_id, file_id)
if file_meta.get("is_folder"):
raise ValueError("Cannot preview a folder")
variants = self._knowledge_preview_variants(file_meta)
variant_keys = {item["key"] for item in variants}
if variant not in {"original", "parsed"}:
raise ValueError("Unsupported preview variant")
filename = file_meta.get("filename") or file_meta.get("original_filename") or file_id
response = {
"source": "knowledge",
"db_id": db_id,
"file_id": file_id,
"filename": filename,
"variant": variant,
"readonly": True,
"available_variants": variants,
}
if variant == "parsed":
markdown_file = file_meta.get("markdown_file")
if not markdown_file or "parsed" not in variant_keys:
return {
**response,
"content": None,
"preview_type": "unsupported",
"supported": False,
"message": "文件尚未生成解析结果",
}
content = await self._read_markdown_from_minio(markdown_file)
return {
**response,
"content": content,
"preview_type": "markdown",
"supported": True,
"message": None,
}
original_path = self._original_file_path(file_meta)
if not original_path or "original" not in variant_keys:
return {
**response,
"content": None,
"preview_type": "unsupported",
"supported": False,
"message": "文件没有可预览的原始内容",
}
preview_type, supported, message = _detect_preview_type(filename, b"")
if preview_type in {"image", "pdf"}:
return {
**response,
"content": None,
"preview_type": preview_type,
"supported": supported,
"message": message,
}
raw_content = await self._read_minio_bytes(original_path)
preview_type, supported, message = _detect_preview_type(filename, raw_content)
if preview_type in {"image", "pdf"} or not supported:
return {
**response,
"content": None,
"preview_type": preview_type,
"supported": supported,
"message": message,
}
try:
content = raw_content.decode("utf-8")
except UnicodeDecodeError:
return {
**response,
"content": None,
"preview_type": "unsupported",
"supported": False,
"message": "当前文件不是 UTF-8 文本,暂不支持预览",
}
return {
**response,
"content": content,
"preview_type": preview_type,
"supported": True,
"message": message,
}
async def get_file_download(self, db_id: str, file_id: str, variant: str = "original") -> dict:
file_meta = self._get_file_meta(db_id, file_id)
if file_meta.get("is_folder"):
raise ValueError("Cannot download a folder")
if variant not in {"original", "parsed"}:
raise ValueError("Unsupported download variant")
filename = file_meta.get("filename") or file_meta.get("original_filename") or file_id
if variant == "parsed":
markdown_file = file_meta.get("markdown_file")
if not markdown_file:
raise ValueError("文件尚未生成解析结果")
return {
"filename": f"{filename}.parsed.md",
"content": await self._read_minio_bytes(markdown_file),
"media_type": "text/markdown; charset=utf-8",
}
original_path = self._original_file_path(file_meta)
if not original_path:
raise ValueError("文件没有可下载的原始内容")
media_type = file_meta.get("content_type") or mimetypes.guess_type(filename)[0] or "application/octet-stream"
return {
"filename": filename,
"content": await self._read_minio_bytes(original_path),
"media_type": media_type,
}
def _build_open_file_window(self, content: str, *, offset: int = 0, limit: int = 800) -> dict[str, Any]:
lines = content.splitlines()
total_lines = len(lines)
start = min(max(int(offset), 0), total_lines)
window_size = min(max(int(limit), 1), 2000)
selected = lines[start : start + window_size]
end = start + len(selected)
return {
"start_line": start + 1 if selected else 0,
"end_line": end,
"total_lines": total_lines,
"offset": start,
"window_size": window_size,
"has_more_before": start > 0,
"has_more_after": end < total_lines,
"next_offset": end if end < total_lines else None,
"content": "\n".join(f"{start + idx + 1:6d}\t{line}" for idx, line in enumerate(selected)),
}
@staticmethod
def build_search_output(resource_id: str, retrieval_results: Any) -> dict[str, Any] | Any:
if not isinstance(retrieval_results, list):
return retrieval_results
results = []
for index, chunk in enumerate(retrieval_results):
if not isinstance(chunk, dict):
continue
metadata = chunk.get("metadata") if isinstance(chunk.get("metadata"), dict) else {}
metadata = {
key: value
for key, value in metadata.items()
if key not in {"filepath", "parsed_path", "path", "markdown_file"}
}
file_id = metadata.get("file_id") or chunk.get("file_id") or chunk.get("full_doc_id") or ""
chunk_id = metadata.get("chunk_id") or chunk.get("chunk_id") or chunk.get("id")
chunk_index = metadata.get("chunk_index")
if chunk_index is None:
chunk_index = chunk.get("chunk_index")
if chunk_index is not None:
metadata.setdefault("chunk_index", chunk_index)
if chunk.get("score") is not None:
metadata.setdefault("score", chunk.get("score"))
if chunk.get("distance") is not None:
metadata.setdefault("distance", chunk.get("distance"))
results.append(
SearchResultSchema(
id=str(chunk_id or f"{file_id}:{index + 1}"),
resource_id=str(resource_id),
file_id=str(file_id or ""),
content=str(chunk.get("content") or ""),
metadata=metadata,
)
)
return SearchOutputSchema(resource_id=str(resource_id), results=results).model_dump()
@staticmethod
def _build_find_file_windows(
content: str,
*,
patterns: list[str],
use_regex: bool = False,
case_sensitive: bool = False,
max_windows: int = 5,
window_size: int = 80,
) -> dict[str, Any]:
patterns = [pattern for pattern in patterns if pattern]
if not patterns:
raise ValueError("请提供至少一个 pattern")
lines = content.splitlines()
flags = 0 if case_sensitive else re.IGNORECASE
if use_regex:
matchers = [re.compile(pattern, flags) for pattern in patterns]
def line_matches(line: str) -> bool:
return any(matcher.search(line) for matcher in matchers)
else:
normalized_patterns = patterns if case_sensitive else [pattern.lower() for pattern in patterns]
def line_matches(line: str) -> bool:
haystack = line if case_sensitive else line.lower()
return any(pattern in haystack for pattern in normalized_patterns)
matched_indexes = [index for index, line in enumerate(lines) if line_matches(line)]
windows: list[FindWindowSchema] = []
covered_until = -1
normalized_window_size = min(max(int(window_size), 1), 200)
half_window = normalized_window_size // 2
for matched_index in matched_indexes:
if matched_index < covered_until:
continue
start = max(matched_index - half_window, 0)
end = min(start + normalized_window_size, len(lines))
start = max(end - normalized_window_size, 0)
matched_lines = [index + 1 for index in matched_indexes if start <= index < end]
selected = lines[start:end]
windows.append(
FindWindowSchema(
start_line=start + 1 if selected else 0,
end_line=end,
matched_lines=matched_lines,
content="\n".join(f"{start + idx + 1:6d}\t{line}" for idx, line in enumerate(selected)),
)
)
covered_until = end
if len(windows) >= max_windows:
break
return FindOutputSchema(
resource_id="",
file_id="",
semantic=False,
match_mode="regex" if use_regex else "keyword",
total_matches=len(matched_indexes),
windows=windows,
).model_dump(exclude={"resource_id", "file_id"})
async def open_file_content(self, db_id: str, file_id: str, offset: int = 0, limit: int = 800) -> dict:
"""按行窗口打开文件解析后的 Markdown 内容"""
file_meta = self.files_meta.get(file_id)
if file_meta is None:
raise Exception(f"文件不存在: {file_id}")
if file_meta.get("database_id") != db_id:
raise Exception(f"文件 {file_id} 不属于知识库 {db_id}")
if file_meta.get("is_folder"):
raise Exception(f"文件 {file_id} 是文件夹")
markdown_file = file_meta.get("markdown_file")
if not markdown_file:
raise Exception(f"文件 {file_id} 没有解析后的 Markdown 内容")
content = await self._read_markdown_from_minio(markdown_file)
return self._build_open_file_window(content, offset=offset, limit=limit)
async def find_file_content(
self,
db_id: str,
file_id: str,
patterns: list[str],
*,
use_regex: bool = False,
case_sensitive: bool = False,
max_windows: int = 5,
window_size: int = 80,
) -> dict:
file_meta = self.files_meta.get(file_id)
if file_meta is None:
raise Exception(f"文件不存在: {file_id}")
if file_meta.get("database_id") != db_id:
raise Exception(f"文件 {file_id} 不属于知识库 {db_id}")
if file_meta.get("is_folder"):
raise Exception(f"文件 {file_id} 是文件夹")
markdown_file = file_meta.get("markdown_file")
if not markdown_file:
raise Exception(f"文件 {file_id} 没有解析后的 Markdown 内容")
content = await self._read_markdown_from_minio(markdown_file)
return self._build_find_file_windows(
content,
patterns=patterns,
use_regex=use_regex,
case_sensitive=case_sensitive,
max_windows=max_windows,
window_size=window_size,
)
@abstractmethod
async def index_file(self, db_id: str, file_id: str, operator_id: str | None = None) -> dict:
"""
Index parsed file (Status: INDEXING -> INDEXED/ERROR_INDEXING)
Args:
db_id: Database ID
file_id: File ID
operator_id: ID of the user performing the operation
Returns:
Updated file metadata
"""
pass
async def create_database(
self,
database_name: str,
description: str,
embedding_model_spec: str | None = None,
llm_model_spec: str | None = None,
**kwargs,
) -> dict:
"""
创建数据库
Args:
database_name: 数据库名称
description: 数据库描述
embedding_model_spec: 嵌入模型 spec
llm_model_spec: LLM 模型 spec
**kwargs: 其他配置参数
Returns:
数据库信息字典
"""
from yuxi.utils import hashstr
kwargs = self.normalize_additional_params(kwargs)
db_id = f"kb_{hashstr(database_name, with_salt=True, length=32)}"
self.databases_meta[db_id] = {
"name": database_name,
"description": description,
"kb_type": self.kb_type,
"embedding_model_spec": embedding_model_spec,
"llm_model_spec": llm_model_spec,
"metadata": kwargs,
"created_at": utc_isoformat(),
"query_params": self._get_default_query_params(db_id),
}
await self._persist_kb(db_id)
# 创建工作目录
working_dir = os.path.join(self.work_dir, db_id)
os.makedirs(working_dir, exist_ok=True)
# 返回数据库信息
db_dict = self.databases_meta[db_id].copy()
db_dict["db_id"] = db_id
db_dict["files"] = {}
return db_dict
async def delete_database(self, db_id: str) -> dict:
"""
删除数据库
Args:
db_id: 数据库ID
Returns:
操作结果
"""
if db_id in self.databases_meta:
from yuxi.knowledge.utils.kb_utils import is_minio_url, parse_minio_url
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
from yuxi.storage.minio import get_minio_client
minio_client = get_minio_client()
# 1. 删除文件元数据中记录的 MinIO 文件
files_to_delete = [fid for fid, finfo in self.files_meta.items() if finfo.get("database_id") == db_id]
for file_id in files_to_delete:
file_path = self.files_meta[file_id].get("path")
if file_path and is_minio_url(file_path):
try:
bucket_name, object_name = parse_minio_url(file_path)
await minio_client.adelete_file(bucket_name, object_name)
except Exception as e:
logger.warning(f"Failed to delete MinIO file {file_path}: {e}")
# 删除解析后的 markdown 文件
parsed_object = f"{db_id}/parsed/{file_id}.md"
await minio_client.adelete_file(minio_client.KB_BUCKETS["parsed"], parsed_object)
del self.files_meta[file_id]
# 2. 并行删除所有知识库 bucket 中该 db_id 下的文件
prefix = f"{db_id}/"
cleanup_buckets = {
minio_client.KB_BUCKETS["parsed"],
minio_client.KB_BUCKETS["documents"],
minio_client.KB_BUCKETS["images"],
}
cleanup_tasks = [
minio_client.adelete_objects_by_prefix(bucket_name, prefix) for bucket_name in cleanup_buckets
]
await asyncio.gather(*cleanup_tasks)
# 3. 删除数据库记录
del self.databases_meta[db_id]
kb_repo = KnowledgeBaseRepository()
await kb_repo.delete(db_id)
await self._save_metadata()
# 删除工作目录
working_dir = os.path.join(self.work_dir, db_id)
if os.path.exists(working_dir):
import shutil
try:
shutil.rmtree(working_dir)
except Exception as e:
logger.error(f"Error deleting working directory {working_dir}: {e}")
return {"message": "删除成功"}
async def create_folder(self, db_id: str, folder_name: str, parent_id: str | None = None) -> dict:
"""Create a folder in the database."""
import uuid
folder_id = f"folder-{uuid.uuid4()}"
self.files_meta[folder_id] = {
"file_id": folder_id,
"filename": folder_name,
"is_folder": True,
"parent_id": parent_id,
"database_id": db_id,
"created_at": utc_isoformat(),
"status": "done",
"path": folder_name,
"file_type": "folder",
}
await self._persist_file(folder_id)
return self.files_meta[folder_id]
@abstractmethod
async def update_content(self, db_id: str, file_ids: list[str], params: dict | None = None) -> list[dict]:
"""
更新内容 - 根据file_ids重新解析文件并更新向量库
Args:
db_id: 数据库ID
file_ids: 文件ID列表
params: 处理参数
Returns:
更新结果列表
"""
pass
@abstractmethod
async def aquery(self, query_text: str, db_id: str, **kwargs) -> list[dict]:
"""
异步查询知识库
Args:
query_text: 查询文本
db_id: 数据库ID
**kwargs: 查询参数
Returns:
一个包含字典的列表每个字典代表一个检索到的文档块
"""
pass
@abstractmethod
def get_query_params_config(self, db_id: str, **kwargs) -> dict:
"""
获取知识库类型的查询参数配置
Args:
db_id: 数据库ID
**kwargs: 额外参数
Returns:
dict: {
"type": "kb_type",
"options": [
{
"key": "param_name",
"label": "参数名称",
"type": "select|number|boolean",
"default": default_value,
"options": [...], # 对于 select 类型
"description": "参数描述",
"min": 1, # 对于 number 类型
"max": 100,
"step": 0.1
},
...
]
}
"""
pass
async def export_data(self, db_id: str, format: str = "zip", **kwargs) -> str:
pass
def _get_query_params(self, db_id: str) -> dict:
"""从实例元数据中加载查询参数"""
if db_id in self.databases_meta:
query_params_meta = self.databases_meta[db_id].get("query_params") or {}
return query_params_meta.get("options", {})
return {}
def _get_default_query_params(self, db_id: str) -> dict[str, Any]:
"""从 get_query_params_config 中提取所有参数的默认值,返回 {"options": {...}}"""
config = self.get_query_params_config(db_id)
defaults = {}
for opt in config.get("options", []):
if "default" in opt:
defaults[opt["key"]] = opt["default"]
return {"options": defaults}
def get_database_info(self, db_id: str, include_files: bool = True) -> dict | None:
"""
获取数据库详细信息
Args:
db_id: 数据库ID
include_files: 是否包含文件信息默认为True
Returns:
数据库信息或None
"""
if db_id not in self.databases_meta:
return None
meta = self.databases_meta[db_id].copy()
meta["db_id"] = db_id
# 检查并修复异常的processing状态
self._check_and_fix_processing_status(db_id)
# 统计文件数量(始终计算,即使不加载文件详情)
db_file_count = sum(1 for file_info in self.files_meta.values() if file_info.get("database_id") == db_id)
meta["row_count"] = db_file_count
# 仅在需要时加载文件详情
if include_files:
db_files = {}
for file_id, file_info in self.files_meta.items():
if file_info.get("database_id") == db_id:
created_at = self._normalize_timestamp(file_info.get("created_at"))
db_files[file_id] = {
"file_id": file_id,
"filename": file_info.get("filename", ""),
"path": file_info.get("path", ""),
"markdown_file": file_info.get("markdown_file", ""),
"type": file_info.get("file_type", ""),
"status": file_info.get("status", "done"),
"created_at": created_at,
"is_folder": file_info.get("is_folder", False),
"parent_id": file_info.get("parent_id", None),
}
# 按创建时间倒序排序文件列表
sorted_files = dict(
sorted(
db_files.items(),
key=lambda item: item[1].get("created_at") or "",
reverse=True,
)
)
meta["files"] = sorted_files
meta["status"] = "已连接"
return meta
def get_databases(self, include_files: bool = False) -> dict:
"""
获取所有数据库信息
Args:
include_files: 是否包含文件信息默认False以减少响应大小
Returns:
数据库列表
"""
# 确保元数据已加载(延迟加载机制)
self._ensure_metadata_loaded()
databases = []
for db_id, meta in self.databases_meta.items():
# 检查并修复异常的processing状态
self._check_and_fix_processing_status(db_id)
db_dict = meta.copy()
db_dict["db_id"] = db_id
# 统计文件数量(始终计算,即使不加载文件详情)
db_file_count = sum(1 for file_info in self.files_meta.values() if file_info.get("database_id") == db_id)
db_dict["row_count"] = db_file_count
# 仅在需要时加载文件详情
if include_files:
db_files = {}
for file_id, file_info in self.files_meta.items():
if file_info.get("database_id") == db_id:
created_at = self._normalize_timestamp(file_info.get("created_at"))
db_files[file_id] = {
"file_id": file_id,
"filename": file_info.get("filename", ""),
"path": file_info.get("path", ""),
"markdown_file": file_info.get("markdown_file", ""),
"type": file_info.get("file_type", ""),
"status": file_info.get("status", "done"),
"created_at": created_at,
"is_folder": file_info.get("is_folder", False),
"parent_id": file_info.get("parent_id", None),
}
# 按创建时间倒序排序文件列表
sorted_files = dict(
sorted(
db_files.items(),
key=lambda item: item[1].get("created_at") or "",
reverse=True,
)
)
db_dict["files"] = sorted_files
db_dict["status"] = "已连接"
databases.append(db_dict)
return {"databases": databases}
@classmethod
def _add_to_processing_queue(cls, file_id: str) -> None:
"""
将文件添加到处理队列
Args:
file_id: 文件ID
"""
with cls._processing_lock:
cls._processing_files.add(file_id)
logger.debug(f"Added file {file_id} to processing queue")
@classmethod
def _remove_from_processing_queue(cls, file_id: str) -> None:
"""
从处理队列中移除文件
Args:
file_id: 文件ID
"""
with cls._processing_lock:
cls._processing_files.discard(file_id)
logger.debug(f"Removed file {file_id} from processing queue")
@classmethod
def _is_file_in_processing_queue(cls, file_id: str) -> bool:
"""
检查文件是否在处理队列中
Args:
file_id: 文件ID
Returns:
bool: 文件是否在处理队列中
"""
with cls._processing_lock:
return file_id in cls._processing_files
def _check_and_fix_processing_status(self, db_id: str) -> None:
"""
检查并修复异常的处理中状态
如果文件状态为处理中但实际不在处理队列中则修改为相应的错误状态
Args:
db_id: 数据库ID
"""
try:
status_changed = False
# 定义需要检查的中间状态及其对应的错误状态
intermediate_states = {
FileStatus.PARSING: FileStatus.ERROR_PARSING,
FileStatus.INDEXING: FileStatus.ERROR_INDEXING,
}
# 检查该数据库下所有中间状态的文件
for file_id, file_info in self.files_meta.items():
if file_info.get("database_id") == db_id:
current_status = file_info.get("status")
if current_status in intermediate_states:
# 检查文件是否真的在处理队列中
if not self._is_file_in_processing_queue(file_id):
error_status = intermediate_states[current_status]
logger.warning(
f"File {file_id} has {current_status} status but is not in processing queue, "
f"marking as {error_status}"
)
self.files_meta[file_id]["status"] = error_status
self.files_meta[file_id]["error"] = (
f"{current_status.capitalize()} interrupted - process not found in queue"
)
self.files_meta[file_id]["updated_at"] = utc_isoformat()
status_changed = True
# 如果有状态变更,保存元数据
if status_changed:
logger.info(f"Fixed interrupted processing status for database {db_id}")
except Exception as e:
logger.error(f"Error checking processing status for database {db_id}: {e}")
async def delete_folder(self, db_id: str, folder_id: str) -> None:
"""
Recursively delete a folder and its content.
Args:
db_id: Database ID
folder_id: Folder ID to delete
"""
# Find all children
children = [
fid
for fid, meta in self.files_meta.items()
if meta.get("database_id") == db_id and meta.get("parent_id") == folder_id
]
for child_id in children:
child_meta = self.files_meta.get(child_id)
if child_meta and child_meta.get("is_folder"):
await self.delete_folder(db_id, child_id)
else:
await self.delete_file(db_id, child_id)
# Delete the folder itself
# We call delete_file which should handle the actual removal.
# Implementations should ensure they handle folder deletion gracefully (e.g. skip vector deletion)
await self.delete_file(db_id, folder_id)
async def move_file(self, db_id: str, file_id: str, new_parent_id: str | None) -> dict:
"""
Move a file or folder to a new parent folder.
Args:
db_id: Database ID
file_id: File/Folder ID to move
new_parent_id: New parent folder ID (None for root)
Returns:
dict: Updated metadata
"""
if file_id not in self.files_meta:
raise ValueError(f"File {file_id} not found")
meta = self.files_meta[file_id]
if meta.get("database_id") != db_id:
raise ValueError(f"File {file_id} does not belong to database {db_id}")
# Basic cycle detection for folders
if meta.get("is_folder") and new_parent_id:
# Check if new_parent_id is a child of file_id (or is file_id itself)
if new_parent_id == file_id:
raise ValueError("Cannot move a folder into itself")
# Walk up the tree from new_parent_id
current = new_parent_id
while current:
parent_meta = self.files_meta.get(current)
if not parent_meta:
break # Should not happen if integrity is maintained
if current == file_id:
raise ValueError("Cannot move a folder into its own subfolder")
current = parent_meta.get("parent_id")
meta["parent_id"] = new_parent_id
await self._persist_file(file_id)
return meta
@abstractmethod
async def delete_file(self, db_id: str, file_id: str) -> None:
"""
删除文件
Args:
db_id: 数据库ID
file_id: 文件ID
"""
pass
@abstractmethod
async def get_file_basic_info(self, db_id: str, file_id: str) -> dict:
"""
获取文件基本信息仅元数据
Args:
db_id: 数据库ID
file_id: 文件ID
Returns:
dict: 包含文件基本信息的字典
"""
pass
@abstractmethod
async def get_file_content(self, db_id: str, file_id: str) -> dict:
"""
获取文件内容信息chunks和lines
Args:
db_id: 数据库ID
file_id: 文件ID
Returns:
dict: 包含文件内容信息的字典
"""
pass
@abstractmethod
async def get_file_info(self, db_id: str, file_id: str) -> dict:
"""
获取文件完整信息基本信息+内容信息
Args:
db_id: 数据库ID
file_id: 文件ID
Returns:
dict: 包含文件信息和chunks的字典
"""
pass
def update_database(
self,
db_id: str,
name: str,
description: str,
llm_model_spec: str | None = None,
update_llm_model_spec: bool = False,
) -> dict:
"""
更新数据库
Args:
db_id: 数据库ID
name: 新名称
description: 新描述
llm_model_spec: LLM 模型 spec可选
Returns:
更新后的数据库信息
"""
if db_id not in self.databases_meta:
raise ValueError(f"数据库 {db_id} 不存在")
self.databases_meta[db_id]["name"] = name
self.databases_meta[db_id]["description"] = description
if update_llm_model_spec:
self.databases_meta[db_id]["llm_model_spec"] = llm_model_spec
return self.get_database_info(db_id)
def get_retrievers(self) -> dict[str, dict]:
"""
获取所有检索器
Returns:
检索器字典
"""
retrievers = {}
for db_id, meta in self.databases_meta.items():
def make_retriever(db_id):
async def retriever(query_text, **kwargs):
results = await self.aquery(query_text, db_id, agent_call=True, **kwargs)
return self.build_search_output(db_id, results)
return retriever
retrievers[db_id] = {
"name": meta["name"],
"description": meta["description"],
"retriever": make_retriever(db_id),
"metadata": meta,
}
return retrievers
async def _load_metadata(self) -> None:
from yuxi.repositories.evaluation_repository import EvaluationRepository
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
eval_repo = EvaluationRepository()
databases = [kb for kb in await kb_repo.get_all() if kb.kb_type == self.kb_type]
self.databases_meta = {
kb.db_id: {
"name": kb.name,
"description": kb.description,
"kb_type": kb.kb_type,
"embedding_model_spec": kb.embedding_model_spec,
"llm_model_spec": kb.llm_model_spec,
"query_params": kb.query_params or self._get_default_query_params(kb.db_id),
"metadata": self.normalize_additional_params(kb.additional_params),
"created_at": utc_isoformat(kb.created_at) if kb.created_at else utc_isoformat(),
}
for kb in databases
}
self.files_meta = {}
for kb in databases:
kb_additional_params = self.databases_meta.get(kb.db_id, {}).get("metadata") or {}
for record in await file_repo.list_by_db_id(kb.db_id):
self.files_meta[record.file_id] = {
"file_id": record.file_id,
"database_id": record.db_id,
"parent_id": record.parent_id,
"filename": record.filename,
"file_type": record.file_type,
"path": record.path,
"markdown_file": record.markdown_file,
"status": record.status,
"content_hash": record.content_hash,
"size": record.file_size,
"content_type": record.content_type,
"processing_params": sanitize_processing_params(
resolve_processing_params(
kb_additional_params=kb_additional_params,
file_processing_params=record.processing_params,
)
),
"is_folder": record.is_folder,
"error": record.error_message,
"created_by": record.created_by,
"updated_by": record.updated_by,
"created_at": utc_isoformat(record.created_at) if record.created_at else None,
"updated_at": utc_isoformat(record.updated_at) if record.updated_at else None,
"original_filename": record.original_filename,
"minio_url": record.minio_url,
}
self.benchmarks_meta = {}
for kb in databases:
benchmarks = await eval_repo.list_benchmarks(kb.db_id)
if not benchmarks:
continue
self.benchmarks_meta[kb.db_id] = {}
for bench in benchmarks:
self.benchmarks_meta[kb.db_id][bench.benchmark_id] = {
"id": bench.benchmark_id,
"benchmark_id": bench.benchmark_id,
"name": bench.name,
"description": bench.description,
"db_id": bench.db_id,
"question_count": bench.question_count,
"has_gold_chunks": bench.has_gold_chunks,
"has_gold_answers": bench.has_gold_answers,
"benchmark_file": bench.data_file_path,
"created_by": bench.created_by,
"created_at": utc_isoformat(bench.created_at) if bench.created_at else None,
"updated_at": utc_isoformat(bench.updated_at) if bench.updated_at else None,
}
logger.info(f"Loaded {self.kb_type} metadata from database for {len(self.databases_meta)} databases")
await self._fill_missing_file_sizes()
async def _fill_missing_file_sizes(self) -> None:
"""为缺少 size 的已有文件从 MinIO 补全大小信息"""
from yuxi.knowledge.utils.kb_utils import is_minio_url, parse_minio_url
from yuxi.storage.minio import get_minio_client
files_to_update: list[str] = []
for file_id, meta in self.files_meta.items():
if meta.get("is_folder"):
continue
if meta.get("size") is not None:
continue
file_path = meta.get("minio_url") or meta.get("path")
if not file_path or not is_minio_url(file_path):
continue
files_to_update.append(file_id)
if not files_to_update:
return
minio_client = get_minio_client()
async def _stat_file(file_id: str, file_path: str) -> tuple[str, int | None]:
bucket_name, obj_name = parse_minio_url(file_path)
try:
return file_id, await minio_client.astat_file(bucket_name, obj_name)
except Exception as exc:
logger.warning(f"Failed to fill size for {file_id}: {exc}")
return file_id, None
results = await asyncio.gather(
*(
_stat_file(fid, self.files_meta[fid].get("minio_url") or self.files_meta[fid].get("path"))
for fid in files_to_update
)
)
updated = 0
for file_id, file_size in results:
if file_size is not None:
self.files_meta[file_id]["size"] = file_size
updated += 1
if updated:
logger.info(f"Filled {updated}/{len(files_to_update)} missing file sizes from MinIO for {self.kb_type}")
for file_id, file_size in results:
if file_size is not None:
await self._persist_file(file_id)
async def _save_metadata(self) -> None:
from yuxi.repositories.evaluation_repository import EvaluationRepository
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
eval_repo = EvaluationRepository()
self._normalize_metadata_state()
for db_id, meta in self.databases_meta.items():
existing = await kb_repo.get_by_id(db_id)
payload = {
"db_id": db_id,
"name": meta.get("name") or db_id,
"description": meta.get("description"),
"kb_type": meta.get("kb_type") or self.kb_type,
"embedding_model_spec": meta.get("embedding_model_spec"),
"llm_model_spec": meta.get("llm_model_spec"),
"query_params": meta.get("query_params"),
"additional_params": meta.get("metadata") or {},
}
if existing is None:
await kb_repo.create(payload)
for file_id, meta in self.files_meta.items():
db_id = meta.get("database_id")
if not db_id:
continue
await file_repo.upsert(
file_id=file_id,
data={
"db_id": db_id,
"parent_id": meta.get("parent_id"),
"filename": meta.get("filename") or "",
"original_filename": meta.get("original_filename"),
"file_type": meta.get("file_type"),
"path": meta.get("path"),
"minio_url": meta.get("minio_url"),
"markdown_file": meta.get("markdown_file"),
"status": meta.get("status"),
"content_hash": meta.get("content_hash"),
"file_size": meta.get("size"),
"content_type": meta.get("content_type"),
"processing_params": sanitize_processing_params(meta.get("processing_params")),
"is_folder": meta.get("is_folder", False),
"error_message": meta.get("error"),
"created_by": str(meta.get("created_by")) if meta.get("created_by") else None,
"updated_by": str(meta.get("updated_by")) if meta.get("updated_by") else None,
},
)
for db_id, benchmarks in self.benchmarks_meta.items():
for benchmark_id, meta in benchmarks.items():
existing = await eval_repo.get_benchmark(benchmark_id)
payload = {
"benchmark_id": benchmark_id,
"db_id": db_id,
"name": meta.get("name") or benchmark_id,
"description": meta.get("description"),
"question_count": int(meta.get("question_count") or 0),
"has_gold_chunks": bool(meta.get("has_gold_chunks")),
"has_gold_answers": bool(meta.get("has_gold_answers")),
"data_file_path": meta.get("benchmark_file"),
"created_by": str(meta.get("created_by")) if meta.get("created_by") else None,
}
if existing is None:
await eval_repo.create_benchmark(payload)
async def _persist_file(self, file_id: str) -> None:
"""只保存单个文件到数据库,避免全量遍历"""
from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository
file_repo = KnowledgeFileRepository()
if file_id not in self.files_meta:
return
meta = self.files_meta[file_id]
db_id = meta.get("database_id")
if not db_id:
return
await file_repo.upsert(
file_id=file_id,
data={
"db_id": db_id,
"parent_id": meta.get("parent_id"),
"filename": meta.get("filename") or "",
"original_filename": meta.get("original_filename"),
"file_type": meta.get("file_type"),
"path": meta.get("path"),
"minio_url": meta.get("minio_url"),
"markdown_file": meta.get("markdown_file"),
"status": meta.get("status"),
"content_hash": meta.get("content_hash"),
"file_size": meta.get("size"),
"content_type": meta.get("content_type"),
"processing_params": sanitize_processing_params(meta.get("processing_params")),
"is_folder": meta.get("is_folder", False),
"error_message": meta.get("error"),
"created_by": str(meta.get("created_by")) if meta.get("created_by") else None,
"updated_by": str(meta.get("updated_by")) if meta.get("updated_by") else None,
},
)
async def _persist_kb(self, db_id: str) -> None:
"""只保存单个知识库到数据库,避免全量遍历"""
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
kb_repo = KnowledgeBaseRepository()
if db_id not in self.databases_meta:
return
meta = self.databases_meta[db_id]
existing = await kb_repo.get_by_id(db_id)
payload = {
"db_id": db_id,
"name": meta.get("name") or db_id,
"description": meta.get("description"),
"kb_type": meta.get("kb_type") or self.kb_type,
"embedding_model_spec": meta.get("embedding_model_spec"),
"llm_model_spec": meta.get("llm_model_spec"),
"query_params": meta.get("query_params"),
"additional_params": meta.get("metadata") or {},
}
if existing is None:
await kb_repo.create(payload)
else:
await kb_repo.update(
db_id,
{
"name": payload["name"],
"description": payload["description"],
"kb_type": payload["kb_type"],
"embedding_model_spec": payload["embedding_model_spec"],
"llm_model_spec": payload["llm_model_spec"],
"query_params": payload["query_params"],
"additional_params": payload["additional_params"],
},
)