fix(knowledge): 修复智能体调用场景下,知识库查询参数不生效的问题。

- 将查询参数从全局元数据迁移到知识库实例元数据
- 废弃 reranker_config 配置,统一通过 query_params.options 管理
- 修改 Milvus 和 LightRAG 实现以使用新的查询参数结构
- 更新前端界面移除旧的 reranker 配置表单
- 添加测试用例验证新的查询参数配置流程
This commit is contained in:
Wenjie Zhang 2026-01-09 01:21:41 +08:00
parent c00949778d
commit 05a67c24f8
7 changed files with 174 additions and 336 deletions

View File

@ -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", []):

View File

@ -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:
"""
获取数据库详细信息

View File

@ -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 = {

View File

@ -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": "选择用于本次查询的重排序模型",
}
)

View File

@ -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:

View File

@ -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)

View File

@ -94,65 +94,6 @@
<InfoCircleOutlined style="margin-left: 8px; color: var(--gray-500); cursor: help;" />
</a-tooltip>
</div>
<div
v-if="['milvus'].includes(newDatabase.kb_type)"
class="reranker-config"
>
<div class="reranker-row">
<div class="reranker-title">
<span>启用重排序</span>
<a-tooltip title="向量检索后使用交叉编码模型对候选文档重新排序,提升召回质量。">
<QuestionCircleOutlined class="hint-icon" />
</a-tooltip>
</div>
<a-switch
v-model:checked="newDatabase.reranker.enabled"
:disabled="rerankerOptions.length === 0"
/>
</div>
<transition name="fade">
<div v-if="newDatabase.reranker.enabled" class="reranker-form">
<div class="form-field">
<label>重排序模型</label>
<a-select
v-model:value="newDatabase.reranker.model"
:options="rerankerOptions"
placeholder="选择重排序模型"
:disabled="rerankerOptions.length === 0"
/>
<p class="field-hint" v-if="rerankerOptions.length === 0">
暂无可用模型请在系统配置中添加
</p>
</div>
<div class="form-grid">
<div class="form-field">
<label>召回数量</label>
<a-input-number
v-model:value="newDatabase.reranker.recall_top_k"
:min="10"
:max="200"
:step="5"
style="width: 100%;"
/>
<p class="field-hint">向量检索阶段保留的候选数量</p>
</div>
<div class="form-field">
<label>最终返回数</label>
<a-input-number
v-model:value="newDatabase.reranker.final_top_k"
:min="1"
:max="100"
style="width: 100%;"
/>
<p class="field-hint">重排序后返回给前端的文档数量</p>
</div>
</div>
</div>
</transition>
</div>
<template #footer>
<a-button key="back" @click="cancelCreateDatabase">取消</a-button>
<a-button key="submit" type="primary" :loading="dbState.creating" @click="handleCreateDatabase">创建</a-button>
@ -279,26 +220,11 @@ const createEmptyDatabaseForm = () => ({
llm_info: {
provider: '',
model_name: ''
},
reranker: {
enabled: false,
model: '',
recall_top_k: 50,
final_top_k: 10,
}
})
const newDatabase = reactive(createEmptyDatabaseForm())
const rerankerOptions = computed(() =>
Object.entries(configStore?.config?.reranker_names || {}).map(([value, info]) => ({
label: info?.name || value,
value
}))
)
const isVectorKb = computed(() => ['milvus'].includes(newDatabase.kb_type))
const llmModelSpec = computed(() => {
const provider = newDatabase.llm_info?.provider || ''
const modelName = newDatabase.llm_info?.model_name || ''
@ -403,9 +329,6 @@ const handleKbTypeChange = (type) => {
console.log('知识库类型改变:', type)
resetNewDatabase()
newDatabase.kb_type = type
if (!['milvus'].includes(type)) {
newDatabase.reranker.enabled = false
}
}
// LLM
@ -438,14 +361,6 @@ const buildRequestData = () => {
if (newDatabase.storage) {
requestData.additional_params.storage = newDatabase.storage
}
if (newDatabase.reranker.enabled) {
requestData.additional_params.reranker_config = {
enabled: true,
model: newDatabase.reranker.model,
recall_top_k: Number(newDatabase.reranker.recall_top_k) || 50,
final_top_k: Number(newDatabase.reranker.final_top_k) || 10,
}
}
}
if (newDatabase.kb_type === 'lightrag') {
@ -477,42 +392,6 @@ const navigateToDatabase = (databaseId) => {
router.push({ path: `/database/${databaseId}` });
};
watch(() => newDatabase.reranker.enabled, (enabled) => {
if (
enabled &&
!newDatabase.reranker.model &&
rerankerOptions.value.length > 0
) {
newDatabase.reranker.model = rerankerOptions.value[0].value
}
})
watch(rerankerOptions, (options) => {
if (!newDatabase.reranker.enabled || options.length === 0) {
return
}
const exists = options.some(option => option.value === newDatabase.reranker.model)
if (!exists) {
newDatabase.reranker.model = options[0].value
}
})
watch(isVectorKb, (isVector) => {
if (!isVector) {
newDatabase.reranker.enabled = false
}
})
watch(
() => newDatabase.reranker.final_top_k,
(value) => {
if (!newDatabase.reranker.enabled) return
if (value > newDatabase.reranker.recall_top_k) {
newDatabase.reranker.recall_top_k = value
}
}
)
watch(() => route.path, (newPath) => {
if (newPath === '/database') {
databaseStore.loadDatabases();
@ -522,7 +401,6 @@ watch(() => route.path, (newPath) => {
onMounted(() => {
loadSupportedKbTypes()
databaseStore.loadDatabases()
// configStore
})
</script>
@ -539,73 +417,6 @@ onMounted(() => {
margin-bottom: 12px;
}
.reranker-config {
border: 1px solid var(--gray-200);
border-radius: 12px;
padding: 16px;
margin-top: 16px;
background: var(--gray-25);
.reranker-row {
display: flex;
align-items: center;
justify-content: space-between;
margin-bottom: 16px;
&:last-child {
margin-bottom: 0;
}
.reranker-title {
display: flex;
align-items: center;
gap: 6px;
font-weight: 500;
color: var(--gray-800);
}
.hint-icon {
color: var(--gray-500);
cursor: help;
}
}
.reranker-form {
display: flex;
flex-direction: column;
gap: 16px;
.form-grid {
display: grid;
grid-template-columns: repeat(2, minmax(0, 1fr));
gap: 16px;
@media (max-width: 768px) {
grid-template-columns: 1fr;
}
}
.form-field {
label {
display: block;
font-size: 14px;
margin-bottom: 8px;
color: var(--gray-700);
}
.field-hint {
margin-top: 6px;
font-size: 12px;
color: var(--gray-500);
&:last-child {
margin-top: 0;
}
}
}
}
}
.kb-type-cards {
display: grid;
grid-template-columns: repeat(3, 1fr);