fix(knowledge): 修复智能体调用场景下,知识库查询参数不生效的问题。
- 将查询参数从全局元数据迁移到知识库实例元数据 - 废弃 reranker_config 配置,统一通过 query_params.options 管理 - 修改 Milvus 和 LightRAG 实现以使用新的查询参数结构 - 更新前端界面移除旧的 reranker 配置表单 - 添加测试用例验证新的查询参数配置流程
This commit is contained in:
parent
c00949778d
commit
05a67c24f8
@ -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", []):
|
||||
|
||||
@ -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:
|
||||
"""
|
||||
获取数据库详细信息
|
||||
|
||||
@ -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 = {
|
||||
|
||||
@ -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": "选择用于本次查询的重排序模型",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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);
|
||||
|
||||
Loading…
Reference in New Issue
Block a user