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 @@
- 暂无可用模型,请在系统配置中添加。 -
-向量检索阶段保留的候选数量
-重排序后返回给前端的文档数量
-