feat: 重构知识库注册机制,简化知识库类型注册并添加描述信息

This commit is contained in:
Wenjie Zhang 2026-05-18 18:07:09 +08:00
parent 213a542a87
commit 5693fdd4a7
8 changed files with 57 additions and 68 deletions

View File

@ -12,10 +12,10 @@ _SKIP_APP_INIT = os.environ.get("YUXI_SKIP_APP_INIT") == "1"
if not _LITE_MODE: if not _LITE_MODE:
# 注册知识库类型 # 注册知识库类型
KnowledgeBaseFactory.register("milvus", MilvusKB, {"description": "基于 Milvus 的生产级向量知识库,适合高性能部署"}) KnowledgeBaseFactory.register(MilvusKB)
KnowledgeBaseFactory.register("dify", DifyKB, {"description": "连接 Dify Dataset 的只读检索知识库"}) KnowledgeBaseFactory.register(DifyKB)
KnowledgeBaseFactory.register("notion", NotionKB, {"description": "连接 Notion Data Source 的只读知识库,支持检索、打开页面和页内查找"}) KnowledgeBaseFactory.register(NotionKB)
# 创建知识库管理器 # 创建知识库管理器
work_dir = os.path.join(config.save_dir, "knowledge_base_data") work_dir = os.path.join(config.save_dir, "knowledge_base_data")

View File

@ -43,6 +43,9 @@ class KBOperationError(KnowledgeBaseException):
class KnowledgeBase(ABC): class KnowledgeBase(ABC):
"""知识库抽象基类,定义统一接口""" """知识库抽象基类,定义统一接口"""
kb_type = ""
name = ""
description = ""
requires_embedding_model = True requires_embedding_model = True
supports_documents = True supports_documents = True
apply_chunk_defaults = True apply_chunk_defaults = True
@ -160,12 +163,6 @@ class KnowledgeBase(ABC):
if normalized: if normalized:
b["updated_at"] = normalized b["updated_at"] = normalized
@property
@abstractmethod
def kb_type(self) -> str:
"""知识库类型标识"""
pass
@classmethod @classmethod
def get_create_params_config(cls) -> dict[str, Any]: def get_create_params_config(cls) -> dict[str, Any]:
"""获取创建知识库时的类型特定参数配置。""" """获取创建知识库时的类型特定参数配置。"""

View File

