feat(evaluation): 新增RAG评估功能模块

实现RAG评估系统,包括基准管理、评估指标计算和结果展示功能

- 添加评估路由和API接口
- 实现评估指标计算工具类
- 新增基准上传和自动生成功能
- 添加评估结果展示组件
- 完善相关文档和测试用例
This commit is contained in:
Wenjie Zhang 2025-12-09 16:21:18 +08:00
parent fd5fca0076
commit 83e0caae7d
22 changed files with 3818 additions and 66 deletions

View File

@ -15,12 +15,17 @@
- 同名文件处理逻辑:遇到同名文件则在上传区域提示,是否删除旧文件
- conversation 待修改为异步的版本
- DBManager 需要将数据库修改为异步的aiosqlite或者异步mysql缓存使用Redis存储
- 【eval】缺少自动生成评估的功能
### Bugs
- 部分异常状态下,智能体的模型名称出现重叠[#279](https://github.com/xerrors/Yuxi-Know/issues/279)
- DeepSeek 官方接口适配会出现问题
- 目前的知识库的图片存在公开访问风险
- 深度分析智能体需要考虑上下文超限的问题
- 【eval】检索评估缺少检索参数的内容
- 【eval】总体评分替换为确定的某个指标最好是在右侧可以选择可行的指标
- 【eval】检索配置应该修改为一个弹窗
- 【eval】当前的检索配置仅保存在本地实际并没有保存在知识库的元数据中
### 新增
- 优化知识库详情页面,更加简洁清晰

View File

@ -5,6 +5,7 @@ from server.routers.chat_router import chat
from server.routers.dashboard_router import dashboard
from server.routers.graph_router import graph
from server.routers.knowledge_router import knowledge
from server.routers.evaluation_router import evaluation
from server.routers.mindmap_router import mindmap
from server.routers.system_router import system
from server.routers.task_router import tasks
@ -17,6 +18,7 @@ router.include_router(auth) # /api/auth/*
router.include_router(chat) # /api/chat/*
router.include_router(dashboard) # /api/dashboard/*
router.include_router(knowledge) # /api/knowledge/*
router.include_router(evaluation) # /api/evaluation/*
router.include_router(mindmap) # /api/mindmap/*
router.include_router(graph) # /api/graph/*
router.include_router(tasks) # /api/tasks/*

View File

@ -0,0 +1,172 @@
import traceback
from fastapi import APIRouter, HTTPException, Depends, File, Form, Body, UploadFile
from src.storage.db.models import User
from server.utils.auth_middleware import get_admin_user
from src.utils import logger
# 创建路由器
evaluation = APIRouter(prefix="/evaluation", tags=["evaluation"])
@evaluation.get("/benchmarks/{benchmark_id}")
async def get_evaluation_benchmark(benchmark_id: str, current_user: User = Depends(get_admin_user)):
"""获取评估基准详情"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
benchmark = await service.get_benchmark_detail(benchmark_id)
return {"message": "success", "data": benchmark}
except Exception as e:
logger.error(f"获取评估基准详情失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"获取评估基准详情失败: {str(e)}")
@evaluation.delete("/benchmarks/{benchmark_id}")
async def delete_evaluation_benchmark(benchmark_id: str, current_user: User = Depends(get_admin_user)):
"""删除评估基准"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
await service.delete_benchmark(benchmark_id)
return {"message": "success", "data": None}
except Exception as e:
logger.error(f"删除评估基准失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"删除评估基准失败: {str(e)}")
@evaluation.get("/{task_id}/results")
async def get_evaluation_results(task_id: str, current_user: User = Depends(get_admin_user)):
"""获取评估结果"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
results = await service.get_evaluation_results(task_id)
return {"message": "success", "data": results}
except Exception as e:
logger.error(f"获取评估结果失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"获取评估结果失败: {str(e)}")
@evaluation.delete("/{task_id}")
async def delete_evaluation_result(task_id: str, current_user: User = Depends(get_admin_user)):
"""删除评估结果"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
await service.delete_evaluation_result(task_id)
return {"message": "success", "data": None}
except Exception as e:
logger.error(f"删除评估结果失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"删除评估结果失败: {str(e)}")
# ============================================================================
# === Knowledge-specific evaluation endpoints ===
# ============================================================================
@evaluation.post("/databases/{db_id}/benchmarks/upload")
async def upload_evaluation_benchmark(
db_id: str,
file: UploadFile = File(...),
name: str = Form(...),
description: str = Form(""),
current_user: User = Depends(get_admin_user),
):
"""上传评估基准文件"""
from src.services.evaluation_service import EvaluationService
try:
# 验证文件格式
if not file.filename.endswith(".jsonl"):
raise HTTPException(status_code=400, detail="仅支持JSONL格式文件")
# 读取文件内容
content = await file.read()
# 调用评估服务处理上传
service = EvaluationService()
result = await service.upload_benchmark(
db_id=db_id,
file_content=content,
filename=file.filename,
name=name,
description=description,
created_by=current_user.user_id,
)
return {"message": "success", "data": result}
except HTTPException:
raise
except Exception as e:
logger.error(f"上传评估基准失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"上传评估基准失败: {str(e)}")
@evaluation.get("/databases/{db_id}/benchmarks")
async def get_evaluation_benchmarks(db_id: str, current_user: User = Depends(get_admin_user)):
"""获取知识库的评估基准列表"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
benchmarks = await service.get_benchmarks(db_id)
return {"message": "success", "data": benchmarks}
except Exception as e:
logger.error(f"获取评估基准列表失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"获取评估基准列表失败: {str(e)}")
@evaluation.post("/databases/{db_id}/benchmarks/generate")
async def generate_evaluation_benchmark(
db_id: str, params: dict = Body(...), current_user: User = Depends(get_admin_user)
):
"""自动生成评估基准"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
result = await service.generate_benchmark(db_id=db_id, params=params, created_by=current_user.user_id)
return {"message": "success", "data": result}
except Exception as e:
logger.error(f"生成评估基准失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"生成评估基准失败: {str(e)}")
@evaluation.post("/databases/{db_id}/run")
async def run_evaluation(db_id: str, params: dict = Body(...), current_user: User = Depends(get_admin_user)):
"""运行RAG评估"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
task_id = await service.run_evaluation(
db_id=db_id,
benchmark_id=params.get("benchmark_id"),
retrieval_config=params.get("retrieval_config", {}),
created_by=current_user.user_id,
)
return {"message": "success", "data": {"task_id": task_id}}
except Exception as e:
logger.error(f"启动评估失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"启动评估失败: {str(e)}")
@evaluation.get("/databases/{db_id}/history")
async def get_evaluation_history(db_id: str, current_user: User = Depends(get_admin_user)):
"""获取知识库的评估历史记录"""
from src.services.evaluation_service import EvaluationService
try:
service = EvaluationService()
history = await service.get_evaluation_history(db_id)
return {"message": "success", "data": history}
except Exception as e:
logger.error(f"获取评估历史失败: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"获取评估历史失败: {str(e)}")

View File

@ -320,7 +320,7 @@ async def add_documents(
item_type = "URL" if content_type == "url" else "文件"
failed_count = len([_p for _p in processed_items if _p.get("status") == "failed"])
success_items = [_p for _p in processed_items if _p.get("status") == "done"]
# success_items = [_p for _p in processed_items if _p.get("status") == "done"]
summary = {
"db_id": db_id,
"item_type": item_type,
@ -550,6 +550,7 @@ async def download_document(db_id: str, doc_id: str, request: Request, current_u
# 根据path类型选择下载方式
from src.knowledge.utils.kb_utils import is_minio_url
if is_minio_url(file_path):
# MinIO下载
logger.debug(f"Downloading from MinIO: {file_path}")
@ -557,6 +558,7 @@ async def download_document(db_id: str, doc_id: str, request: Request, current_u
try:
# 使用通用函数解析MinIO URL
from src.knowledge.utils.kb_utils import parse_minio_url
bucket_name, object_name = parse_minio_url(file_path)
logger.debug(f"Parsed bucket_name: {bucket_name}, object_name: {object_name}")
@ -683,6 +685,37 @@ async def query_test(
return {"message": f"测试查询失败: {e}", "status": "failed"}
@knowledge.put("/databases/{db_id}/query-params")
async def update_knowledge_base_query_params(
db_id: str, params: dict = Body(...), current_user: User = Depends(get_admin_user)
):
"""更新知识库查询参数配置"""
try:
# 获取知识库实例
kb_instance = knowledge_base.get_kb(db_id)
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] = {}
# 保存查询参数到元数据
if "query_params" not in knowledge_base.global_databases_meta[db_id]:
knowledge_base.global_databases_meta[db_id]["query_params"] = {}
knowledge_base.global_databases_meta[db_id]["query_params"].update(params)
knowledge_base._save_global_metadata()
return {"message": "success", "data": params}
except Exception as e:
logger.error(f"更新知识库查询参数失败: {e}")
raise HTTPException(status_code=500, detail=f"更新查询参数失败: {str(e)}")
@knowledge.get("/databases/{db_id}/query-params")
async def get_knowledge_base_query_params(db_id: str, current_user: User = Depends(get_admin_user)):
"""获取知识库类型特定的查询参数"""
@ -1136,6 +1169,7 @@ async def upload_file(
# 直接上传到MinIO添加时间戳区分版本
import time
timestamp = int(time.time() * 1000)
minio_filename = f"{basename}_{timestamp}{ext}"
@ -1146,7 +1180,7 @@ async def upload_file(
bucket_name = "default-uploads"
# 上传到MinIO
minio_url = await aupload_file_to_minio(bucket_name, minio_filename, file_bytes, ext.lstrip('.'))
minio_url = await aupload_file_to_minio(bucket_name, minio_filename, file_bytes, ext.lstrip("."))
# 检测同名文件(基于原始文件名)
same_name_files = await knowledge_base.get_same_name_files(db_id, filename)
@ -1154,16 +1188,16 @@ async def upload_file(
return {
"message": "File successfully uploaded",
"file_path": minio_url, # MinIO路径作为主要路径
"minio_path": minio_url, # MinIO路径
"file_path": minio_url, # MinIO路径作为主要路径
"minio_path": minio_url, # MinIO路径
"db_id": db_id,
"content_hash": content_hash,
"filename": filename, # 原始文件名(小写)
"original_filename": basename, # 原始文件名(去掉后缀)
"minio_filename": minio_filename, # MinIO中的文件名带时间戳
"bucket_name": bucket_name, # MinIO存储桶名称
"same_name_files": same_name_files, # 同名文件列表
"has_same_name": has_same_name # 是否包含同名文件标志
"filename": filename, # 原始文件名(小写)
"original_filename": basename, # 原始文件名(去掉后缀)
"minio_filename": minio_filename, # MinIO中的文件名带时间戳
"bucket_name": bucket_name, # MinIO存储桶名称
"same_name_files": same_name_files, # 同名文件列表
"has_same_name": has_same_name, # 是否包含同名文件标志
}

View File

@ -41,6 +41,12 @@ class Task:
data = asdict(self)
return data
def to_summary_dict(self) -> dict[str, Any]:
data = asdict(self)
data.pop("payload", None)
data.pop("result", None)
return data
@classmethod
def from_dict(cls, data: dict[str, Any]) -> "Task":
return cls(
@ -160,7 +166,7 @@ class Tasker:
}
return {
"tasks": [task.to_dict() for task in limited_tasks],
"tasks": [task.to_summary_dict() for task in limited_tasks],
"summary": summary,
}

View File

@ -95,7 +95,7 @@ class MilvusKB(KnowledgeBase):
if not (metadata := self.databases_meta.get(db_id)):
raise ValueError(f"Database {db_id} not found")
# embed_info = metadata.get("embed_info", {})
# 获取嵌入模型信息
if not (embed_info := metadata.get("embed_info")):
logger.error(f"Embedding info not found for database {db_id}, using default model")
embed_info = config.embed_model_names[config.embed_model]
@ -109,45 +109,58 @@ class MilvusKB(KnowledgeBase):
# 检查嵌入模型是否匹配
description = collection.description
expected_model = getattr(embed_info, "name", "default") if embed_info else "default"
expected_model = embed_info["name"] if embed_info else "default"
if expected_model not in description:
logger.warning(f"Collection {collection_name} model mismatch, recreating...")
logger.warning(
f"Collection {collection_name} model mismatch: "
f"expected='{expected_model}', found_in_description='{description}'"
)
utility.drop_collection(collection_name, using=self.connection_alias)
raise Exception("Model mismatch, recreating collection")
return self._create_new_collection(collection_name, embed_info, db_id)
logger.info(f"Retrieved existing collection: {collection_name}")
return collection
else:
raise Exception("Collection not found, creating new one")
logger.info(f"Collection {collection_name} not found, creating new one")
return self._create_new_collection(collection_name, embed_info, db_id)
except Exception:
# 创建新集合
embedding_dim = embed_info.get("dimension", 1024)
model_name = embed_info.get("name", "default")
except (connections.MilvusException, RuntimeError) as e:
logger.error(f"Error checking collection {collection_name}: {e}")
raise
except Exception as e:
logger.error(f"Unexpected error while managing collection {collection_name}: {e}")
logger.debug(f"Traceback: {traceback.format_exc()}")
raise
# 定义集合Schema
fields = [
FieldSchema(name="id", dtype=DataType.VARCHAR, max_length=100, is_primary=True),
FieldSchema(name="content", dtype=DataType.VARCHAR, max_length=65535),
FieldSchema(name="source", dtype=DataType.VARCHAR, max_length=500),
FieldSchema(name="chunk_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="file_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="chunk_index", dtype=DataType.INT64),
FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=embedding_dim),
]
def _create_new_collection(self, collection_name: str, embed_info: Any, db_id: str) -> Collection:
"""创建新的 Milvus 集合"""
embedding_dim = embed_info.get("dimension", 1024)
model_name = embed_info.get("name", "default")
schema = CollectionSchema(
fields=fields, description=f"Knowledge base collection for {db_id} using {model_name}"
)
# 定义集合Schema
fields = [
FieldSchema(name="id", dtype=DataType.VARCHAR, max_length=100, is_primary=True),
FieldSchema(name="content", dtype=DataType.VARCHAR, max_length=65535),
FieldSchema(name="source", dtype=DataType.VARCHAR, max_length=500),
FieldSchema(name="chunk_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="file_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="chunk_index", dtype=DataType.INT64),
FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=embedding_dim),
]
# 创建集合
collection = Collection(name=collection_name, schema=schema, using=self.connection_alias)
schema = CollectionSchema(
fields=fields, description=f"Knowledge base collection for {db_id} using {model_name}"
)
# 创建索引
index_params = {"metric_type": "COSINE", "index_type": "IVF_FLAT", "params": {"nlist": 1024}}
collection.create_index("embedding", index_params)
# 创建集合
collection = Collection(name=collection_name, schema=schema, using=self.connection_alias)
logger.info(f"Created new Milvus collection: {collection_name}: {model_name=}, {embedding_dim=}")
# 创建索引
index_params = {"metric_type": "COSINE", "index_type": "IVF_FLAT", "params": {"nlist": 1024}}
collection.create_index("embedding", index_params)
logger.info(f"Created new Milvus collection: {collection_name} '{model_name=}', {embedding_dim=}")
return collection

View File

@ -402,9 +402,8 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) -
- params['_zip_images_info']: 图片信息列表
- params['_zip_content_hash']: 内容哈希值
"""
import tempfile
import aiofiles
import os
import tempfile
# 检测是否是MinIO URL
from src.knowledge.utils.kb_utils import is_minio_url
@ -474,7 +473,7 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) -
elif file_ext == ".docx":
text = _extract_docx_markdown_with_images(file_path_obj, params=params)
result = f"" + text
result = "" + text
elif file_ext == ".doc":
text = _extract_word_text(file_path_obj)
@ -500,7 +499,7 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) -
df = pd.read_csv(file_path_obj)
# 将每一行数据与表头组合成独立的表格
markdown_content = f""
markdown_content = ""
for index, row in df.iterrows():
# 创建包含表头和当前行的小表格
@ -515,7 +514,7 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) -
import pandas as pd
from openpyxl import load_workbook
markdown_content = f""
markdown_content = ""
# 使用 openpyxl 加载工作簿以正确处理合并单元格
wb = load_workbook(file_path_obj, data_only=True)
@ -604,7 +603,7 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) -
# 尝试作为文本文件读取
raise ValueError(f"Unsupported file type: {file_ext}")
except Exception as e:
except Exception:
# 清理临时文件
if is_minio_url(file_path) and os.path.exists(actual_file_path):
try:

View File

@ -408,13 +408,15 @@ class KnowledgeBaseManager:
current_filename = file_info.get("filename", "")
if current_filename.lower() == filename.lower():
same_name_files.append({
"file_id": file_id,
"filename": current_filename,
"size": file_info.get("size", 0),
"created_at": file_info.get("created_at", ""),
"content_hash": file_info.get("content_hash", "")
})
same_name_files.append(
{
"file_id": file_id,
"filename": current_filename,
"size": file_info.get("size", 0),
"created_at": file_info.get("created_at", ""),
"content_hash": file_info.get("content_hash", ""),
}
)
# 按上传时间降序排序
same_name_files.sort(key=lambda x: x.get("created_at", ""), reverse=True)

View File

@ -165,7 +165,8 @@ async def prepare_item_metadata(item: str, content_type: str, db_id: str, params
# 如果文件名包含时间戳,提取原始文件名
import re
timestamp_pattern = r'^(.+)_(\d{13})(\.[^.]+)$'
timestamp_pattern = r"^(.+)_(\d{13})(\.[^.]+)$"
match = re.match(timestamp_pattern, filename)
if match:
original_filename = match.group(1) + match.group(3)
@ -347,10 +348,10 @@ def parse_minio_url(file_path: str) -> tuple[str, str]:
parsed_url = urlparse(file_path)
# 从URL路径中提取对象名称去掉开头的斜杠
object_name = parsed_url.path.lstrip('/')
object_name = parsed_url.path.lstrip("/")
# 分离bucket名称和对象名称
path_parts = object_name.split('/', 1)
path_parts = object_name.split("/", 1)
if len(path_parts) > 1:
bucket_name = path_parts[0]
object_name = path_parts[1]

View File

@ -0,0 +1,597 @@
import asyncio
import glob
import json
import os
import uuid
from datetime import datetime
from typing import Any
from server.services.tasker import TaskContext, tasker
from src.knowledge import knowledge_base
from src.models import select_model
from src.utils import logger
from src.utils.evaluation_metrics import EvaluationMetricsCalculator
class EvaluationService:
"""RAG评估服务 - 基于文件存储的版本"""
def __init__(self):
# 使用环境变量 DATA_DIR 或默认 'saves'
self.data_dir = os.environ.get("DATA_DIR", "saves")
self.root_eval_dir = os.path.join(self.data_dir, "evaluation")
def _get_benchmark_dir(self, db_id: str) -> str:
path = os.path.join(self.root_eval_dir, db_id, "benchmarks")
os.makedirs(path, exist_ok=True)
return path
def _get_result_dir(self, db_id: str) -> str:
path = os.path.join(self.root_eval_dir, db_id, "results")
os.makedirs(path, exist_ok=True)
return path
def _find_benchmark_location(self, benchmark_id: str) -> tuple:
"""
高效查找基准文件位置返回 (db_id, meta_file_path)
避免全局搜索先检查是否有索引映射
"""
# 由于当前文件结构限制,仍然需要搜索
# 但可以优化搜索顺序和错误处理
try:
# 搜索所有 DB 目录找到该 benchmark
pattern = os.path.join(self.root_eval_dir, "*", "benchmarks", f"{benchmark_id}.meta.json")
matches = glob.glob(pattern)
if not matches:
raise ValueError(f"评估基准 {benchmark_id} 不存在")
meta_file_path = matches[0]
# 从路径推断 db_id (parent of parent)
# path: .../{db_id}/benchmarks/{bid}.meta.json
db_id = os.path.basename(os.path.dirname(os.path.dirname(meta_file_path)))
return db_id, meta_file_path
except Exception as e:
logger.error(f"查找基准文件失败: {e}")
raise
def _find_result_location(self, task_id: str) -> tuple:
"""
高效查找评估结果文件位置返回 (db_id, result_file_path)
"""
try:
pattern = os.path.join(self.root_eval_dir, "*", "results", f"{task_id}.json")
matches = glob.glob(pattern)
if not matches:
raise ValueError(f"评估结果 {task_id} 不存在")
result_file_path = matches[0]
# 从路径推断 db_id
db_id = os.path.basename(os.path.dirname(os.path.dirname(result_file_path)))
return db_id, result_file_path
except Exception as e:
logger.error(f"查找评估结果文件失败: {e}")
raise
async def upload_benchmark(
self, db_id: str, file_content: bytes, filename: str, name: str, description: str, created_by: str
) -> dict[str, Any]:
"""上传评估基准文件"""
try:
content_str = file_content.decode("utf-8")
questions = []
has_gold_chunks = False
has_gold_answers = False
# 解析 JSONL
for line_num, line in enumerate(content_str.strip().split("\n"), 1):
if not line.strip():
continue
try:
item = json.loads(line)
if "query" not in item:
raise ValueError(f"{line_num}行缺少必需的'query'字段")
if item.get("gold_chunk_ids"):
has_gold_chunks = True
if item.get("gold_answer"):
has_gold_answers = True
questions.append(item)
except json.JSONDecodeError as e:
raise ValueError(f"{line_num}行JSON格式错误: {str(e)}")
if not questions:
raise ValueError("文件中没有有效的问题数据")
benchmark_id = f"benchmark_{uuid.uuid4().hex[:8]}"
benchmark_dir = self._get_benchmark_dir(db_id)
# 保存数据文件 (.jsonl)
data_file_path = os.path.join(benchmark_dir, f"{benchmark_id}.jsonl")
with open(data_file_path, "w", encoding="utf-8") as f:
f.write(content_str)
# 保存元数据文件 (.meta.json)
meta = {
"id": benchmark_id, # 前端期望字段可能是 id
"benchmark_id": benchmark_id,
"name": name,
"description": description,
"db_id": db_id,
"question_count": len(questions),
"has_gold_chunks": has_gold_chunks,
"has_gold_answers": has_gold_answers,
"created_by": created_by,
"created_at": datetime.utcnow().isoformat(),
"updated_at": datetime.utcnow().isoformat(),
}
meta_file_path = os.path.join(benchmark_dir, f"{benchmark_id}.meta.json")
with open(meta_file_path, "w", encoding="utf-8") as f:
json.dump(meta, f, ensure_ascii=False, indent=2)
return meta
except Exception as e:
logger.error(f"上传评估基准失败: {e}")
raise
async def get_benchmarks(self, db_id: str) -> list[dict[str, Any]]:
"""获取知识库的评估基准列表"""
try:
benchmark_dir = self._get_benchmark_dir(db_id)
benchmarks = []
# 查找所有 .meta.json 文件
meta_files = glob.glob(os.path.join(benchmark_dir, "*.meta.json"))
for meta_file in meta_files:
try:
with open(meta_file, encoding="utf-8") as f:
meta = json.load(f)
benchmarks.append(meta)
except Exception as e:
logger.error(f"Failed to load benchmark meta {meta_file}: {e}")
# 按创建时间倒序
benchmarks.sort(key=lambda x: x.get("created_at", ""), reverse=True)
return benchmarks
except Exception as e:
logger.error(f"获取评估基准列表失败: {e}")
raise
async def get_benchmark_detail(self, benchmark_id: str) -> dict[str, Any]:
"""获取评估基准详情 (包含问题列表)"""
try:
# 使用优化的查找方法
db_id, meta_file_path = self._find_benchmark_location(benchmark_id)
with open(meta_file_path, encoding="utf-8") as f:
found_meta = json.load(f)
# 加载数据文件
data_file_path = os.path.join(os.path.dirname(meta_file_path), f"{benchmark_id}.jsonl")
questions = []
if os.path.exists(data_file_path):
with open(data_file_path, encoding="utf-8") as f:
for line in f:
if line.strip():
questions.append(json.loads(line))
found_meta["questions"] = questions
return found_meta
except Exception as e:
logger.error(f"获取评估基准详情失败: {e}")
raise
async def delete_benchmark(self, benchmark_id: str) -> None:
"""删除评估基准"""
try:
# 使用优化的查找方法
_, meta_file_path = self._find_benchmark_location(benchmark_id)
data_file_path = meta_file_path.replace(".meta.json", ".jsonl")
if os.path.exists(meta_file_path):
os.remove(meta_file_path)
if os.path.exists(data_file_path):
os.remove(data_file_path)
logger.info(f"成功删除评估基准: {benchmark_id}")
except Exception as e:
logger.error(f"删除评估基准失败: {e}")
raise
async def delete_evaluation_result(self, task_id: str) -> None:
"""删除评估结果"""
try:
# 使用优化的查找方法
_, result_file_path = self._find_result_location(task_id)
# 删除结果文件
os.remove(result_file_path)
logger.info(f"成功删除评估结果: {task_id}")
except Exception as e:
logger.error(f"删除评估结果失败: {e}")
raise
async def generate_benchmark(self, db_id: str, params: dict[str, Any], created_by: str) -> dict[str, Any]:
"""自动生成评估基准 (Stub - Temporarily Disabled)"""
# 保持与之前的逻辑一致:暂不支持自动生成
# 我们可以保留接口但只返回错误,或者像之前一样进入 task 然后报错
task_id = f"gen_benchmark_{uuid.uuid4().hex[:8]}"
await tasker.enqueue(
name="生成评估基准(Disabled)",
task_type="benchmark_generation",
payload={"task_id": task_id, "db_id": db_id, "created_by": created_by, **params},
coroutine=self._generate_benchmark_task,
)
return {"task_id": task_id, "message": "基准生成任务已提交"}
async def _generate_benchmark_task(self, context: TaskContext):
"""生成任务实现"""
await context.set_progress(0, "初始化")
raise NotImplementedError("自动生成基准功能暂时不可用,请手动上传基准文件。")
async def run_evaluation(
self, db_id: str, benchmark_id: str, retrieval_config: dict[str, Any], created_by: str
) -> str:
"""运行RAG评估"""
try:
task_id = f"eval_{uuid.uuid4().hex[:8]}"
# 获取基准元数据以验证是否存在
# 使用优化的查找方法
_, meta_file_path = self._find_benchmark_location(benchmark_id)
with open(meta_file_path, encoding="utf-8") as f:
benchmark_meta = json.load(f)
# 初始化结果文件 (Status: running)
result_dir = self._get_result_dir(db_id)
result_file_path = os.path.join(result_dir, f"{task_id}.json")
initial_result = {
"id": task_id, # for compatibility
"task_id": task_id,
"benchmark_id": benchmark_id,
"db_id": db_id,
"retrieval_config": retrieval_config,
"metrics": {},
"status": "running",
"total_questions": benchmark_meta.get("question_count", 0),
"completed_questions": 0,
"started_at": datetime.utcnow().isoformat(),
"completed_at": None,
"interim_results": [],
}
with open(result_file_path, "w", encoding="utf-8") as f:
json.dump(initial_result, f, ensure_ascii=False, indent=2)
await tasker.enqueue(
name=f"RAG评估({benchmark_meta.get('name')})",
task_type="rag_evaluation",
payload={
"task_id": task_id,
"db_id": db_id,
"benchmark_id": benchmark_id,
"retrieval_config": retrieval_config,
"created_by": created_by,
},
coroutine=self._run_evaluation_task,
)
return task_id
except Exception as e:
logger.error(f"启动评估失败: {e}")
raise
async def _run_evaluation_task(self, context: TaskContext):
"""运行评估任务"""
try:
task = context._tasker._tasks.get(context.task_id)
if not task:
raise ValueError("Task not found")
payload = task.payload
task_id = payload["task_id"]
db_id = payload["db_id"]
benchmark_id = payload["benchmark_id"]
retrieval_config = payload["retrieval_config"]
# 加载基准数据
await context.set_progress(5, "加载基准数据")
# 这里我们需要重新找到 benchmark file因为 payload 里可能没有完整路径
try:
_, meta_path = self._find_benchmark_location(benchmark_id)
data_path = meta_path.replace(".meta.json", ".jsonl")
except ValueError:
raise ValueError("Benchmark file not found")
with open(meta_path, encoding="utf-8") as f:
benchmark_meta = json.load(f)
benchmark_data = []
with open(data_path, encoding="utf-8") as f:
for line in f:
if line.strip():
benchmark_data.append(json.loads(line))
# 开始评估
kb_instance = knowledge_base.get_kb(db_id)
if not kb_instance:
raise ValueError(f"Knowledge Base {db_id} not found")
if kb_instance.kb_type == "lightrag":
raise ValueError("暂不支持对 LightRAG 类型的知识库进行 RAG 评估")
# 初始化 Judge LLM
judge_llm = None
if benchmark_meta.get("has_gold_answers"):
# 优先使用配置中的 judge_llm否则回退到 answer_llm或者默认
judge_model_spec = retrieval_config.get("judge_llm") or retrieval_config.get("answer_llm")
if judge_model_spec:
try:
logger.debug(f"Initializing Judge LLM: {judge_model_spec}")
judge_llm = select_model(model_spec=judge_model_spec)
except Exception as e:
logger.error(f"Failed to load judge LLM: {e}")
total_questions = len(benchmark_data)
interim_results = []
all_retrieval_metrics = []
all_answer_metrics = []
# 更新结果文件 helper
result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json")
def update_result_file(status="running", completed=0, metrics=None, interim=None, final_score=None):
try:
if os.path.exists(result_file_path):
with open(result_file_path, encoding="utf-8") as f:
data = json.load(f)
else:
data = {} # Should have been created in run_evaluation
data["status"] = status
data["completed_questions"] = completed
if metrics:
data["metrics"] = metrics
if interim is not None:
data["interim_results"] = interim
if final_score is not None:
data["overall_score"] = final_score
if status in ["completed", "failed"]:
data["completed_at"] = datetime.utcnow().isoformat()
with open(result_file_path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
except Exception as e:
logger.error(f"Failed to update result file: {e}")
for i, question_data in enumerate(benchmark_data):
# 检查任务是否被取消
await context.raise_if_cancelled()
progress = 10 + (i / total_questions) * 80
await context.set_progress(progress, f"评估 {i + 1}/{total_questions}")
# 执行查询
query_result = await kb_instance.aquery(question_data["query"], db_id, **retrieval_config)
# 处理结果
if isinstance(query_result, dict):
generated_answer = query_result.get("answer", "")
retrieved_chunks = query_result.get("retrieved_chunks", [])
else:
retrieved_chunks = query_result if isinstance(query_result, list) else []
generated_answer = ""
# 如果没有生成的答案,但有检索结果且配置了 LLM则生成答案
if not generated_answer and retrieved_chunks and retrieval_config.get("answer_llm"):
logger.debug(f"使用 LLM {retrieval_config.get('answer_llm')} 生成答案...")
try:
# 从配置中获取 LLM
model_spec = retrieval_config["answer_llm"]
llm = select_model(model_spec=model_spec)
# 构建上下文
context_docs = []
for idx, chunk in enumerate(retrieved_chunks[:5]): # 使用前5个最相关的文档
content = chunk.get("content", "")
if content:
context_docs.append(f"文档 {idx + 1}:\n{content}")
context_text = "\\n\\n".join(context_docs)
# 构建提示词
prompt = (
f"基于以下上下文信息,请回答用户的问题。\n\n"
f"上下文信息:{context_text}\n\n"
f"用户问题:{question_data["query"]}\n\n"
"请根据上下文信息准确回答问题。如果上下文中没有相关信息,请说明。\n\n"
)
# 生成答案 - 使用 asyncio.to_thread 避免阻塞事件循环
response = await asyncio.to_thread(llm.call, prompt, stream=False)
generated_answer = response.content if response else ""
logger.debug(f"LLM 生成的答案长度: {len(generated_answer) if generated_answer else 0}")
except Exception as e:
logger.error(f"LLM 生成答案失败: {e}")
generated_answer = ""
# 计算指标
current_metrics = {}
retrieval_scores = {}
answer_scores = {}
if benchmark_meta.get("has_gold_chunks") and question_data.get("gold_chunk_ids"):
retrieval_scores = EvaluationMetricsCalculator.calculate_retrieval_metrics(
retrieved_chunks, question_data["gold_chunk_ids"]
)
current_metrics.update(retrieval_scores)
all_retrieval_metrics.append(retrieval_scores)
if benchmark_meta.get("has_gold_answers") and question_data.get("gold_answer"):
if judge_llm:
# 评判过程包含 LLM 调用,使用 asyncio.to_thread 避免阻塞
answer_scores = await asyncio.to_thread(
EvaluationMetricsCalculator.calculate_answer_metrics,
query=question_data["query"],
generated_answer=generated_answer,
gold_answer=question_data["gold_answer"],
judge_llm=judge_llm,
)
current_metrics.update(answer_scores)
all_answer_metrics.append(answer_scores)
else:
logger.warning("需要计算答案指标但未配置 Judge LLM")
interim_results.append(
{
"query": question_data["query"],
"gold_chunk_ids": question_data.get("gold_chunk_ids"),
"gold_answer": question_data.get("gold_answer"),
"generated_answer": generated_answer,
"retrieved_chunks": retrieved_chunks,
"metrics": current_metrics,
}
)
# 计算当前累计指标
current_overall_metrics = {}
if all_retrieval_metrics:
keys = all_retrieval_metrics[0].keys()
for k in keys:
current_overall_metrics[k] = sum(m.get(k, 0) for m in all_retrieval_metrics) / len(
all_retrieval_metrics
)
if all_answer_metrics:
scores = [m.get("score", 0) for m in all_answer_metrics]
current_overall_metrics["answer_correctness"] = sum(scores) / len(scores) if scores else 0.0
# 更新 Tasker 的 result 以便实时获取当前指标
await context.set_result(
{
"current_metrics": current_overall_metrics,
"completed_questions": i + 1,
"total_questions": total_questions,
}
)
# 定期更新文件 (每5个或最后一个)
if (i + 1) % 5 == 0 or (i + 1) == total_questions:
update_result_file(completed=i + 1, interim=interim_results)
# 最终计算
await context.set_progress(95, "计算最终指标")
# 汇总指标
overall_metrics = {}
# 检索指标平均值
if all_retrieval_metrics:
keys = all_retrieval_metrics[0].keys()
for k in keys:
overall_metrics[k] = sum(m.get(k, 0) for m in all_retrieval_metrics) / len(all_retrieval_metrics)
# 答案指标平均值
if all_answer_metrics:
scores = [m.get("score", 0) for m in all_answer_metrics]
overall_metrics["answer_correctness"] = sum(scores) / len(scores) if scores else 0.0
overall_score = EvaluationMetricsCalculator.calculate_overall_score(
all_retrieval_metrics, all_answer_metrics
)
overall_metrics["overall_score"] = overall_score
update_result_file(
status="completed",
completed=total_questions,
metrics=overall_metrics,
interim=interim_results,
final_score=overall_score,
)
await context.set_progress(100, "完成")
except Exception as e:
logger.error(f"Task failed: {e}")
# Try to update status to failed
try:
# Need to find the file path again or pass it around.
# Re-deriving from payload if available
if "payload" in locals():
path = os.path.join(self._get_result_dir(payload["db_id"]), f"{payload['task_id']}.json")
if os.path.exists(path):
with open(path, encoding="utf-8") as f:
d = json.load(f)
d["status"] = "failed"
d["error"] = str(e)
with open(path, "w", encoding="utf-8") as f:
json.dump(d, f, ensure_ascii=False, indent=2)
except Exception as e:
logger.error(f"Error updating result file: {e}")
pass
await context.set_message(f"Error: {str(e)}")
raise
async def get_evaluation_results(self, task_id: str) -> dict[str, Any]:
"""获取评估结果"""
try:
# 使用优化的查找方法
_, result_file_path = self._find_result_location(task_id)
with open(result_file_path, encoding="utf-8") as f:
return json.load(f)
except ValueError:
# 可能是内存中的任务状态?如果文件没创建(极早失败),检查 tasker
task = await tasker.get_task(task_id)
if task:
return {"task_id": task_id, "status": task.status, "progress": task.progress, "message": task.message}
raise ValueError(f"Result not found for task {task_id}")
async def get_evaluation_history(self, db_id: str) -> list[dict[str, Any]]:
"""获取知识库的评估历史记录"""
try:
result_dir = self._get_result_dir(db_id)
history = []
# 查找所有 .json 文件
result_files = glob.glob(os.path.join(result_dir, "*.json"))
for result_file in result_files:
try:
with open(result_file, encoding="utf-8") as f:
data = json.load(f)
# 只返回摘要信息不返回详细的interim_results
summary = {
"task_id": data.get("task_id"),
"benchmark_id": data.get("benchmark_id"),
"status": data.get("status"),
"started_at": data.get("started_at"),
"completed_at": data.get("completed_at"),
"total_questions": data.get("total_questions"),
"completed_questions": data.get("completed_questions"),
"overall_score": data.get("overall_score"),
# 也可以带上部分 metrics 摘要
"metrics": data.get("metrics"),
}
history.append(summary)
except Exception as e:
logger.error(f"Failed to load result file {result_file}: {e}")
# 按开始时间倒序
history.sort(key=lambda x: x.get("started_at", ""), reverse=True)
return history
except Exception as e:
logger.error(f"获取评估历史失败: {e}")
raise

View File

@ -0,0 +1,152 @@
"""
RAG评估指标计算工具
简化版只保留Recall/F1检索 LLM Judge答案准确性
"""
import json
import textwrap
from typing import Any
from src.utils import logger
class RetrievalMetrics:
"""检索评估指标计算"""
@staticmethod
def precision_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
"""计算Precision@K"""
if not retrieved_ids[:k]:
return 0.0
retrieved_set = set(retrieved_ids[:k])
relevant_set = set(relevant_ids)
return len(retrieved_set & relevant_set) / k
@staticmethod
def recall_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
"""计算Recall@K"""
if not relevant_ids:
return 0.0
retrieved_set = set(retrieved_ids[:k])
relevant_set = set(relevant_ids)
return len(retrieved_set & relevant_set) / len(relevant_set)
@staticmethod
def f1_score_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
"""计算F1@K"""
precision = RetrievalMetrics.precision_at_k(retrieved_ids, relevant_ids, k)
recall = RetrievalMetrics.recall_at_k(retrieved_ids, relevant_ids, k)
if precision + recall == 0:
return 0.0
return 2 * precision * recall / (precision + recall)
class AnswerMetrics:
"""答案评估指标计算"""
@staticmethod
def judge_correctness(query: str, generated_answer: str, gold_answer: str, judge_llm: Any) -> dict[str, Any]:
"""
使用LLM判断生成的答案是否正确
"""
if not generated_answer:
return {"score": 0.0, "reasoning": "未生成答案"}
if not gold_answer:
return {"score": 0.0, "reasoning": "无参考答案"}
prompt = textwrap.dedent(f"""你是一个公正的评判者请评估AI生成的答案相对于标准答案的准确性。
问题{query}
标准答案
{gold_answer}
AI生成的答案
{generated_answer}
请判断AI生成的答案是否在事实层面与标准答案一致
忽略措辞标点符号或格式上的细微差异
只关注核心事实是否准确包含
请返回以下JSON格式的结果不要包含其他文本
{{
"score": 1.0, // 如果答案正确返回 1.0 错误返回 0.0
"reasoning": "简要说明判定理由"
}}
""")
try:
response = judge_llm.call(prompt, stream=False)
content = response.content.strip()
# 尝试清理可能的 markdown 代码块
if content.startswith("```json"):
content = content[7:]
if content.endswith("```"):
content = content[:-3]
content = content.strip()
result = json.loads(content)
return {"score": float(result.get("score", 0.0)), "reasoning": result.get("reasoning", "")}
except Exception as e:
logger.error(f"LLM 评判失败: {e}")
return {"score": 0.0, "reasoning": f"评判出错: {str(e)}"}
class EvaluationMetricsCalculator:
"""综合评估指标计算器"""
@staticmethod
def calculate_retrieval_metrics(
retrieved_chunks: list[dict[str, Any]], gold_chunk_ids: list[str], k_values: list[int] = [1, 3, 5, 10]
) -> dict[str, float]:
"""计算检索指标 (Recall, F1)"""
if not retrieved_chunks or not gold_chunk_ids:
return {}
# 提取 ID
retrieved_ids = []
for chunk in retrieved_chunks:
chunk_id = chunk.get("chunk_id") or chunk.get("metadata", {}).get("chunk_id")
retrieved_ids.append(str(chunk_id) if chunk_id else "")
metrics = {}
for k in k_values:
metrics[f"recall@{k}"] = RetrievalMetrics.recall_at_k(retrieved_ids, gold_chunk_ids, k)
metrics[f"f1@{k}"] = RetrievalMetrics.f1_score_at_k(retrieved_ids, gold_chunk_ids, k)
return metrics
@staticmethod
def calculate_answer_metrics(
query: str, generated_answer: str, gold_answer: str, judge_llm: Any = None
) -> dict[str, Any]:
"""计算答案指标 (LLM Judge)"""
if not judge_llm:
return {}
return AnswerMetrics.judge_correctness(query, generated_answer, gold_answer, judge_llm)
@staticmethod
def calculate_overall_score(
retrieval_metrics_list: list[dict[str, float]], answer_metrics_list: list[dict[str, Any]]
) -> float:
"""计算整体平均分"""
total_score = 0.0
count = 0
# 简单的平均策略将所有retrieval metric的值和answer metric的score一起平均
# 用户可能希望分开看但calculate_overall_score返回一个单值。
# 计算检索平均分
for m in retrieval_metrics_list:
if m:
total_score += sum(m.values()) / len(m)
count += 1
# 计算答案平均分
for m in answer_metrics_list:
if "score" in m:
total_score += m["score"]
count += 1
return total_score / count if count > 0 else 0.0

212
test/test_manual_eval.py Normal file
View File

@ -0,0 +1,212 @@
import requests
import json
import time
import sys
import os
# 添加评估指标测试功能
def test_evaluation_metrics():
"""测试评估指标计算"""
print("\n" + "=" * 50)
print("测试评估指标计算")
print("=" * 50)
try:
from src.utils.evaluation_metrics import EvaluationMetricsCalculator
# 测试检索指标
retrieved_chunks = [
{"content": "test1", "metadata": {"chunk_id": "file_bbb147_chunk_0"}},
{"content": "test2", "metadata": {"chunk_id": "file_bbb147_chunk_1"}},
{"content": "test3", "metadata": {"chunk_id": "file_bbb147_chunk_2"}},
]
gold_chunk_ids = ["file_bbb147_chunk_0", "file_bbb147_chunk_2"]
print("测试检索指标...")
retrieval_metrics = EvaluationMetricsCalculator.calculate_retrieval_metrics(retrieved_chunks, gold_chunk_ids)
print(f"检索指标结果: {retrieval_metrics}")
# 测试答案指标
# generated_answer = "该研究以数据语义化—知识结构化—可信推理的技术主线"
# gold_answer = "该研究以数据语义化—知识结构化—可信推理的技术主线,遵循数据—知识—推理—应用的演化逻辑"
print("\n测试答案指标需要Judge LLM跳过实际LLM调用...")
# 由于需要judge_llm这里只测试检索指标
print("跳过答案指标测试需要配置Judge LLM")
print("评估指标计算测试完成!")
return True
except Exception as e:
print(f"评估指标测试失败: {e}")
return False
BASE_URL = "http://localhost:5050"
USERNAME = "zwj"
PASSWORD = "zwj12138"
DB_ID = "kb_5e343066eb4713959698ae6ca16843a0"
def get_token():
try:
resp = requests.post(f"{BASE_URL}/api/auth/token", data={"username": USERNAME, "password": PASSWORD})
if resp.status_code != 200:
print(f"Login failed: {resp.text}")
return None
return resp.json()["access_token"]
except Exception as e:
print(f"Connection failed: {e}")
return None
def main():
# 首先测试评估指标计算
test_evaluation_metrics()
token = get_token()
if not token:
sys.exit(1)
headers = {"Authorization": f"Bearer {token}"}
print(f"\n1. Trying to retrieve a real chunk from {DB_ID}...")
chunk_content = "This is a fallback test query."
chunk_id = "unknown"
try:
# 尝试查询 (Fixed payload structure)
resp = requests.post(
f"{BASE_URL}/api/knowledge/databases/{DB_ID}/query",
headers=headers,
json={"query": "人工智能", "meta": {"top_k": 5}},
)
if resp.status_code == 200:
data = resp.json()
# 结果在 result 字段中
results = data.get("result", [])
first_chunk = None
if isinstance(results, list) and len(results) > 0:
first_chunk = results[0]
elif isinstance(results, dict) and "retrieved_chunks" in results:
if results["retrieved_chunks"]:
first_chunk = results["retrieved_chunks"][0]
if first_chunk:
print(f"Found chunk: {str(first_chunk.get('content', ''))[:50]}...")
chunk_content = first_chunk.get("content", chunk_content)
chunk_id = first_chunk.get("chunk_id", chunk_id) or first_chunk.get("id", chunk_id)
else:
print("No chunks found in retrieval results. Using fallback.")
print(f"Raw results: {results}")
else:
print(f"Query failed: {resp.text}")
except Exception as e:
print(f"Query Exception: {e}")
print(f"\n2. Creating benchmark with query based on chunk: {chunk_id}")
# 构造 Benchmark 数据
benchmark_data = {
"query": chunk_content, # 使用 Chunk 内容作为查询,理论上应该能召回它自己
"gold_chunk_ids": [str(chunk_id)], # Ensure string
"gold_answer": "This is a gold answer.",
}
# 生成临时文件
with open("temp_benchmark.jsonl", "w") as f:
f.write(json.dumps(benchmark_data, ensure_ascii=False) + "\n")
print("\n3. Uploading benchmark...")
try:
files = {"file": ("temp_benchmark.jsonl", open("temp_benchmark.jsonl", "rb"), "application/jsonlines")}
# Fixed: use params for name/description
params = {"name": "Manual Eval Test", "description": "Test generated from script"}
resp = requests.post(
f"{BASE_URL}/api/evaluation/databases/{DB_ID}/benchmarks/upload",
headers=headers,
params=params,
files=files,
)
if resp.status_code != 200:
print(f"Upload failed: {resp.text}")
sys.exit(1)
benchmark = resp.json()
# Benchmark response format depends on Service impl.
# My filesystem impl returns the metadata dict directly.
# But wait, KnowledgeRouter might wrap it?
# router returns: return {"message": "上传成功", "data": result} (Assume standard wrapper)
# Let's check router impl.
# From router code read earlier:
# result = await service.upload_benchmark(...)
# return {"message": "上传成功", "data": result}
benchmark_id = benchmark["data"]["benchmark_id"]
print(f"Benchmark uploaded: {benchmark_id}")
except Exception as e:
print(f"Upload Exception: {e}")
sys.exit(1)
print("\n4. Running evaluation...")
try:
payload = {"benchmark_id": benchmark_id, "retrieval_config": {"top_k": 5}}
resp = requests.post(f"{BASE_URL}/api/evaluation/databases/{DB_ID}/run", headers=headers, json=payload)
if resp.status_code != 200:
print(f"Run evaluation failed: {resp.text}")
sys.exit(1)
task_id = resp.json()["data"]["task_id"]
print(f"Evaluation task started: {task_id}")
# 轮询状态
while True:
resp = requests.get(f"{BASE_URL}/api/evaluation/{task_id}/progress", headers=headers)
if resp.status_code != 200:
print(f"Get progress failed: {resp.text}")
break
progress = resp.json()
# The progress endpoint returns {task_id, status, ...} based on my service impl?
# Router wrapper: return {"message": "success", "data": result}
data = progress.get("data", progress) # Handle wrapper if exists
status = data["status"]
current_progress = data.get("progress", 0)
print(f"Status: {status}, Progress: {current_progress}%")
if status in ["completed", "failed"]:
break
time.sleep(2)
if status == "completed":
print("\nEvaluation Completed!")
# 获取结果
resp = requests.get(f"{BASE_URL}/api/evaluation/{task_id}/results", headers=headers)
print(json.dumps(resp.json(), indent=2, ensure_ascii=False))
else:
print("\nEvaluation Failed!")
except Exception as e:
print(f"Evaluation Exception: {e}")
# 清理
if os.path.exists("temp_benchmark.jsonl"):
os.remove("temp_benchmark.jsonl")
if __name__ == "__main__":
main()

View File

@ -163,7 +163,7 @@ export function apiPost(url, data = {}, options = {}, requiresAuth = true, respo
url,
{
method: 'POST',
body: JSON.stringify(data),
body: data instanceof FormData ? data : JSON.stringify(data),
...options
},
requiresAuth,
@ -195,7 +195,7 @@ export function apiPut(url, data = {}, options = {}, requiresAuth = true, respon
url,
{
method: 'PUT',
body: JSON.stringify(data),
body: data instanceof FormData ? data : JSON.stringify(data),
...options
},
requiresAuth,

View File

@ -162,6 +162,16 @@ export const queryApi = {
return apiAdminGet(`/api/knowledge/databases/${dbId}/query-params`)
},
/**
* 更新知识库查询参数
* @param {string} dbId - 知识库ID
* @param {Object} params - 查询参数
* @returns {Promise} - 更新结果
*/
updateKnowledgeBaseQueryParams: async (dbId, params) => {
return apiAdminPut(`/api/knowledge/databases/${dbId}/query-params`, params)
},
/**
* 生成知识库的测试问题
* @param {string} dbId - 知识库ID
@ -297,3 +307,114 @@ export const embeddingApi = {
return apiAdminGet('/api/knowledge/embedding-models/status')
}
}
// =============================================================================
// === RAG评估分组 ===
// =============================================================================
export const evaluationApi = {
/**
* 上传评估基准文件
* @param {string} dbId - 知识库ID
* @param {File} file - JSONL文件
* @param {Object} metadata - 基准元数据
* @returns {Promise} - 上传结果
*/
uploadBenchmark: async (dbId, file, metadata = {}) => {
const formData = new FormData()
formData.append('file', file)
formData.append('name', metadata.name || '')
formData.append('description', metadata.description || '')
// 调试:打印 FormData 内容
console.log('FormData 内容:')
for (let [key, value] of formData.entries()) {
console.log(key, value)
}
console.log('file type:', file ? file.type : 'undefined')
console.log('file name:', file ? file.name : 'undefined')
// 直接传递 FormDataapiAdminPost 会正确处理
return apiAdminPost(`/api/evaluation/databases/${dbId}/benchmarks/upload`, formData)
},
/**
* 获取评估基准列表
* @param {string} dbId - 知识库ID
* @returns {Promise} - 基准列表
*/
getBenchmarks: async (dbId) => {
return apiAdminGet(`/api/evaluation/databases/${dbId}/benchmarks`)
},
/**
* 获取评估基准详情
* @param {string} benchmarkId - 基准ID
* @returns {Promise} - 基准详情
*/
getBenchmark: async (benchmarkId) => {
return apiAdminGet(`/api/evaluation/benchmarks/${benchmarkId}`)
},
/**
* 删除评估基准
* @param {string} benchmarkId - 基准ID
* @returns {Promise} - 删除结果
*/
deleteBenchmark: async (benchmarkId) => {
return apiAdminDelete(`/api/evaluation/benchmarks/${benchmarkId}`)
},
/**
* 自动生成评估基准
* @param {string} dbId - 知识库ID
* @param {Object} params - 生成参数
* @param {number} params.count - 生成问题数量
* @param {boolean} params.include_answers - 是否生成答案
* @param {Object} params.llm_config - LLM配置
* @returns {Promise} - 生成结果
*/
generateBenchmark: async (dbId, params) => {
return apiAdminPost(`/api/evaluation/databases/${dbId}/benchmarks/generate`, params)
},
/**
* 运行RAG评估
* @param {string} dbId - 知识库ID
* @param {Object} params - 评估参数
* @param {string} params.benchmark_id - 基准ID
* @param {Object} params.retrieval_config - 检索配置
* @returns {Promise} - 评估任务ID
*/
runEvaluation: async (dbId, params) => {
return apiAdminPost(`/api/evaluation/databases/${dbId}/run`, params)
},
/**
* 获取评估结果
* @param {string} taskId - 任务ID
* @returns {Promise} - 评估结果
*/
getEvaluationResults: async (taskId) => {
return apiAdminGet(`/api/evaluation/${taskId}/results`)
},
/**
* 删除评估结果
* @param {string} taskId - 任务ID
* @returns {Promise} - 删除结果
*/
deleteEvaluationResult: async (taskId) => {
return apiAdminDelete(`/api/evaluation/${taskId}`)
},
/**
* 获取知识库的评估历史记录
* @param {string} dbId - 知识库ID
* @returns {Promise} - 评估历史列表
*/
getEvaluationHistory: async (dbId) => {
return apiAdminGet(`/api/evaluation/databases/${dbId}/history`)
}
}

View File

@ -0,0 +1,597 @@
<template>
<div class="evaluation-benchmarks-container">
<!-- 操作栏 -->
<div class="benchmarks-header">
<div class="header-left">
<span class="total-count">{{ benchmarks.length }} 个基准</span>
</div>
<div class="header-right">
<a-button type="primary" @click="showUploadModal">
<template #icon><UploadOutlined /></template>
上传基准
</a-button>
<a-button @click="showGenerateModal">
<template #icon><RobotOutlined /></template>
自动生成
</a-button>
</div>
</div>
<!-- 基准列表 -->
<div class="benchmarks-list">
<div v-if="!loading && benchmarks.length === 0" class="empty-state">
<div class="empty-icon">📋</div>
<div class="empty-title">暂无评估基准</div>
<div class="empty-description">上传或生成评估基准开始使用</div>
</div>
<div v-else-if="loading" class="loading-state">
<a-spin size="large" />
</div>
<div v-else class="benchmark-list-content">
<div
v-for="benchmark in benchmarks"
:key="benchmark.benchmark_id"
class="benchmark-item"
@click="previewBenchmark(benchmark)"
>
<!-- 主要内容 -->
<div class="benchmark-main">
<div class="benchmark-header">
<h4 class="benchmark-name">{{ benchmark.name }}</h4>
<div class="benchmark-actions">
<a-button type="text" size="small" @click.stop="previewBenchmark(benchmark)">
<EyeOutlined />
</a-button>
<a-button type="text" size="small" danger @click.stop="deleteBenchmark(benchmark)">
<DeleteOutlined />
</a-button>
</div>
</div>
<p class="benchmark-desc">{{ benchmark.description || '暂无描述' }}</p>
<!-- 标签区域 -->
<div class="benchmark-meta">
<div class="meta-row">
<span
v-if="benchmark.has_gold_chunks && benchmark.has_gold_answers"
class="type-badge type-both"
>
检索 + 问答
</span>
<span
v-else-if="benchmark.has_gold_chunks"
class="type-badge type-retrieval"
>
检索评估
</span>
<span
v-else-if="benchmark.has_gold_answers"
class="type-badge type-answer"
>
问答评估
</span>
<span v-else class="type-badge type-query">仅查询</span>
<span
:class="['tag', benchmark.has_gold_chunks ? 'tag-yes' : 'tag-no']"
>
{{ benchmark.has_gold_chunks ? '✓' : '✗' }} 黄金Chunk
</span>
<span
:class="['tag', benchmark.has_gold_answers ? 'tag-yes' : 'tag-no']"
>
{{ benchmark.has_gold_answers ? '✓' : '✗' }} 黄金答案
</span>
</div>
</div>
</div>
<!-- 底部信息 -->
<div class="benchmark-footer">
<span class="benchmark-time">{{ formatDate(benchmark.created_at) }}</span>
<span class="benchmark-id">{{ benchmark.benchmark_id.slice(0, 8) }}</span>
<span class="benchmark-count">{{ benchmark.question_count }} 个问题</span>
</div>
</div>
</div>
</div>
<!-- 上传模态框 -->
<BenchmarkUploadModal
v-model:visible="uploadModalVisible"
:database-id="databaseId"
@success="onUploadSuccess"
/>
<!-- 生成模态框 -->
<BenchmarkGenerateModal
v-model:visible="generateModalVisible"
:database-id="databaseId"
@success="onGenerateSuccess"
/>
<!-- 预览模态框 -->
<a-modal
v-model:open="previewModalVisible"
title="评估基准详情"
width="800px"
:footer="null"
>
<div v-if="previewData" class="preview-content">
<div class="preview-header">
<h3>{{ previewData.name }}</h3>
<div class="preview-meta">
<span class="meta-item">
<span class="meta-label">问题数:</span>
{{ previewData.question_count }}
</span>
<span class="meta-item">
<span class="meta-label">黄金Chunk:</span>
<span :class="previewData.has_gold_chunks ? 'status-yes' : 'status-no'">
{{ previewData.has_gold_chunks ? '有' : '无' }}
</span>
</span>
<span class="meta-item">
<span class="meta-label">黄金答案:</span>
<span :class="previewData.has_gold_answers ? 'status-yes' : 'status-no'">
{{ previewData.has_gold_answers ? '有' : '无' }}
</span>
</span>
</div>
</div>
<div class="preview-questions" v-if="previewQuestions.length > 0">
<h4>问题示例 (前5条)</h4>
<div class="question-list">
<div
v-for="(item, index) in previewQuestions.slice(0, 5)"
:key="index"
class="question-item"
>
<div class="question-header">
<span class="question-num">Q{{ index + 1 }}</span>
</div>
<div class="question-body">
<p class="question-text">{{ item.query }}</p>
<div v-if="item.gold_chunk_ids" class="question-chunk">
黄金Chunk: {{ item.gold_chunk_ids.slice(0, 3).join(', ') }}
<span v-if="item.gold_chunk_ids.length > 3">...等{{ item.gold_chunk_ids.length }}</span>
</div>
<div v-if="item.gold_answer" class="question-answer">
黄金答案: {{ item.gold_answer.slice(0, 150) }}
<span v-if="item.gold_answer.length > 150">...</span>
</div>
</div>
</div>
</div>
</div>
</div>
</a-modal>
</div>
</template>
<script setup>
import { ref, reactive, onMounted } from 'vue';
import { message, Modal } from 'ant-design-vue';
import {
UploadOutlined,
RobotOutlined,
EyeOutlined,
DeleteOutlined,
CheckCircleOutlined,
CloseCircleOutlined
} from '@ant-design/icons-vue';
import { evaluationApi } from '@/apis/knowledge_api';
import BenchmarkUploadModal from './modals/BenchmarkUploadModal.vue';
import BenchmarkGenerateModal from './modals/BenchmarkGenerateModal.vue';
const props = defineProps({
databaseId: {
type: String,
required: true
}
});
const emit = defineEmits(['refresh']);
//
const loading = ref(false);
const benchmarks = ref([]);
const uploadModalVisible = ref(false);
const generateModalVisible = ref(false);
const previewModalVisible = ref(false);
const previewData = ref(null);
const previewQuestions = ref([]);
//
const loadBenchmarks = async () => {
if (!props.databaseId) return;
loading.value = true;
try {
const response = await evaluationApi.getBenchmarks(props.databaseId);
if (response && response.message === 'success' && Array.isArray(response.data)) {
benchmarks.value = response.data;
} else {
console.error('响应格式不符合预期:', response);
message.error('基准数据格式错误');
}
} catch (error) {
console.error('加载评估基准失败:', error);
message.error('加载评估基准失败');
} finally {
loading.value = false;
}
};
//
const showUploadModal = () => {
uploadModalVisible.value = true;
};
//
const showGenerateModal = () => {
generateModalVisible.value = true;
};
//
const onUploadSuccess = () => {
loadBenchmarks();
message.success('基准上传成功');
//
emit('refresh');
};
//
const onGenerateSuccess = () => {
loadBenchmarks();
message.success('基准生成成功');
//
emit('refresh');
};
//
const previewBenchmark = async (benchmark) => {
try {
const response = await evaluationApi.getBenchmark(benchmark.benchmark_id);
if (response.message === 'success') {
previewData.value = response.data;
previewQuestions.value = response.data.questions || [];
previewModalVisible.value = true;
}
} catch (error) {
console.error('获取基准详情失败:', error);
message.error('获取基准详情失败');
}
};
//
const deleteBenchmark = (benchmark) => {
Modal.confirm({
title: '确认删除',
content: `确定要删除评估基准"${benchmark.name}"吗?此操作不可恢复。`,
okText: '确定',
cancelText: '取消',
onOk: async () => {
try {
const response = await evaluationApi.deleteBenchmark(benchmark.benchmark_id);
if (response.message === 'success') {
message.success('删除成功');
loadBenchmarks();
}
} catch (error) {
console.error('删除基准失败:', error);
message.error('删除基准失败');
}
}
});
};
//
const formatDate = (dateStr) => {
if (!dateStr) return '-';
const date = new Date(dateStr);
return date.toLocaleDateString('zh-CN', {
year: 'numeric',
month: '2-digit',
day: '2-digit',
hour: '2-digit',
minute: '2-digit'
});
};
//
onMounted(() => {
loadBenchmarks();
});
</script>
<style lang="less" scoped>
.evaluation-benchmarks-container {
height: 100%;
display: flex;
flex-direction: column;
}
.benchmarks-header {
display: flex;
justify-content: space-between;
align-items: center;
padding: 12px 0;
margin-bottom: 12px;
.total-count {
font-size: 13px;
color: var(--color-text-secondary);
}
.header-right {
display: flex;
gap: 8px;
}
}
.benchmarks-list {
flex: 1;
overflow-y: auto;
}
.benchmark-list-content {
display: flex;
flex-direction: column;
gap: 6px;
}
.benchmark-item {
padding: 12px;
border: 1px solid var(--gray-200);
border-radius: 8px;
background: var(--color-bg-container);
cursor: pointer;
transition: all 0.2s;
&:hover {
border-color: var(--color-primary-100);
box-shadow: 0 1px 2px var(--shadow-2);
background: var(--gray-10);
}
&:active {
transform: scale(0.998);
}
}
.benchmark-main {
margin-bottom: 8px;
}
.benchmark-header {
display: flex;
justify-content: space-between;
align-items: flex-start;
margin-bottom: 6px;
.benchmark-name {
margin: 0;
font-size: 15px;
font-weight: 600;
color: var(--gray-1000);
flex: 1;
}
.benchmark-actions {
display: flex;
gap: 4px;
}
}
.benchmark-desc {
margin: 0 0 8px;
font-size: 13px;
color: var(--color-text-secondary);
line-height: 1.5;
display: -webkit-box;
-webkit-line-clamp: 2;
-webkit-box-orient: vertical;
overflow: hidden;
}
.benchmark-meta {
margin-bottom: 8px;
}
.meta-row {
display: flex;
align-items: center;
gap: 6px;
flex-wrap: wrap;
}
.tag {
display: inline-flex;
align-items: center;
padding: 1px 6px;
border-radius: 3px;
font-size: 11px;
font-weight: 500;
background: var(--main-50);
color: var(--color-text-tertiary);
&.tag-yes {
// background: var(--color-success-50);
color: var(--main-500);
}
}
.type-badge {
padding: 1px 6px;
border-radius: 3px;
font-size: 11px;
font-weight: 500;
&.type-both {
background: var(--color-accent-50);
color: var(--color-accent-700);
}
&.type-retrieval {
background: var(--color-info-50);
color: var(--color-info-700);
}
&.type-answer {
background: var(--color-warning-50);
color: var(--color-warning-700);
}
&.type-query {
background: var(--gray-100);
color: var(--gray-700);
}
}
.benchmark-footer {
display: flex;
justify-content: space-between;
align-items: center;
padding-top: 8px;
border-top: 1px solid var(--gray-150);
font-size: 11px;
color: var(--color-text-tertiary);
.benchmark-id {
font-family: 'SF Mono', 'Monaco', 'Consolas', monospace;
}
.benchmark-count {
color: var(--color-primary-700);
font-weight: 500;
}
}
.empty-state {
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
height: 300px;
text-align: center;
.empty-icon {
font-size: 48px;
margin-bottom: 16px;
opacity: 0.5;
}
.empty-title {
font-size: 18px;
font-weight: 500;
color: var(--gray-900);
margin-bottom: 8px;
}
.empty-description {
font-size: 14px;
color: var(--gray-600);
}
}
.loading-state {
display: flex;
align-items: center;
justify-content: center;
height: 200px;
}
//
.preview-content {
.preview-header {
margin-bottom: 24px;
padding-bottom: 16px;
border-bottom: 1px solid var(--gray-200);
h3 {
margin: 0 0 12px;
font-size: 20px;
font-weight: 600;
color: var(--gray-1000);
}
.preview-meta {
display: flex;
gap: 24px;
.meta-item {
font-size: 14px;
.meta-label {
color: var(--color-text-tertiary);
margin-right: 4px;
}
.status-yes {
color: var(--color-success-700);
font-weight: 500;
}
.status-no {
color: var(--color-text-tertiary);
}
}
}
}
.preview-questions {
h4 {
margin: 0 0 16px;
font-size: 16px;
font-weight: 600;
color: var(--gray-900);
}
.question-list {
display: flex;
flex-direction: column;
gap: 16px;
}
.question-item {
padding: 16px;
background: var(--gray-50);
border-radius: 8px;
border: 1px solid var(--gray-200);
.question-header {
margin-bottom: 8px;
.question-num {
font-size: 14px;
font-weight: 600;
color: var(--gray-700);
}
}
.question-body {
.question-text {
margin: 0 0 12px;
font-size: 14px;
line-height: 1.6;
color: var(--gray-800);
}
.question-chunk,
.question-answer {
margin: 8px 0;
font-size: 13px;
color: var(--gray-600);
}
}
}
}
}
</style>

View File

@ -82,7 +82,7 @@
<!-- 重新分块参数配置模态框 -->
<a-modal
v-model:visible="rechunkModalVisible"
v-model:open="rechunkModalVisible"
title="重新分块参数配置"
:confirm-loading="rechunkModalLoading"
width="600px"

File diff suppressed because it is too large Load Diff

View File

@ -46,11 +46,12 @@
</a-select>
<a-select
v-else-if="param.type === 'boolean'"
v-model:value="meta[param.key]"
:value="computedMeta[param.key]"
@update:value="(value) => updateMeta(param.key, value)"
style="width: 100%;"
>
<a-select-option :value="true">启用</a-select-option>
<a-select-option :value="false">关闭</a-select-option>
<a-select-option value="true">启用</a-select-option>
<a-select-option value="false">关闭</a-select-option>
</a-select>
<a-input-number
v-else-if="param.type === 'number'"
@ -89,6 +90,32 @@ const error = ref('');
const queryParams = ref([]);
const meta = reactive({});
//
const computedMeta = computed(() => {
const result = {};
for (const key in meta) {
const param = queryParams.value.find(p => p.key === key);
if (param?.type === 'boolean') {
// select
result[key] = meta[key].toString();
} else {
result[key] = meta[key];
}
}
return result;
});
//
const updateMeta = (key, value) => {
const param = queryParams.value.find(p => p.key === key);
if (param?.type === 'boolean') {
//
meta[key] = value === 'true';
} else {
meta[key] = value;
}
};
//
const loadQueryParams = async () => {
try {
@ -104,7 +131,12 @@ const loadQueryParams = async () => {
// meta
queryParams.value.forEach(param => {
if (param.default !== undefined) {
meta[param.key] = param.default;
// 使
if (param.type === 'boolean') {
meta[param.key] = Boolean(param.default);
} else {
meta[param.key] = param.default;
}
}
});
@ -128,6 +160,17 @@ const loadSavedConfig = () => {
if (saved) {
try {
const savedConfig = JSON.parse(saved);
//
queryParams.value.forEach(param => {
if (param.type === 'boolean' && savedConfig[param.key] !== undefined) {
//
if (typeof savedConfig[param.key] === 'string') {
savedConfig[param.key] = savedConfig[param.key] === 'true';
}
}
});
Object.assign(meta, savedConfig);
} catch (e) {
console.warn('Failed to parse saved config:', e);
@ -150,16 +193,29 @@ const resetToDefaults = () => {
};
//
const saveConfig = () => {
const saveConfig = async () => {
// include_distances true
meta['include_distances'] = true;
// localStorage
localStorage.setItem(`search-config-${props.databaseId}`, JSON.stringify(meta));
//
try {
const { knowledgeApi } = await import('@/apis/knowledge_api');
const response = await knowledgeApi.updateKnowledgeBaseQueryParams(props.databaseId, meta);
if (response.message === 'success') {
message.success('配置已保存到知识库');
} else {
message.warning('配置已保存到本地,但同步到知识库失败');
}
} catch (error) {
console.error('保存配置到知识库失败:', error);
message.warning('配置已保存到本地,但同步到知识库失败');
}
// store
Object.assign(store.meta, meta);
message.success('配置已保存');
};
//

View File

@ -59,7 +59,7 @@
<a-progress
:percent="Math.round(task.progress || 0)"
:status="progressStatus(task.status)"
stroke-width=6
:stroke-width="6"
/>
<!-- <span class="task-card-progress-value">{{ Math.round(task.progress || 0) }}%</span> -->
</div>

View File

@ -0,0 +1,318 @@
<template>
<a-modal
v-model:open="visible"
title="自动生成评估基准"
width="600px"
:confirmLoading="generating"
@ok="handleGenerate"
@cancel="handleCancel"
>
<a-form
ref="formRef"
:model="formState"
:rules="rules"
layout="vertical"
>
<a-form-item label="基准名称" name="name">
<a-input
v-model:value="formState.name"
placeholder="请输入评估基准名称"
/>
</a-form-item>
<a-form-item label="描述" name="description">
<a-textarea
v-model:value="formState.description"
placeholder="请输入评估基准描述(可选)"
:rows="3"
</a-textarea>
</a-form-item>
<a-form-item label="生成参数" name="params">
<a-row :gutter="16">
<a-col :span="12">
<a-form-item label="问题数量" name="count" :labelCol="{ span: 24 }" :wrapperCol="{ span: 24 }">
<a-input-number
v-model:value="formState.count"
:min="1"
:max="100"
style="width: 100%"
placeholder="生成问题数量"
/>
</a-form-item>
</a-col>
<a-col :span="12">
<a-form-item label="每个问题生成答案数量" name="answers_per_question" :labelCol="{ span: 24 }" :wrapperCol="{ span: 24 }">
<a-input-number
v-model:value="formState.answers_per_question"
:min="0"
:max="3"
style="width: 100%"
placeholder="每个问题生成答案数量"
/>
</a-form-item>
</a-col>
</a-row>
<a-row :gutter="16">
<a-col :span="24">
<a-form-item label="黄金答案生成策略" name="answer_generation_strategy" :labelCol="{ span: 24 }" :wrapperCol="{ span: 24 }">
<a-select
v-model:value="formState.answer_generation_strategy"
placeholder="选择答案生成策略"
>
<a-select-option value="chunk_based">基于检索的文档块生成答案</a-select-option>
<a-select-option value="llm_based">基于LLM的知识理解生成答案</a-select-option>
</a-select>
</a-form-item>
</a-col>
</a-row>
<a-row :gutter="16">
<a-col :span="12">
<a-form-item label="采样数量" name="sample_count" :labelCol="{ span: 24 }" :wrapperCol="{ span: 24 }">
<a-input-number
v-model:value="formState.sample_count"
:min="1"
:max="1000"
style="width: 100%"
placeholder="采样文档块数量"
/>
</a-form-item>
</a-col>
<a-col :span="12">
<a-form-item label="相似度阈值" name="similarity_threshold" :labelCol="{ span: 24 }" :wrapperCol="{ span: 24 }">
<a-input-number
v-model:value="formState.similarity_threshold"
:min="0"
:max="1"
:step="0.1"
style="width: 100%"
placeholder="文档块相似度阈值"
/>
</a-form-item>
</a-col>
</a-row>
</a-form-item>
<a-form-item label="LLM配置" name="llm_config">
<a-card size="small" title="配置参数">
<a-form-item label="LLM模型配置" name="llm_model">
<ModelSelectorComponent
:model_spec="formState.llm_model_spec"
placeholder="选择用于生成问题的LLM模型"
size="default"
@select-model="handleSelectLLMModel"
/>
</a-form-item>
<a-row :gutter="16">
<a-col :span="12">
<a-form-item label="Temperature" name="temperature" :labelCol="{ span: 24 }" :wrapperCol="{ span: 24 }">
<a-input-number
v-model:value="formState.llm_config.temperature"
:min="0"
:max="2"
:step="0.1"
style="width: 100%"
placeholder="控制生成内容的随机性"
/>
</a-form-item>
</a-col>
<a-col :span="12">
<a-form-item label="Max Tokens" name="max_tokens" :labelCol="{ span: 24 }" :wrapperCol="{ span: 24 }">
<a-input-number
v-model:value="formState.llm_config.max_tokens"
:min="100"
:max="4000"
:step="100"
style="width: 100%"
placeholder="生成内容的最大长度"
/>
</a-form-item>
</a-col>
</a-row>
</a-card>
</a-form-item>
<a-form-item>
<a-alert
message="生成说明"
type="info"
show-icon
>
<template #description>
<ul>
<li>系统将从知识库中随机采样文档块</li>
<li>基于采样的文档块生成相关查询问题</li>
<li>可选择为每个问题生成对应的黄金答案</li>
<li>生成过程可能需要几分钟请耐心等待</li>
</ul>
</template>
</a-alert>
</a-form-item>
</a-form>
</a-modal>
</template>
<script setup>
import { ref, reactive, computed, watch } from 'vue';
import { message } from 'ant-design-vue';
import { evaluationApi } from '@/apis/knowledge_api';
import ModelSelectorComponent from '@/components/ModelSelectorComponent.vue';
const props = defineProps({
visible: {
type: Boolean,
default: false
},
databaseId: {
type: String,
required: true
}
});
const emit = defineEmits(['update:visible', 'success']);
//
const formRef = ref();
const generating = ref(false);
const formState = reactive({
name: '',
description: '',
count: 10,
answers_per_question: 1,
answer_generation_strategy: 'chunk_based',
sample_count: 100,
similarity_threshold: 0.7,
llm_model_spec: '',
llm_config: {
model: '',
temperature: 0.7,
max_tokens: 1000
}
});
//
const rules = {
name: [
{ required: true, message: '请输入基准名称', trigger: 'blur' },
{ min: 2, max: 100, message: '基准名称长度应在2-100个字符之间', trigger: 'blur' }
],
count: [
{ required: true, message: '请输入生成问题数量', trigger: 'blur' }
]
};
// visible
const visible = computed({
get: () => props.visible,
set: (val) => emit('update:visible', val)
});
//
const handleGenerate = async () => {
try {
//
await formRef.value.validate();
generating.value = true;
const params = {
name: formState.name,
description: formState.description,
count: formState.count,
answers_per_question: formState.answers_per_question,
answer_generation_strategy: formState.answer_generation_strategy,
sample_count: formState.sample_count,
similarity_threshold: formState.similarity_threshold,
llm_config: {
...formState.llm_config,
model_spec: formState.llm_model_spec
}
};
const response = await evaluationApi.generateBenchmark(props.databaseId, params);
if (response.message === 'success') {
message.success('生成任务已提交,请稍后查看结果');
handleCancel();
emit('success');
} else {
message.error(response.message || '生成失败');
}
} catch (error) {
console.error('生成失败:', error);
message.error('生成失败');
} finally {
generating.value = false;
}
};
//
const handleCancel = () => {
visible.value = false;
resetForm();
};
//
const resetForm = () => {
formRef.value?.resetFields();
Object.assign(formState, {
name: '',
description: '',
count: 10,
answers_per_question: 1,
answer_generation_strategy: 'chunk_based',
sample_count: 100,
similarity_threshold: 0.7,
llm_model_spec: '',
llm_config: {
model: '',
temperature: 0.7,
max_tokens: 1000
}
});
generating.value = false;
};
// LLM
const handleSelectLLMModel = (modelSpec) => {
formState.llm_model_spec = modelSpec;
formState.llm_config.model = modelSpec;
};
// visible
watch(visible, (val) => {
if (!val) {
resetForm();
}
});
</script>
<style lang="less" scoped>
:deep(.ant-card) {
.ant-card-head {
min-height: auto;
padding: 0 12px;
border-bottom: 1px solid var(--gray-200);
.ant-card-head-title {
font-size: 14px;
padding: 8px 0;
}
}
}
:deep(.ant-alert) {
ul {
margin: 8px 0;
padding-left: 20px;
li {
margin-bottom: 4px;
}
}
}
</style>

View File

@ -0,0 +1,284 @@
<template>
<a-modal
v-model:open="visible"
title="上传评估基准"
width="600px"
:confirmLoading="uploading"
@ok="handleUpload"
@cancel="handleCancel"
>
<a-form
ref="formRef"
:model="formState"
:rules="rules"
layout="vertical"
>
<a-form-item label="基准名称" name="name">
<a-input
v-model:value="formState.name"
placeholder="请输入评估基准名称"
/>
</a-form-item>
<a-form-item label="描述" name="description">
<a-textarea
v-model:value="formState.description"
placeholder="请输入评估基准描述(可选)"
:rows="3"
/>
</a-form-item>
<a-form-item label="基准文件" name="file">
<a-upload-dragger
v-model:fileList="fileList"
name="file"
:multiple="false"
accept=".jsonl"
:before-upload="beforeUpload"
@remove="handleRemove"
>
<p class="ant-upload-text">
<FileTextOutlined />
点击或拖拽文件到此区域上传
</p>
<p class="ant-upload-hint">
仅支持 JSONL 格式文件.jsonl
</p>
</a-upload-dragger>
</a-form-item>
<a-form-item label="文件格式说明">
<a-alert
message="JSONL文件格式要求"
type="info"
show-icon
>
<template #description>
<div class="format-info">
<p>每行一个JSON对象包含以下字段</p>
<ul>
<li><code>query</code> (必需): 查询问题</li>
<li><code>gold_chunk_ids</code> (可选): 黄金文档块ID列表</li>
<li><code>gold_answer</code> (可选): 黄金答案</li>
</ul>
<p>示例</p>
<pre class="format-example">{"query": "什么是人工智能?", "gold_chunk_ids": ["chunk_001"], "gold_answer": "人工智能是..."}</pre>
</div>
</template>
</a-alert>
</a-form-item>
</a-form>
</a-modal>
</template>
<script setup>
import { ref, reactive, computed, watch } from 'vue';
import { message } from 'ant-design-vue';
import { FileTextOutlined } from '@ant-design/icons-vue';
import { evaluationApi } from '@/apis/knowledge_api';
const props = defineProps({
visible: {
type: Boolean,
default: false
},
databaseId: {
type: String,
required: true
}
});
const emit = defineEmits(['update:visible', 'success']);
//
const formRef = ref();
const fileList = ref([]);
const uploading = ref(false);
const formState = reactive({
name: '',
description: '',
file: null
});
//
const rules = {
name: [
{ required: true, message: '请输入基准名称', trigger: 'blur' },
{ min: 2, max: 100, message: '基准名称长度应在2-100个字符之间', trigger: 'blur' }
],
file: [
{ required: true, message: '请选择基准文件', trigger: 'change' }
]
};
// visible
const visible = computed({
get: () => props.visible,
set: (val) => emit('update:visible', val)
});
//
const beforeUpload = (file) => {
//
if (!file.name.endsWith('.jsonl')) {
message.error('仅支持 JSONL 格式文件');
return false;
}
// 10MB
const isLt10M = file.size / 1024 / 1024 < 10;
if (!isLt10M) {
message.error('文件大小不能超过 10MB');
return false;
}
//
const reader = new FileReader();
reader.onload = (e) => {
try {
const content = e.target.result;
const lines = content.trim().split('\n');
//
if (lines.length === 0) {
message.error('文件不能为空');
return false;
}
// JSON
for (let i = 0; i < Math.min(5, lines.length); i++) {
const line = lines[i].trim();
if (line) {
JSON.parse(line);
}
}
formState.file = file;
} catch (error) {
message.error('文件格式错误请检查JSONL格式');
return false;
}
};
reader.readAsText(file);
//
return false;
};
//
const handleRemove = () => {
formState.file = null;
};
//
const handleUpload = async () => {
try {
//
await formRef.value.validate();
if (!formState.file) {
message.error('请选择基准文件');
return;
}
uploading.value = true;
const response = await evaluationApi.uploadBenchmark(
props.databaseId,
formState.file,
{
name: formState.name,
description: formState.description
}
);
if (response.message === 'success') {
message.success('上传成功');
handleCancel();
emit('success');
} else {
message.error(response.message || '上传失败');
}
} catch (error) {
console.error('上传失败:', error);
message.error('上传失败');
} finally {
uploading.value = false;
}
};
//
const handleCancel = () => {
visible.value = false;
resetForm();
};
//
const resetForm = () => {
formRef.value?.resetFields();
fileList.value = [];
formState.file = null;
uploading.value = false;
};
// visible
watch(visible, (val) => {
if (!val) {
resetForm();
}
});
</script>
<style lang="less" scoped>
.format-info {
p {
margin-bottom: 8px;
}
ul {
margin: 8px 0;
padding-left: 20px;
li {
margin-bottom: 4px;
}
}
code {
background-color: var(--gray-100);
padding: 2px 4px;
border-radius: 3px;
font-family: 'Monaco', 'Consolas', monospace;
font-size: 13px;
}
.format-example {
background-color: var(--gray-100);
padding: 8px;
border-radius: 4px;
margin-top: 8px;
overflow-x: auto;
font-family: 'Monaco', 'Consolas', monospace;
font-size: 12px;
line-height: 1.4;
}
}
:deep(.ant-upload-dragger) {
.ant-upload-text {
font-size: 16px;
color: var(--gray-700);
.anticon {
font-size: 48px;
color: var(--gray-400);
margin-bottom: 16px;
}
}
.ant-upload-hint {
color: var(--gray-500);
}
}
</style>

View File

@ -44,6 +44,30 @@
<a-tab-pane key="config" tab="检索配置">
<SearchConfigTab :database-id="databaseId" />
</a-tab-pane>
<a-tab-pane key="evaluation" tab="RAG评估Beta">
<RAGEvaluationTab
v-if="databaseId"
:database-id="databaseId"
@switch-to-benchmarks="activeTab = 'benchmarks'"
/>
</a-tab-pane>
<a-tab-pane key="benchmarks" tab="评估基准管理Beta">
<div class="benchmark-management-container">
<div class="benchmark-content">
<EvaluationBenchmarks
v-if="databaseId"
:database-id="databaseId"
@benchmark-selected="(benchmark) => {
//
activeTab = 'evaluation';
}"
@refresh="() => {
//
}"
/>
</div>
</div>
</a-tab-pane>
</a-tabs>
</div>
</div>
@ -62,6 +86,8 @@ import KnowledgeGraphSection from '@/components/KnowledgeGraphSection.vue';
import QuerySection from '@/components/QuerySection.vue';
import SearchConfigTab from '@/components/SearchConfigTab.vue';
import MindMapSection from '@/components/MindMapSection.vue';
import RAGEvaluationTab from '@/components/RAGEvaluationTab.vue';
import EvaluationBenchmarks from '@/components/EvaluationBenchmarks.vue';
const route = useRoute();
const store = useDatabaseStore();
@ -529,4 +555,19 @@ const handleMouseUp = () => {
overflow: hidden;
}
}
//
.benchmark-management-container {
height: 100%;
background: var(--gray-0);
display: flex;
flex-direction: column;
}
.benchmark-content {
flex: 1;
overflow: hidden;
min-height: 0;
padding: 12px 16px;
}
</style>