diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index e235d6d5..76eddf54 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -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) diff --git a/src/knowledge/base.py b/src/knowledge/base.py index 72d0a736..c652114b 100644 --- a/src/knowledge/base.py +++ b/src/knowledge/base.py @@ -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 diff --git a/src/knowledge/implementations/lightrag.py b/src/knowledge/implementations/lightrag.py index e310bba5..97be526b 100644 --- a/src/knowledge/implementations/lightrag.py +++ b/src/knowledge/implementations/lightrag.py @@ -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 原生功能导出知识库数据。 diff --git a/src/knowledge/implementations/milvus.py b/src/knowledge/implementations/milvus.py index ce068bdd..39073787 100644 --- a/src/knowledge/implementations/milvus.py +++ b/src/knowledge/implementations/milvus.py @@ -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: diff --git a/uv.lock b/uv.lock index e43c0ab1..6d0abbfc 100644 --- a/uv.lock +++ b/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" },