refactor(knowledge): 将知识库检索参数从 router 移动到对应的实现文件中
This commit is contained in:
parent
c8bdef40d5
commit
98c73a90c9
@ -1,23 +1,23 @@
|
||||
import aiofiles
|
||||
import asyncio
|
||||
import os
|
||||
import traceback
|
||||
import textwrap
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from urllib.parse import quote, unquote
|
||||
|
||||
import aiofiles
|
||||
from fastapi import APIRouter, Body, Depends, File, HTTPException, Query, Request, UploadFile
|
||||
from fastapi.responses import FileResponse
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
from src.storage.db.models import User
|
||||
from server.utils.auth_middleware import get_admin_user
|
||||
from server.services.tasker import TaskContext, tasker
|
||||
from server.utils.auth_middleware import get_admin_user
|
||||
from src import config, knowledge_base
|
||||
from src.knowledge.indexing import SUPPORTED_FILE_EXTENSIONS, is_supported_file_extension, process_file_to_markdown
|
||||
from src.knowledge.utils import calculate_content_hash
|
||||
from src.models.embed import test_embedding_model_status, test_all_embedding_models_status
|
||||
from src.storage.minio.client import aupload_file_to_minio, get_minio_client, StorageError
|
||||
from src.models.embed import test_all_embedding_models_status, test_embedding_model_status
|
||||
from src.storage.db.models import User
|
||||
from src.storage.minio.client import StorageError, aupload_file_to_minio, get_minio_client
|
||||
from src.utils import logger
|
||||
|
||||
knowledge = APIRouter(prefix="/knowledge", tags=["knowledge"])
|
||||
@ -466,11 +466,9 @@ async def index_documents(
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to update params for {file_id}: {e}")
|
||||
param_update_failed.add(file_id)
|
||||
processed_items.append({
|
||||
"file_id": file_id,
|
||||
"status": "failed",
|
||||
"error": f"参数更新失败: {str(e)}"
|
||||
})
|
||||
processed_items.append(
|
||||
{"file_id": file_id, "status": "failed", "error": f"参数更新失败: {str(e)}"}
|
||||
)
|
||||
|
||||
for idx, file_id in enumerate(file_ids, 1):
|
||||
await context.raise_if_cancelled()
|
||||
@ -582,6 +580,7 @@ async def delete_document(db_id: str, doc_id: str, current_user: User = Depends(
|
||||
logger.error(f"删除文档失败 {e}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=400, detail=f"删除文档失败: {e}")
|
||||
|
||||
|
||||
@knowledge.get("/databases/{db_id}/documents/{doc_id}/download")
|
||||
async def download_document(db_id: str, doc_id: str, request: Request, current_user: User = Depends(get_admin_user)):
|
||||
"""下载原始文件 - 根据path类型选择本地或MinIO下载"""
|
||||
@ -786,171 +785,45 @@ async def update_knowledge_base_query_params(
|
||||
async def get_knowledge_base_query_params(db_id: str, current_user: User = Depends(get_admin_user)):
|
||||
"""获取知识库类型特定的查询参数"""
|
||||
try:
|
||||
# 获取数据库信息
|
||||
db_info = knowledge_base.get_database_info(db_id)
|
||||
if not db_info:
|
||||
raise HTTPException(status_code=404, detail="Database not found")
|
||||
# 获取知识库实例
|
||||
kb_instance = knowledge_base._get_kb_for_database(db_id)
|
||||
|
||||
kb_type = db_info.get("kb_type", "lightrag")
|
||||
metadata = db_info.get("metadata", {}) or {}
|
||||
reranker_config = metadata.get("reranker_config", {}) or {}
|
||||
reranker_enabled = bool(reranker_config.get("enabled", False))
|
||||
# 调用知识库实例的方法获取配置
|
||||
params = kb_instance.get_query_params_config(
|
||||
db_id=db_id,
|
||||
reranker_names=config.reranker_names, # 传递动态配置
|
||||
)
|
||||
|
||||
# 根据知识库类型返回不同的查询参数
|
||||
if kb_type == "lightrag":
|
||||
params = {
|
||||
"type": "lightrag",
|
||||
"options": [
|
||||
{
|
||||
"key": "mode",
|
||||
"label": "检索模式",
|
||||
"type": "select",
|
||||
"default": "mix",
|
||||
"options": [
|
||||
{"value": "local", "label": "Local", "description": "上下文相关信息"},
|
||||
{"value": "global", "label": "Global", "description": "全局知识"},
|
||||
{"value": "hybrid", "label": "Hybrid", "description": "本地和全局混合"},
|
||||
{"value": "naive", "label": "Naive", "description": "基本搜索"},
|
||||
{"value": "mix", "label": "Mix", "description": "知识图谱和向量检索混合"},
|
||||
],
|
||||
},
|
||||
{
|
||||
"key": "only_need_context",
|
||||
"label": "只使用上下文",
|
||||
"type": "boolean",
|
||||
"default": True,
|
||||
"description": "只返回上下文,不生成回答",
|
||||
},
|
||||
{
|
||||
"key": "only_need_prompt",
|
||||
"label": "只使用提示",
|
||||
"type": "boolean",
|
||||
"default": False,
|
||||
"description": "只返回提示,不进行检索",
|
||||
},
|
||||
{
|
||||
"key": "top_k",
|
||||
"label": "TopK",
|
||||
"type": "number",
|
||||
"default": 10,
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"description": "返回的最大结果数量",
|
||||
},
|
||||
],
|
||||
}
|
||||
elif kb_type == "milvus":
|
||||
top_k_default = reranker_config.get("final_top_k", 10)
|
||||
params_list = [
|
||||
{
|
||||
"key": "top_k",
|
||||
"label": "TopK",
|
||||
"type": "number",
|
||||
"default": top_k_default,
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"description": "返回的最大结果数量",
|
||||
},
|
||||
{
|
||||
"key": "similarity_threshold",
|
||||
"label": "相似度阈值",
|
||||
"type": "number",
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.1,
|
||||
"description": "过滤相似度低于此值的结果",
|
||||
},
|
||||
{
|
||||
"key": "include_distances",
|
||||
"label": "显示相似度",
|
||||
"type": "boolean",
|
||||
"default": True,
|
||||
"description": "在结果中显示相似度分数",
|
||||
},
|
||||
{
|
||||
"key": "metric_type",
|
||||
"label": "距离度量类型",
|
||||
"type": "select",
|
||||
"default": "COSINE",
|
||||
"options": [
|
||||
{"value": "COSINE", "label": "余弦相似度", "description": "适合文本语义相似度"},
|
||||
{"value": "L2", "label": "欧几里得距离", "description": "适合数值型数据"},
|
||||
{"value": "IP", "label": "内积", "description": "适合标准化向量"},
|
||||
],
|
||||
"description": "向量相似度计算方法",
|
||||
},
|
||||
{
|
||||
"key": "use_reranker",
|
||||
"label": "启用重排序",
|
||||
"type": "boolean",
|
||||
"default": reranker_enabled,
|
||||
"description": "是否使用精排模型对检索结果进行重排序",
|
||||
},
|
||||
{
|
||||
"key": "recall_top_k",
|
||||
"label": "召回数量",
|
||||
"type": "number",
|
||||
"default": reranker_config.get("recall_top_k", 50),
|
||||
"min": 10,
|
||||
"max": 200,
|
||||
"description": "启用重排序时向量检索的候选数量",
|
||||
},
|
||||
]
|
||||
|
||||
if config.reranker_names:
|
||||
params_list.append(
|
||||
{
|
||||
"key": "reranker_model",
|
||||
"label": "重排序模型",
|
||||
"type": "select",
|
||||
"default": reranker_config.get("model", ""),
|
||||
"options": [
|
||||
{"label": info.name, "value": model_id} for model_id, info in config.reranker_names.items()
|
||||
],
|
||||
"description": "覆盖默认配置,选择用于本次查询的重排序模型",
|
||||
}
|
||||
)
|
||||
|
||||
params = {"type": "milvus", "options": params_list}
|
||||
else:
|
||||
# 未知类型,返回基本参数
|
||||
params = {
|
||||
"type": "unknown",
|
||||
"options": [
|
||||
{
|
||||
"key": "top_k",
|
||||
"label": "TopK",
|
||||
"type": "number",
|
||||
"default": 10,
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"description": "返回的最大结果数量",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
# 获取用户保存的配置
|
||||
saved_options = {}
|
||||
try:
|
||||
if db_id in knowledge_base.global_databases_meta:
|
||||
query_params_meta = knowledge_base.global_databases_meta[db_id].get("query_params", {})
|
||||
saved_options = query_params_meta.get("options", {})
|
||||
except Exception as saved_error:
|
||||
logger.warning(f"获取保存的配置失败: {saved_error}")
|
||||
|
||||
# 将保存的值合并到默认配置中
|
||||
# 获取用户保存的配置并合并
|
||||
saved_options = _get_saved_query_options(db_id)
|
||||
if saved_options:
|
||||
for option in params.get("options", []):
|
||||
key = option.get("key")
|
||||
if key in saved_options:
|
||||
option["default"] = saved_options[key]
|
||||
params = _merge_saved_options(params, saved_options)
|
||||
|
||||
return {"params": params, "message": "success"}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取知识库查询参数失败 {e}, {traceback.format_exc()}")
|
||||
return {"message": f"获取知识库查询参数失败 {e}", "params": {}}
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _get_saved_query_options(db_id: str) -> dict:
|
||||
"""获取用户保存的查询参数配置"""
|
||||
try:
|
||||
if db_id in knowledge_base.global_databases_meta:
|
||||
query_params_meta = knowledge_base.global_databases_meta[db_id].get("query_params", {})
|
||||
return query_params_meta.get("options", {})
|
||||
except Exception as e:
|
||||
logger.warning(f"获取保存的配置失败: {e}")
|
||||
return {}
|
||||
|
||||
|
||||
def _merge_saved_options(params: dict, saved_options: dict) -> dict:
|
||||
"""将用户保存的配置合并到默认配置中"""
|
||||
for option in params.get("options", []):
|
||||
key = option.get("key")
|
||||
if key in saved_options:
|
||||
option["default"] = saved_options[key]
|
||||
return params
|
||||
|
||||
|
||||
# =============================================================================
|
||||
@ -1000,9 +873,10 @@ async def generate_sample_questions(
|
||||
生成的问题列表
|
||||
"""
|
||||
try:
|
||||
from src.models import select_model
|
||||
import json
|
||||
|
||||
from src.models import select_model
|
||||
|
||||
# 从请求体中提取参数
|
||||
count = request_body.get("count", 10)
|
||||
|
||||
|
||||
@ -476,6 +476,36 @@ class KnowledgeBase(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_query_params_config(self, db_id: str, **kwargs) -> dict:
|
||||
"""
|
||||
获取知识库类型的查询参数配置
|
||||
|
||||
Args:
|
||||
db_id: 数据库ID
|
||||
**kwargs: 额外参数(如 reranker_names 等)
|
||||
|
||||
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
|
||||
|
||||
|
||||
@ -553,6 +553,49 @@ class LightRagKB(KnowledgeBase):
|
||||
|
||||
return {**basic_info, **content_info}
|
||||
|
||||
def get_query_params_config(self, db_id: str, **kwargs) -> dict:
|
||||
"""获取 LightRAG 知识库的查询参数配置"""
|
||||
options = [
|
||||
{
|
||||
"key": "mode",
|
||||
"label": "检索模式",
|
||||
"type": "select",
|
||||
"default": "mix",
|
||||
"options": [
|
||||
{"value": "local", "label": "Local", "description": "上下文相关信息"},
|
||||
{"value": "global", "label": "Global", "description": "全局知识"},
|
||||
{"value": "hybrid", "label": "Hybrid", "description": "本地和全局混合"},
|
||||
{"value": "naive", "label": "Naive", "description": "基本搜索"},
|
||||
{"value": "mix", "label": "Mix", "description": "知识图谱和向量检索混合"},
|
||||
],
|
||||
},
|
||||
{
|
||||
"key": "only_need_context",
|
||||
"label": "只使用上下文",
|
||||
"type": "boolean",
|
||||
"default": True,
|
||||
"description": "只返回上下文,不生成回答",
|
||||
},
|
||||
{
|
||||
"key": "only_need_prompt",
|
||||
"label": "只使用提示",
|
||||
"type": "boolean",
|
||||
"default": False,
|
||||
"description": "只返回提示,不进行检索",
|
||||
},
|
||||
{
|
||||
"key": "top_k",
|
||||
"label": "TopK",
|
||||
"type": "number",
|
||||
"default": 10,
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"description": "返回的最大结果数量",
|
||||
},
|
||||
]
|
||||
|
||||
return {"type": "lightrag", "options": options}
|
||||
|
||||
async def export_data(self, db_id: str, format: str = "csv", **kwargs) -> str:
|
||||
"""
|
||||
使用 LightRAG 原生功能导出知识库数据。
|
||||
|
||||
@ -689,6 +689,88 @@ class MilvusKB(KnowledgeBase):
|
||||
# Call base method to delete local files and metadata
|
||||
return super().delete_database(db_id)
|
||||
|
||||
def get_query_params_config(self, db_id: str, **kwargs) -> dict:
|
||||
"""获取 Milvus 知识库的查询参数配置"""
|
||||
# 从 metadata 中获取 reranker 配置
|
||||
db_meta = self.databases_meta.get(db_id, {})
|
||||
metadata = db_meta.get("metadata", {}) or {}
|
||||
reranker_config = metadata.get("reranker_config", {}) or {}
|
||||
reranker_enabled = bool(reranker_config.get("enabled", False))
|
||||
|
||||
# 构建 Milvus 特定参数
|
||||
options = [
|
||||
{
|
||||
"key": "top_k",
|
||||
"label": "TopK",
|
||||
"type": "number",
|
||||
"default": reranker_config.get("final_top_k", 10),
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"description": "返回的最大结果数量",
|
||||
},
|
||||
{
|
||||
"key": "similarity_threshold",
|
||||
"label": "相似度阈值",
|
||||
"type": "number",
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.1,
|
||||
"description": "过滤相似度低于此值的结果",
|
||||
},
|
||||
{
|
||||
"key": "include_distances",
|
||||
"label": "显示相似度",
|
||||
"type": "boolean",
|
||||
"default": True,
|
||||
"description": "在结果中显示相似度分数",
|
||||
},
|
||||
{
|
||||
"key": "metric_type",
|
||||
"label": "距离度量类型",
|
||||
"type": "select",
|
||||
"default": "COSINE",
|
||||
"options": [
|
||||
{"value": "COSINE", "label": "余弦相似度", "description": "适合文本语义相似度"},
|
||||
{"value": "L2", "label": "欧几里得距离", "description": "适合数值型数据"},
|
||||
{"value": "IP", "label": "内积", "description": "适合标准化向量"},
|
||||
],
|
||||
"description": "向量相似度计算方法",
|
||||
},
|
||||
{
|
||||
"key": "use_reranker",
|
||||
"label": "启用重排序",
|
||||
"type": "boolean",
|
||||
"default": reranker_enabled,
|
||||
"description": "是否使用精排模型对检索结果进行重排序",
|
||||
},
|
||||
{
|
||||
"key": "recall_top_k",
|
||||
"label": "召回数量",
|
||||
"type": "number",
|
||||
"default": reranker_config.get("recall_top_k", 50),
|
||||
"min": 10,
|
||||
"max": 200,
|
||||
"description": "启用重排序时向量检索的候选数量",
|
||||
},
|
||||
]
|
||||
|
||||
# 动态添加 reranker 模型选择
|
||||
reranker_names = kwargs.get("reranker_names", {})
|
||||
if reranker_names:
|
||||
options.append(
|
||||
{
|
||||
"key": "reranker_model",
|
||||
"label": "重排序模型",
|
||||
"type": "select",
|
||||
"default": reranker_config.get("model", ""),
|
||||
"options": [{"label": info.name, "value": model_id} for model_id, info in reranker_names.items()],
|
||||
"description": "覆盖默认配置,选择用于本次查询的重排序模型",
|
||||
}
|
||||
)
|
||||
|
||||
return {"type": "milvus", "options": options}
|
||||
|
||||
def __del__(self):
|
||||
"""清理连接"""
|
||||
try:
|
||||
|
||||
11
uv.lock
11
uv.lock
@ -702,6 +702,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/ae/3a/dbeec9d1ee0844c679f6bb5d6ad4e9f198b1224f4e7a32825f47f6192b0c/cffi-2.0.0-cp314-cp314t-win_arm64.whl", hash = "sha256:0a1527a803f0a659de1af2e1fd700213caba79377e27e4693648c2923da066f9", size = 184195, upload-time = "2025-09-08T23:23:43.004Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "chardet"
|
||||
version = "5.2.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/f3/0d/f7b6ab21ec75897ed80c17d79b15951a719226b9fababf1e40ea74d69079/chardet-5.2.0.tar.gz", hash = "sha256:1b3b6ff479a8c414bc3fa2c0852995695c4a026dcd6d0633b2dd092ca39c1cf7", size = 2069618, upload-time = "2023-08-01T19:23:02.662Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/38/6f/f5fbc992a329ee4e0f288c1fe0e2ad9485ed064cac731ed2fe47dcc38cbf/chardet-5.2.0-py3-none-any.whl", hash = "sha256:e1cf59446890a00105fe7b7912492ea04b6e6f06d4b742b2c788469e34c82970", size = 199385, upload-time = "2023-08-01T19:23:00.661Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "charset-normalizer"
|
||||
version = "3.4.4"
|
||||
@ -7334,6 +7343,7 @@ dependencies = [
|
||||
{ name = "aiohttp" },
|
||||
{ name = "aiosqlite" },
|
||||
{ name = "asyncpg" },
|
||||
{ name = "chardet" },
|
||||
{ name = "colorlog" },
|
||||
{ name = "dashscope" },
|
||||
{ name = "deepagents" },
|
||||
@ -7408,6 +7418,7 @@ requires-dist = [
|
||||
{ name = "aiohttp", specifier = ">=3.9.0" },
|
||||
{ name = "aiosqlite", specifier = ">=0.20.0" },
|
||||
{ name = "asyncpg", specifier = ">=0.30.0" },
|
||||
{ name = "chardet", specifier = ">=5.0.0" },
|
||||
{ name = "colorlog", specifier = ">=6.9.0" },
|
||||
{ name = "dashscope", specifier = ">=1.23.2" },
|
||||
{ name = "deepagents", specifier = ">=0.2.5" },
|
||||
|
||||
Loading…
Reference in New Issue
Block a user