@ -8,25 +8,21 @@ class KnowledgeBaseFactory:
# 注册的知识库类型映射 {kb_type: kb_class} # 注册的知识库类型映射 {kb_type: kb_class}
_kb_types: dict[str, type[KnowledgeBase]] = {} _kb_types: dict[str, type[KnowledgeBase]] = {}
# 每种类型的默认配置
_default_configs: dict[str, dict] = {}
@classmethod @classmethod
def register(cls, kb_type: str, kb_class: type[KnowledgeBase], default_config: dict = None): def register(cls, kb_class: type[KnowledgeBase]):
""" """
注册知识库类型 注册知识库类型
Args: Args:
kb_type: 知识库类型标识
kb_class: 知识库类 kb_class: 知识库类
default_config: 默认配置
""" """
if not issubclass(kb_class, KnowledgeBase): if not issubclass(kb_class, KnowledgeBase):
raise ValueError("Knowledge base class must inherit from KnowledgeBase") raise ValueError("Knowledge base class must inherit from KnowledgeBase")
if not kb_class.kb_type:
raise ValueError("Knowledge base class must define kb_type")
cls._kb_types[kb_type] = kb_class cls._kb_types[kb_class.kb_type] = kb_class
cls._default_configs[kb_type] = default_config or {} # logger.info(f"Registered knowledge base type: {kb_class.kb_type}")
# logger.info(f"Registered knowledge base type: {kb_type}")
@classmethod @classmethod
def create(cls, kb_type: str, work_dir: str, **kwargs) -> KnowledgeBase: def create(cls, kb_type: str, work_dir: str, **kwargs) -> KnowledgeBase:
@ -50,13 +46,9 @@ class KnowledgeBaseFactory:
kb_class = cls._kb_types[kb_type] kb_class = cls._kb_types[kb_type]
# 合并默认配置和用户配置
config = cls._default_configs[kb_type].copy()
config.update(kwargs)
try: try:
# 创建实例 # 创建实例
instance = kb_class(work_dir, **config) instance = kb_class(work_dir, **kwargs)
logger.info(f"Created {kb_type} knowledge base instance at {work_dir}") logger.info(f"Created {kb_type} knowledge base instance at {work_dir}")
return instance return instance
except Exception as e: except Exception as e:
@ -74,9 +66,8 @@ class KnowledgeBaseFactory:
result = {} result = {}
for kb_type, kb_class in cls._kb_types.items(): for kb_type, kb_class in cls._kb_types.items():
result[kb_type] = { result[kb_type] = {
"class_name": kb_class.__name__, "name": kb_class.name,
"description": kb_class.__doc__ or "", "description": kb_class.description,
"default_config": cls._default_configs[kb_type],
"requires_embedding_model": kb_class.requires_embedding_model, "requires_embedding_model": kb_class.requires_embedding_model,
"supports_documents": kb_class.supports_documents, "supports_documents": kb_class.supports_documents,
"create_params": kb_class.get_create_params_config(), "create_params": kb_class.get_create_params_config(),
@ -111,16 +102,3 @@ class KnowledgeBaseFactory:
是否支持 是否支持
""" """
return kb_type in cls._kb_types return kb_type in cls._kb_types
@classmethod
def get_default_config(cls, kb_type: str) -> dict:
"""
获取指定类型的默认配置
Args:
kb_type: 知识库类型
Returns:
默认配置字典
"""
return cls._default_configs.get(kb_type, {}).copy()

View File

@ -12,14 +12,14 @@ DIFY_REQUIRED_PARAMS = ("dify_api_url", "dify_token", "dify_dataset_id")
class DifyKB(ReadOnlyConnectors): class DifyKB(ReadOnlyConnectors):
"""基于 Dify Dataset Retrieve API 的只读检索知识库实现""" """基于 Dify Dataset Retrieve API 的只读检索知识库实现"""
kb_type = "dify"
name = "Dify"
description = "连接 Dify Dataset 的只读检索知识库"
def __init__(self, work_dir: str, **kwargs): def __init__(self, work_dir: str, **kwargs):
del kwargs del kwargs
super().__init__(work_dir) super().__init__(work_dir)
@property
def kb_type(self) -> str:
return "dify"
@classmethod @classmethod
def get_create_params_config(cls) -> dict[str, Any]: def get_create_params_config(cls) -> dict[str, Any]:
return { return {

View File

@ -239,6 +239,10 @@ def _retrieval_config_options() -> list[dict[str, Any]]:
class MilvusKB(KnowledgeBase): class MilvusKB(KnowledgeBase):
"""基于 Milvus 的生产级向量库""" """基于 Milvus 的生产级向量库"""
kb_type = "milvus"
name = "Milvus"
description = "基于 Milvus 的生产级向量知识库,适合高性能部署"
def __init__(self, work_dir: str, **kwargs): def __init__(self, work_dir: str, **kwargs):
""" """
初始化 Milvus 知识库 初始化 Milvus 知识库
@ -273,11 +277,6 @@ class MilvusKB(KnowledgeBase):
logger.info("MilvusKB initialized") logger.info("MilvusKB initialized")
@property
def kb_type(self) -> str:
"""知识库类型标识"""
return "milvus"
def _init_connection(self): def _init_connection(self):
"""初始化 Milvus 连接""" """初始化 Milvus 连接"""
try: try:

View File

@ -170,15 +170,15 @@ class _NotionClient:
class NotionKB(ReadOnlyConnectors): class NotionKB(ReadOnlyConnectors):
"""连接 Notion Data Source 的只读知识库实现""" """连接 Notion Data Source 的只读知识库实现"""
kb_type = "notion"
name = "Notion"
description = "连接 Notion Data Source 的只读知识库,支持检索、打开页面和页内查找"
def __init__(self, work_dir: str, **kwargs): def __init__(self, work_dir: str, **kwargs):
del kwargs del kwargs
super().__init__(work_dir) super().__init__(work_dir)
self._page_markdown_cache: dict[tuple[str, str, str, str], tuple[float, str]] = {} self._page_markdown_cache: dict[tuple[str, str, str, str], tuple[float, str]] = {}
@property
def kb_type(self) -> str:
return "notion"
@classmethod @classmethod
def get_create_params_config(cls) -> dict[str, Any]: def get_create_params_config(cls) -> dict[str, Any]:
return { return {

View File

@ -402,6 +402,9 @@ async def test_get_knowledge_base_types(test_client, admin_headers):
payload = response.json() payload = response.json()
assert payload["message"] == "success" assert payload["message"] == "success"
assert "kb_types" in payload assert "kb_types" in payload
assert "default_config" not in payload["kb_types"]["dify"]
assert payload["kb_types"]["dify"]["name"] == "Dify"
assert payload["kb_types"]["dify"]["description"] == "连接 Dify Dataset 的只读检索知识库"
assert payload["kb_types"]["dify"]["requires_embedding_model"] is False assert payload["kb_types"]["dify"]["requires_embedding_model"] is False
assert payload["kb_types"]["dify"]["supports_documents"] is False assert payload["kb_types"]["dify"]["supports_documents"] is False
assert [option["key"] for option in payload["kb_types"]["dify"]["create_params"]["options"]] == [ assert [option["key"] for option in payload["kb_types"]["dify"]["create_params"]["options"]] == [
@ -409,6 +412,12 @@ async def test_get_knowledge_base_types(test_client, admin_headers):
"dify_token", "dify_token",
"dify_dataset_id", "dify_dataset_id",
] ]
assert "default_config" not in payload["kb_types"]["notion"]
assert payload["kb_types"]["notion"]["name"] == "Notion"
assert (
payload["kb_types"]["notion"]["description"]
== "连接 Notion Data Source 的只读知识库,支持检索、打开页面和页内查找"
)
assert payload["kb_types"]["notion"]["requires_embedding_model"] is False assert payload["kb_types"]["notion"]["requires_embedding_model"] is False
assert payload["kb_types"]["notion"]["supports_documents"] is False assert payload["kb_types"]["notion"]["supports_documents"] is False
assert [option["key"] for option in payload["kb_types"]["notion"]["create_params"]["options"]] == [ assert [option["key"] for option in payload["kb_types"]["notion"]["create_params"]["options"]] == [

View File

@ -25,7 +25,7 @@
</a-select> </a-select>
</template> </template>
<template #actions> <template #actions>
<a-button type="primary" @click="state.openNewDatabaseModel = true"> <a-button type="primary" :disabled="!kbTypes.length" @click="state.openNewDatabaseModel = true">
<PlusOutlined /> 新建知识库 <PlusOutlined /> 新建知识库
</a-button> </a-button>
</template> </template>
@ -159,6 +159,7 @@
key="submit" key="submit"
type="primary" type="primary"
:loading="dbState.creating" :loading="dbState.creating"
:disabled="!selectedKbTypeInfo"
@click="handleCreateDatabase" @click="handleCreateDatabase"
>创建</a-button >创建</a-button
> >
@ -175,7 +176,12 @@
<div v-else-if="!databases || databases.length === 0" class="empty-state"> <div v-else-if="!databases || databases.length === 0" class="empty-state">
<h3 class="empty-title">暂无知识库</h3> <h3 class="empty-title">暂无知识库</h3>
<p class="empty-description">创建您的第一个知识库开始管理文档和知识</p> <p class="empty-description">创建您的第一个知识库开始管理文档和知识</p>
<a-button type="primary" size="large" @click="state.openNewDatabaseModel = true"> <a-button
type="primary"
size="large"
:disabled="!kbTypes.length"
@click="state.openNewDatabaseModel = true"
>
<template #icon> <template #icon>
<PlusOutlined /> <PlusOutlined />
</template> </template>
@ -276,7 +282,7 @@ const createEmptyDatabaseForm = () => ({
name: '', name: '',
description: '', description: '',
embedding_model_spec: configStore.config?.embed_model, embedding_model_spec: configStore.config?.embed_model,
kb_type: 'milvus', kb_type: '',
storage: '', storage: '',
chunk_preset_id: 'general', chunk_preset_id: 'general',
additional_params: {} additional_params: {}
@ -300,8 +306,7 @@ const createParamOptions = computed(
() => selectedKbTypeInfo.value?.create_params?.options || [] () => selectedKbTypeInfo.value?.create_params?.options || []
) )
const getKbTypeDescription = (typeInfo) => const getKbTypeDescription = (typeInfo) => typeInfo?.description || ''
typeInfo?.default_config?.description || typeInfo?.description || ''
const resetCreateParamValues = () => { const resetCreateParamValues = () => {
newDatabase.additional_params = {} newDatabase.additional_params = {}
@ -320,25 +325,21 @@ const resetCreateParamValues = () => {
const loadSupportedKbTypes = async () => { const loadSupportedKbTypes = async () => {
try { try {
const data = await typeApi.getKnowledgeBaseTypes() const data = await typeApi.getKnowledgeBaseTypes()
supportedKbTypes.value = data.kb_types supportedKbTypes.value = data.kb_types || {}
newDatabase.kb_type = kbTypes.value[0] || ''
resetCreateParamValues() resetCreateParamValues()
console.log('支持的知识库类型:', supportedKbTypes.value)
} catch (error) { } catch (error) {
console.error('加载知识库类型失败:', error) console.error('加载知识库类型失败:', error)
// supportedKbTypes.value = {}
supportedKbTypes.value = { newDatabase.kb_type = ''
milvus: { resetCreateParamValues()
description: '基于 Milvus 的生产级向量知识库,支持文档检索和图谱构建', message.error('加载知识库类型失败,请稍后重试')
class_name: 'MilvusKB',
requires_embedding_model: true,
create_params: { options: [] }
}
}
} }
} }
const resetNewDatabase = () => { const resetNewDatabase = () => {
Object.assign(newDatabase, createEmptyDatabaseForm()) Object.assign(newDatabase, createEmptyDatabaseForm())
newDatabase.kb_type = kbTypes.value[0] || ''
resetCreateParamValues() resetCreateParamValues()
// //
shareConfig.value = { shareConfig.value = {
@ -430,6 +431,11 @@ const buildRequestData = () => {
// //
const handleCreateDatabase = async () => { const handleCreateDatabase = async () => {
if (!selectedKbTypeInfo.value) {
message.error('知识库类型加载失败,无法创建知识库')
return
}
for (const field of createParamOptions.value) { for (const field of createParamOptions.value) {
if (!field.required) continue if (!field.required) continue
const value = newDatabase.additional_params[field.key] const value = newDatabase.additional_params[field.key]