refactor(kb): 知识库模块重构,添加只读连接器支持
- 新增 ReadOnlyConnectors 基类,支持只读外部检索 - 添加知识库类型特定参数配置方法 - 重构 dify 知识库实现 - 优化 additional_params 处理逻辑
This commit is contained in:
parent
96dfaf619a
commit
4181604a71
@ -41,6 +41,10 @@ class KBOperationError(KnowledgeBaseException):
|
||||
class KnowledgeBase(ABC):
|
||||
"""知识库抽象基类,定义统一接口"""
|
||||
|
||||
requires_embedding_model = True
|
||||
supports_documents = True
|
||||
apply_chunk_defaults = True
|
||||
|
||||
# 类级别的处理队列,跟踪所有正在处理的文件
|
||||
_processing_files = set()
|
||||
_processing_lock = None
|
||||
@ -76,7 +80,7 @@ class KnowledgeBase(ABC):
|
||||
self.databases_meta = {}
|
||||
for db_id, meta in global_databases_meta.items():
|
||||
if meta.get("kb_type") == self.kb_type:
|
||||
normalized_additional_params = ensure_chunk_defaults_in_additional_params(meta.get("additional_params"))
|
||||
normalized_additional_params = self.normalize_additional_params(meta.get("additional_params"))
|
||||
self.databases_meta[db_id] = {
|
||||
"name": meta.get("name"),
|
||||
"description": meta.get("description"),
|
||||
@ -160,6 +164,24 @@ class KnowledgeBase(ABC):
|
||||
"""知识库类型标识"""
|
||||
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:
|
||||
"""
|
||||
@ -708,7 +730,7 @@ class KnowledgeBase(ABC):
|
||||
"""
|
||||
from yuxi.utils import hashstr
|
||||
|
||||
kwargs = ensure_chunk_defaults_in_additional_params(kwargs)
|
||||
kwargs = self.normalize_additional_params(kwargs)
|
||||
|
||||
db_id = f"kb_{hashstr(database_name, with_salt=True, length=32)}"
|
||||
|
||||
@ -1286,7 +1308,7 @@ class KnowledgeBase(ABC):
|
||||
"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": ensure_chunk_defaults_in_additional_params(kb.additional_params),
|
||||
"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
|
||||
|
||||
@ -77,9 +77,28 @@ class KnowledgeBaseFactory:
|
||||
"class_name": kb_class.__name__,
|
||||
"description": kb_class.__doc__ or "",
|
||||
"default_config": cls._default_configs[kb_type],
|
||||
"requires_embedding_model": kb_class.requires_embedding_model,
|
||||
"supports_documents": kb_class.supports_documents,
|
||||
"create_params": kb_class.get_create_params_config(),
|
||||
}
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def get_kb_class(cls, kb_type: str) -> type[KnowledgeBase]:
|
||||
"""
|
||||
获取指定类型的知识库类。
|
||||
|
||||
Args:
|
||||
kb_type: 知识库类型
|
||||
|
||||
Returns:
|
||||
知识库类
|
||||
"""
|
||||
if kb_type not in cls._kb_types:
|
||||
available_types = list(cls._kb_types.keys())
|
||||
raise KBNotFoundError(f"Unknown knowledge base type: {kb_type}. Available types: {available_types}")
|
||||
return cls._kb_types[kb_type]
|
||||
|
||||
@classmethod
|
||||
def is_type_supported(cls, kb_type: str) -> bool:
|
||||
"""
|
||||
|
||||
@ -7,5 +7,6 @@
|
||||
|
||||
from .dify import DifyKB
|
||||
from .milvus import MilvusKB
|
||||
from .read_only_connectors import ReadOnlyConnectors
|
||||
|
||||
__all__ = ["MilvusKB", "DifyKB"]
|
||||
__all__ = ["MilvusKB", "DifyKB", "ReadOnlyConnectors"]
|
||||
|
||||
@ -3,11 +3,13 @@ from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from yuxi.knowledge.base import KnowledgeBase
|
||||
from yuxi.knowledge.implementations.read_only_connectors import ReadOnlyConnectors
|
||||
from yuxi.utils import logger
|
||||
|
||||
DIFY_REQUIRED_PARAMS = ("dify_api_url", "dify_token", "dify_dataset_id")
|
||||
|
||||
class DifyKB(KnowledgeBase):
|
||||
|
||||
class DifyKB(ReadOnlyConnectors):
|
||||
"""基于 Dify Dataset Retrieve API 的只读检索知识库实现"""
|
||||
|
||||
def __init__(self, work_dir: str, **kwargs):
|
||||
@ -18,75 +20,48 @@ class DifyKB(KnowledgeBase):
|
||||
def kb_type(self) -> str:
|
||||
return "dify"
|
||||
|
||||
async def _create_kb_instance(self, db_id: str, config: dict) -> Any:
|
||||
return None
|
||||
@classmethod
|
||||
def get_create_params_config(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"options": [
|
||||
{
|
||||
"key": "dify_api_url",
|
||||
"label": "Dify API URL",
|
||||
"type": "text",
|
||||
"required": True,
|
||||
"placeholder": "例如: https://api.dify.ai/v1",
|
||||
"description": "Dify API 地址,必须以 /v1 结尾",
|
||||
},
|
||||
{
|
||||
"key": "dify_token",
|
||||
"label": "Dify Token",
|
||||
"type": "password",
|
||||
"required": True,
|
||||
"placeholder": "请输入 Dify API Token",
|
||||
},
|
||||
{
|
||||
"key": "dify_dataset_id",
|
||||
"label": "Dataset ID",
|
||||
"type": "text",
|
||||
"required": True,
|
||||
"placeholder": "请输入 Dify dataset_id",
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
async def _initialize_kb_instance(self, instance: Any) -> None:
|
||||
return None
|
||||
@classmethod
|
||||
def validate_additional_params(cls, additional_params: dict | None) -> dict:
|
||||
params = dict(additional_params or {})
|
||||
missing_fields = [field for field in DIFY_REQUIRED_PARAMS if not str(params.get(field) or "").strip()]
|
||||
if missing_fields:
|
||||
raise ValueError(f"Dify 参数缺失: {', '.join(missing_fields)}")
|
||||
|
||||
@staticmethod
|
||||
def _readonly_error() -> ValueError:
|
||||
return ValueError("Dify 知识库为只读检索类型,不支持该操作")
|
||||
|
||||
async def add_file_record(
|
||||
self, db_id: str, item: str, params: dict | None = None, operator_id: str | None = None
|
||||
) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def parse_file(self, db_id: str, file_id: str, operator_id: str | None = None) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def update_file_params(self, db_id: str, file_id: str, params: dict, operator_id: str | None = None) -> None:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def create_folder(self, db_id: str, folder_name: str, parent_id: str | None = None) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def move_file(self, db_id: str, file_id: str, new_parent_id: str | None) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def delete_folder(self, db_id: str, folder_id: str) -> None:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def index_file(self, db_id: str, file_id: str, operator_id: str | None = None) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def update_content(self, db_id: str, file_ids: list[str], params: dict | None = None) -> list[dict]:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def delete_file(self, db_id: str, file_id: str) -> None:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def get_file_basic_info(self, db_id: str, file_id: str) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def get_file_content(self, db_id: str, file_id: str) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def open_file_content(self, db_id: str, file_id: str, offset: int = 0, limit: int = 800) -> dict:
|
||||
del offset, limit
|
||||
raise self._readonly_error()
|
||||
|
||||
async def get_file_info(self, db_id: str, file_id: str) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def list_file_tree(
|
||||
self,
|
||||
db_id: str,
|
||||
parent_id: str | None = None,
|
||||
recursive: bool = False,
|
||||
files_only: bool = False,
|
||||
) -> dict:
|
||||
del db_id, parent_id, recursive, files_only
|
||||
raise ValueError("Dify 知识库不支持文件树预览")
|
||||
|
||||
async def read_file_preview(self, db_id: str, file_id: str, variant: str = "parsed") -> dict:
|
||||
del db_id, file_id, variant
|
||||
raise ValueError("Dify 知识库不支持文件预览")
|
||||
|
||||
async def get_file_download(self, db_id: str, file_id: str, variant: str = "original") -> dict:
|
||||
del db_id, file_id, variant
|
||||
raise ValueError("Dify 知识库不支持文件下载")
|
||||
params["dify_api_url"] = str(params.get("dify_api_url") or "").strip()
|
||||
params["dify_token"] = str(params.get("dify_token") or "").strip()
|
||||
params["dify_dataset_id"] = str(params.get("dify_dataset_id") or "").strip()
|
||||
if not params["dify_api_url"].endswith("/v1"):
|
||||
raise ValueError("Dify api_url 必须以 /v1 结尾")
|
||||
return params
|
||||
|
||||
async def aquery(self, query_text: str, db_id: str, agent_call: bool = False, **kwargs) -> list[dict]:
|
||||
del agent_call
|
||||
|
||||
@ -0,0 +1,92 @@
|
||||
from typing import Any
|
||||
|
||||
from yuxi.knowledge.base import KnowledgeBase
|
||||
|
||||
|
||||
class ReadOnlyConnectors(KnowledgeBase):
|
||||
"""只读外部检索连接器基类。
|
||||
|
||||
这类知识库只负责保存连接参数和执行 Query,不承载文档上传、解析、索引和文件预览能力。
|
||||
"""
|
||||
|
||||
requires_embedding_model = False
|
||||
supports_documents = False
|
||||
apply_chunk_defaults = False
|
||||
|
||||
@staticmethod
|
||||
def _readonly_error() -> ValueError:
|
||||
return ValueError("只读检索连接器不支持该操作")
|
||||
|
||||
async def _create_kb_instance(self, db_id: str, config: dict) -> Any:
|
||||
del db_id, config
|
||||
return None
|
||||
|
||||
async def _initialize_kb_instance(self, instance: Any) -> None:
|
||||
del instance
|
||||
return None
|
||||
|
||||
async def add_file_record(
|
||||
self, db_id: str, item: str, params: dict | None = None, operator_id: str | None = None
|
||||
) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def parse_file(self, db_id: str, file_id: str, operator_id: str | None = None) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def update_file_params(self, db_id: str, file_id: str, params: dict, operator_id: str | None = None) -> None:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def create_folder(self, db_id: str, folder_name: str, parent_id: str | None = None) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def move_file(self, db_id: str, file_id: str, new_parent_id: str | None) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def delete_folder(self, db_id: str, folder_id: str) -> None:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def index_file(
|
||||
self,
|
||||
db_id: str,
|
||||
file_id: str,
|
||||
operator_id: str | None = None,
|
||||
params: dict | None = None,
|
||||
) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def update_content(self, db_id: str, file_ids: list[str], params: dict | None = None) -> list[dict]:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def delete_file(self, db_id: str, file_id: str) -> None:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def get_file_basic_info(self, db_id: str, file_id: str) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def get_file_content(self, db_id: str, file_id: str) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def open_file_content(self, db_id: str, file_id: str, offset: int = 0, limit: int = 800) -> dict:
|
||||
del offset, limit
|
||||
raise self._readonly_error()
|
||||
|
||||
async def get_file_info(self, db_id: str, file_id: str) -> dict:
|
||||
raise self._readonly_error()
|
||||
|
||||
async def list_file_tree(
|
||||
self,
|
||||
db_id: str,
|
||||
parent_id: str | None = None,
|
||||
recursive: bool = False,
|
||||
files_only: bool = False,
|
||||
) -> dict:
|
||||
del db_id, parent_id, recursive, files_only
|
||||
raise ValueError("只读检索连接器不支持文件树预览")
|
||||
|
||||
async def read_file_preview(self, db_id: str, file_id: str, variant: str = "parsed") -> dict:
|
||||
del db_id, file_id, variant
|
||||
raise ValueError("只读检索连接器不支持文件预览")
|
||||
|
||||
async def get_file_download(self, db_id: str, file_id: str, variant: str = "original") -> dict:
|
||||
del db_id, file_id, variant
|
||||
raise ValueError("只读检索连接器不支持文件下载")
|
||||
@ -2,10 +2,7 @@ import asyncio
|
||||
import os
|
||||
|
||||
from yuxi.knowledge.base import KBNotFoundError, KnowledgeBase
|
||||
from yuxi.knowledge.chunking.ragflow_like.presets import (
|
||||
deep_merge,
|
||||
ensure_chunk_defaults_in_additional_params,
|
||||
)
|
||||
from yuxi.knowledge.chunking.ragflow_like.presets import deep_merge
|
||||
from yuxi.knowledge.factory import KnowledgeBaseFactory
|
||||
from yuxi.storage.postgres.models_business import User
|
||||
from yuxi.utils import logger
|
||||
@ -182,7 +179,8 @@ class KnowledgeBaseManager:
|
||||
|
||||
# 补充 share_config 和 additional_params
|
||||
db_info["share_config"] = row.share_config or {"is_shared": True, "accessible_departments": []}
|
||||
db_info["additional_params"] = ensure_chunk_defaults_in_additional_params(row.additional_params)
|
||||
db_info["additional_params"] = kb_instance.normalize_additional_params(row.additional_params)
|
||||
db_info["created_by"] = row.created_by
|
||||
all_databases.append(db_info)
|
||||
return {"databases": all_databases}
|
||||
|
||||
@ -317,6 +315,7 @@ class KnowledgeBaseManager:
|
||||
embedding_model_spec: str | None = None,
|
||||
llm_model_spec: str | None = None,
|
||||
share_config: dict | None = None,
|
||||
created_by: str | None = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
"""
|
||||
@ -329,6 +328,7 @@ class KnowledgeBaseManager:
|
||||
embedding_model_spec: 嵌入模型 spec
|
||||
llm_model_spec: LLM 模型 spec
|
||||
share_config: 共享配置
|
||||
created_by: 创建者 uid
|
||||
**kwargs: 其他配置参数
|
||||
|
||||
Returns:
|
||||
@ -346,9 +346,8 @@ class KnowledgeBaseManager:
|
||||
if share_config is None:
|
||||
share_config = {"is_shared": True, "accessible_departments": []}
|
||||
|
||||
kwargs = ensure_chunk_defaults_in_additional_params(kwargs)
|
||||
|
||||
kb_instance = self._get_or_create_kb_instance(kb_type)
|
||||
kwargs = kb_instance.normalize_additional_params(kwargs)
|
||||
db_info = await kb_instance.create_database(
|
||||
database_name,
|
||||
description,
|
||||
@ -361,7 +360,7 @@ class KnowledgeBaseManager:
|
||||
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
|
||||
|
||||
kb_repo = KnowledgeBaseRepository()
|
||||
updated = await kb_repo.update(db_id, {"share_config": share_config})
|
||||
updated = await kb_repo.update(db_id, {"share_config": share_config, "created_by": created_by})
|
||||
if updated is None:
|
||||
await kb_repo.create(
|
||||
{
|
||||
@ -373,6 +372,7 @@ class KnowledgeBaseManager:
|
||||
"llm_model_spec": db_info.get("llm_model_spec"),
|
||||
"additional_params": kwargs.copy(),
|
||||
"share_config": share_config,
|
||||
"created_by": created_by,
|
||||
}
|
||||
)
|
||||
|
||||
@ -455,7 +455,7 @@ class KnowledgeBaseManager:
|
||||
}
|
||||
|
||||
# 添加数据库中的附加字段
|
||||
db_info["additional_params"] = ensure_chunk_defaults_in_additional_params(kb.additional_params)
|
||||
db_info["additional_params"] = kb_instance.normalize_additional_params(kb.additional_params)
|
||||
db_info["share_config"] = kb.share_config or {"is_shared": True, "accessible_departments": []}
|
||||
db_info["mindmap"] = kb.mindmap
|
||||
db_info["sample_questions"] = kb.sample_questions or []
|
||||
@ -637,7 +637,7 @@ class KnowledgeBaseManager:
|
||||
if current_graph_config.get("locked") and "graph_build_config" in additional_params:
|
||||
raise ValueError("图谱抽取配置已锁定,请使用图谱重置接口重新配置")
|
||||
|
||||
merged_additional_params = ensure_chunk_defaults_in_additional_params(
|
||||
merged_additional_params = kb_instance.normalize_additional_params(
|
||||
deep_merge(current_additional_params, additional_params)
|
||||
)
|
||||
update_data["additional_params"] = merged_additional_params
|
||||
|
||||
Loading…
Reference in New Issue
Block a user