diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index 76eddf54..35841605 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -2,7 +2,6 @@ import asyncio import os import textwrap import traceback -from collections.abc import Mapping from urllib.parse import quote, unquote import aiofiles @@ -93,53 +92,23 @@ async def create_database( additional_params = {**(additional_params or {})} additional_params["auto_generate_questions"] = False # 默认不生成问题 - def normalize_reranker_config(kb: str, params: dict) -> None: + def remove_reranker_config(kb: str, params: dict) -> None: + """ + 移除 reranker_config(已废弃) + 所有 reranker 参数现在通过 query_params.options 配置 + """ reranker_cfg = params.get("reranker_config") - if kb not in {"milvus"}: - if kb == "lightrag" and reranker_cfg: - logger.warning("LightRAG does not support reranker, ignoring reranker_config") - params.pop("reranker_config", None) - return - - if not reranker_cfg: - params["reranker_config"] = { - "enabled": False, - "model": "", - "recall_top_k": 50, - "final_top_k": 10, - } - return - - if not isinstance(reranker_cfg, Mapping): - raise HTTPException(status_code=400, detail="reranker_config must be an object") - - reranker_enabled = bool(reranker_cfg.get("enabled", False)) - model = (reranker_cfg.get("model") or "").strip() - recall_top_k = max(1, int(reranker_cfg.get("recall_top_k", 50))) - final_top_k = max(1, int(reranker_cfg.get("final_top_k", 10))) - - if reranker_enabled: - if not model: - raise HTTPException(status_code=400, detail="reranker_config.model is required when enabled") - if model not in config.reranker_names: - raise HTTPException(status_code=400, detail=f"Unsupported reranker model: {model}") - if final_top_k > recall_top_k: - logger.warning( - f"final_top_k ({final_top_k}) cannot exceed recall_top_k ({recall_top_k}); " - "adjusting recall_top_k to match final_top_k" + if reranker_cfg: + if kb == "milvus": + logger.info( + "reranker_config is deprecated, please use query_params.options instead" ) - recall_top_k = final_top_k - else: - model = model if model in config.reranker_names else "" + else: + logger.warning(f"{kb} does not support reranker, ignoring reranker_config") + # 移除 reranker_config,不再保存 + params.pop("reranker_config", None) - params["reranker_config"] = { - "enabled": reranker_enabled, - "model": model, - "recall_top_k": recall_top_k, - "final_top_k": final_top_k, - } - - normalize_reranker_config(kb_type, additional_params) + remove_reranker_config(kb_type, additional_params) embed_info = config.embed_model_names[embed_model_name] database_info = await knowledge_base.create_database( @@ -754,23 +723,16 @@ async def update_knowledge_base_query_params( if not kb_instance: raise HTTPException(status_code=404, detail="Knowledge base not found") - # 更新知识库元数据中的查询参数 + # 更新实例元数据中的查询参数 async with knowledge_base._metadata_lock: - # 确保知识库元数据存在 - if db_id not in knowledge_base.global_databases_meta: - knowledge_base.global_databases_meta[db_id] = {} + # 确保 db_id 在实例的 databases_meta 中 + if db_id not in kb_instance.databases_meta: + raise HTTPException(status_code=404, detail="Database not found in instance metadata") - # 初始化 query_params 结构 - if "query_params" not in knowledge_base.global_databases_meta[db_id]: - knowledge_base.global_databases_meta[db_id]["query_params"] = {} - - # 将参数保存到 options 下,与评估服务期望的结构一致 - if "options" not in knowledge_base.global_databases_meta[db_id]["query_params"]: - knowledge_base.global_databases_meta[db_id]["query_params"]["options"] = {} - - # 更新 options - knowledge_base.global_databases_meta[db_id]["query_params"]["options"].update(params) - knowledge_base._save_global_metadata() + # 使用 setdefault 简化嵌套字典的初始化 + options = kb_instance.databases_meta[db_id].setdefault("query_params", {}).setdefault("options", {}) + options.update(params) + kb_instance._save_metadata() logger.info(f"更新知识库 {db_id} 查询参数: {params}") @@ -794,8 +756,8 @@ async def get_knowledge_base_query_params(db_id: str, current_user: User = Depen reranker_names=config.reranker_names, # 传递动态配置 ) - # 获取用户保存的配置并合并 - saved_options = _get_saved_query_options(db_id) + # 获取用户保存的配置并合并(从实例 metadata 读取) + saved_options = kb_instance._get_query_params(db_id) if saved_options: params = _merge_saved_options(params, saved_options) @@ -805,18 +767,6 @@ async def get_knowledge_base_query_params(db_id: str, current_user: User = Depen logger.error(f"获取知识库查询参数失败 {e}, {traceback.format_exc()}") 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", []): diff --git a/src/knowledge/base.py b/src/knowledge/base.py index 707a7de1..1092c7bf 100644 --- a/src/knowledge/base.py +++ b/src/knowledge/base.py @@ -526,6 +526,13 @@ class KnowledgeBase(ABC): logger.warning("query is deprecated, use aquery instead") return asyncio.run(self.aquery(query_text, db_id, **kwargs)) + 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", {}) + return query_params_meta.get("options", {}) + return {} + def get_database_info(self, db_id: str) -> dict | None: """ 获取数据库详细信息 diff --git a/src/knowledge/implementations/lightrag.py b/src/knowledge/implementations/lightrag.py index f91cd67a..e0cf1216 100644 --- a/src/knowledge/implementations/lightrag.py +++ b/src/knowledge/implementations/lightrag.py @@ -448,7 +448,9 @@ class LightRagKB(KnowledgeBase): } # 过滤 kwargs,只保留 QueryParam 支持的参数 - filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + query_params = self._get_query_params(db_id) + query_params = query_params | kwargs + filtered_kwargs = {k: v for k, v in query_params.items() if k in valid_params} # 设置查询参数 params_dict = { diff --git a/src/knowledge/implementations/milvus.py b/src/knowledge/implementations/milvus.py index 7efe12c6..704cc18c 100644 --- a/src/knowledge/implementations/milvus.py +++ b/src/knowledge/implementations/milvus.py @@ -464,25 +464,26 @@ class MilvusKB(KnowledgeBase): if not collection: raise ValueError(f"Database {db_id} not found") + query_params = self._get_query_params(db_id) + # 合并查询参数:kwargs(临时参数)优先级高于 query_params(持久化参数) + # 这样允许用户在单次查询中临时覆盖持久化配置 + merged_kwargs = {**query_params, **kwargs} + try: - db_meta = self.databases_meta.get(db_id, {}) - db_metadata = db_meta.get("metadata", {}) or {} - reranker_config = db_metadata.get("reranker_config", {}) or {} + # 查询参数(从 merged_kwargs 读取) + logger.debug(f"Query params: {merged_kwargs}") + final_top_k = int(merged_kwargs.get("final_top_k", 10)) + final_top_k = max(final_top_k, 1) + similarity_threshold = float(merged_kwargs.get("similarity_threshold", 0.2)) + metric_type = merged_kwargs.get("metric_type", "COSINE") + include_distances = bool(merged_kwargs.get("include_distances", True)) - requested_top_k = int(kwargs.get("top_k", reranker_config.get("final_top_k", 30))) - requested_top_k = max(requested_top_k, 1) - similarity_threshold = float(kwargs.get("similarity_threshold", 0.2)) - metric_type = kwargs.get("metric_type", "COSINE") - include_distances = bool(kwargs.get("include_distances", True)) - - use_reranker = bool(kwargs.get("use_reranker", reranker_config.get("enabled", False))) + use_reranker = bool(merged_kwargs.get("use_reranker", False)) if use_reranker: - recall_top_k = int(kwargs.get("recall_top_k", reranker_config.get("recall_top_k", 50))) - recall_top_k = max(recall_top_k, requested_top_k) - final_top_k = requested_top_k + recall_top_k = int(merged_kwargs.get("recall_top_k", 50)) + recall_top_k = max(recall_top_k, final_top_k) else: - recall_top_k = requested_top_k - final_top_k = requested_top_k + recall_top_k = final_top_k embed_info = self.databases_meta[db_id].get("embed_info", {}) embedding_function = self._get_embedding_function(embed_info) @@ -492,7 +493,7 @@ class MilvusKB(KnowledgeBase): # 构建过滤表达式 expr = None - if file_name := kwargs.get("file_name"): + if file_name := merged_kwargs.get("file_name"): # 使用 like 支持模糊匹配 # 注意:需要转义双引号以防止注入 safe_file_name = file_name.replace('"', '\\"') @@ -537,35 +538,43 @@ class MilvusKB(KnowledgeBase): logger.debug(f"Milvus query response: {len(retrieved_chunks)} chunks found (after similarity filtering)") - if use_reranker and retrieved_chunks: + if not use_reranker: + return retrieved_chunks[:final_top_k] + + # 使用重排序模型 + reranker_model = merged_kwargs.get("reranker_model") + if not reranker_model: + raise ValueError( + "Reranker model must be specified when use_reranker=True. " + "Please provide reranker_model in query parameters." + ) + + try: + from src.models.rerank import get_reranker + + reranker = get_reranker(reranker_model) try: - reranker_model = kwargs.get("reranker_model", reranker_config.get("model")) - if not reranker_model: - logger.warning("Reranker enabled but no model specified, skipping reranking") - else: - from src.models.rerank import get_reranker + rerank_start = time.time() + documents_text = [chunk["content"] for chunk in retrieved_chunks] + rerank_scores = await reranker.acompute_score([query_text, documents_text], normalize=True) - reranker = get_reranker(reranker_model) - try: - rerank_start = time.time() - documents_text = [chunk["content"] for chunk in retrieved_chunks] - rerank_scores = await reranker.acompute_score([query_text, documents_text], normalize=True) + for chunk, rerank_score in zip(retrieved_chunks, rerank_scores): + chunk["rerank_score"] = float(rerank_score) - for chunk, rerank_score in zip(retrieved_chunks, rerank_scores): - chunk["rerank_score"] = float(rerank_score) + retrieved_chunks.sort( + key=lambda item: item.get("rerank_score", item.get("score", 0.0)), reverse=True + ) + elapsed = time.time() - rerank_start + logger.info( + f"Reranking completed for {db_id} in {elapsed:.3f}s with model {reranker_model}" + ) + finally: + await reranker.aclose() - retrieved_chunks.sort( - key=lambda item: item.get("rerank_score", item.get("score", 0.0)), reverse=True - ) - elapsed = time.time() - rerank_start - logger.info( - f"Reranking completed for {db_id} in {elapsed:.3f}s with model {reranker_model}" - ) - finally: - await reranker.aclose() - except Exception as exc: # noqa: BLE001 - logger.error(f"Reranking failed: {exc}, falling back to vector scores") + except Exception as exc: # noqa: BLE001 + logger.error(f"Reranking failed: {exc}, falling back to vector scores") + # 统一返回结果 return retrieved_chunks[:final_top_k] except Exception as e: @@ -691,22 +700,16 @@ class MilvusKB(KnowledgeBase): 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 特定参数 + # 构建 Milvus 特定参数(不再从 reranker_config 读取) options = [ { - "key": "top_k", - "label": "TopK", + "key": "final_top_k", + "label": "最终返回数", "type": "number", - "default": reranker_config.get("final_top_k", 10), + "default": 10, "min": 1, "max": 100, - "description": "返回的最大结果数量", + "description": "重排序后返回给前端的文档数量", }, { "key": "similarity_threshold", @@ -741,17 +744,17 @@ class MilvusKB(KnowledgeBase): "key": "use_reranker", "label": "启用重排序", "type": "boolean", - "default": reranker_enabled, + "default": False, "description": "是否使用精排模型对检索结果进行重排序", }, { "key": "recall_top_k", "label": "召回数量", "type": "number", - "default": reranker_config.get("recall_top_k", 50), + "default": 50, "min": 10, "max": 200, - "description": "启用重排序时向量检索的候选数量", + "description": "向量检索时保留的候选数量(启用重排序时有效)", }, ] @@ -763,9 +766,9 @@ class MilvusKB(KnowledgeBase): "key": "reranker_model", "label": "重排序模型", "type": "select", - "default": reranker_config.get("model", ""), + "default": "", "options": [{"label": info.name, "value": model_id} for model_id, info in reranker_names.items()], - "description": "覆盖默认配置,选择用于本次查询的重排序模型", + "description": "选择用于本次查询的重排序模型", } ) diff --git a/src/knowledge/manager.py b/src/knowledge/manager.py index 35c9d3df..e5ac85b8 100644 --- a/src/knowledge/manager.py +++ b/src/knowledge/manager.py @@ -43,6 +43,12 @@ class KnowledgeBaseManager: # 初始化已存在的知识库实例 self._initialize_existing_kbs() + # 迁移 query_params 到 instance metadata + try: + self._migrate_all_query_params() + except Exception as e: + logger.warning(f"Failed to migrate query_params: {e}") + logger.info("KnowledgeBaseManager initialized") # 在后台运行数据一致性检测(不阻塞初始化) @@ -323,6 +329,36 @@ class KnowledgeBaseManager: kb_instance = self._get_kb_for_database(db_id) return kb_instance.query(query_text, db_id, **kwargs) + def _migrate_all_query_params(self): + """将所有 query_params 从 global metadata 迁移到 instance metadata""" + migration_count = 0 + + for db_id, global_meta in list(self.global_databases_meta.items()): + if "query_params" not in global_meta: + continue + + kb_type = global_meta.get("kb_type", "lightrag") + kb_instance = self.kb_instances.get(kb_type) + + if not kb_instance or db_id not in kb_instance.databases_meta: + logger.warning(f"Cannot migrate query_params for {db_id}, skipping") + continue + + # 检查是否已迁移 + if "query_params" in kb_instance.databases_meta[db_id]: + # 已经迁移过,直接清理 global metadata 并跳过 + del global_meta["query_params"] + continue + + # 执行迁移 + kb_instance.databases_meta[db_id]["query_params"] = global_meta["query_params"] + del global_meta["query_params"] + migration_count += 1 + + if migration_count > 0: + self._save_global_metadata() + logger.info(f"Successfully migrated query_params for {migration_count} databases") + def get_database_info(self, db_id: str) -> dict | None: """获取数据库详细信息""" try: diff --git a/test/api/test_knowledge_router.py b/test/api/test_knowledge_router.py index b0ce41f1..e5d3eeb2 100644 --- a/test/api/test_knowledge_router.py +++ b/test/api/test_knowledge_router.py @@ -54,20 +54,14 @@ async def test_knowledge_routes_enforce_permissions(test_client, standard_user, async def test_admin_can_create_vector_db_with_reranker(test_client, admin_headers): + """测试创建向量库并配置 reranker 参数(通过 query_params.options)""" db_name = f"pytest_rerank_{uuid.uuid4().hex[:6]}" payload = { "database_name": db_name, "description": "Vector DB with reranker", "embed_model_name": "siliconflow/BAAI/bge-m3", "kb_type": "milvus", - "additional_params": { - "reranker_config": { - "enabled": True, - "model": "siliconflow/BAAI/bge-reranker-v2-m3", - "recall_top_k": 25, - "final_top_k": 8, - } - }, + "additional_params": {}, } create_response = await test_client.post("/api/knowledge/databases", json=payload, headers=admin_headers) @@ -77,14 +71,7 @@ async def test_admin_can_create_vector_db_with_reranker(test_client, admin_heade db_id = db_payload["db_id"] try: - info_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers) - assert info_response.status_code == 200, info_response.text - info_payload = info_response.json() - - reranker_config = info_payload.get("metadata", {}).get("reranker_config", {}) - assert reranker_config.get("enabled") is True - assert reranker_config.get("model") == "siliconflow/BAAI/bge-reranker-v2-m3" - + # 获取查询参数配置 params_response = await test_client.get(f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers) assert params_response.status_code == 200, params_response.text @@ -92,7 +79,49 @@ async def test_admin_can_create_vector_db_with_reranker(test_client, admin_heade options = params_payload.get("params", {}).get("options", []) option_keys = {option.get("key") for option in options} + # 验证新的参数名称 + assert "final_top_k" in option_keys assert "use_reranker" in option_keys assert "recall_top_k" in option_keys + assert "reranker_model" in option_keys + + # 验证参数配置 + final_top_k_option = next((opt for opt in options if opt.get("key") == "final_top_k"), None) + assert final_top_k_option is not None + assert final_top_k_option.get("default") == 10 + + use_reranker_option = next((opt for opt in options if opt.get("key") == "use_reranker"), None) + assert use_reranker_option is not None + assert use_reranker_option.get("default") is False + + # 保存查询参数(模拟前端配置) + update_params = { + "final_top_k": 5, + "use_reranker": True, + "recall_top_k": 20, + } + update_response = await test_client.put( + f"/api/knowledge/databases/{db_id}/query-params", + json=update_params, + headers=admin_headers + ) + assert update_response.status_code == 200, update_response.text + + # 再次获取参数,验证保存成功 + params_response2 = await test_client.get(f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers) + assert params_response2.status_code == 200, params_response2.text + + params_payload2 = params_response2.json() + options2 = params_payload2.get("params", {}).get("options", []) + + # 验证保存的值 + final_top_k_option2 = next((opt for opt in options2 if opt.get("key") == "final_top_k"), None) + assert final_top_k_option2 is not None + assert final_top_k_option2.get("default") == 5 # 保存的值 + + use_reranker_option2 = next((opt for opt in options2 if opt.get("key") == "use_reranker"), None) + assert use_reranker_option2 is not None + assert use_reranker_option2.get("default") is True # 保存的值 + finally: await test_client.delete(f"/api/knowledge/databases/{db_id}", headers=admin_headers) diff --git a/web/src/views/DataBaseView.vue b/web/src/views/DataBaseView.vue index 0296a3f9..828bc266 100644 --- a/web/src/views/DataBaseView.vue +++ b/web/src/views/DataBaseView.vue @@ -94,65 +94,6 @@ - -
-
-
- 启用重排序 - - - -
- -
- - -
-
- - -

- 暂无可用模型,请在系统配置中添加。 -

-
- -
-
- - -

向量检索阶段保留的候选数量

-
-
- - -

重排序后返回给前端的文档数量

-
-
-
-
-