diff --git a/docs/latest/changelog/roadmap.md b/docs/latest/changelog/roadmap.md index 07b4aac9..8bcbb7ea 100644 --- a/docs/latest/changelog/roadmap.md +++ b/docs/latest/changelog/roadmap.md @@ -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】当前的检索配置仅保存在本地,实际并没有保存在知识库的元数据中 ### 新增 - 优化知识库详情页面,更加简洁清晰 diff --git a/server/routers/__init__.py b/server/routers/__init__.py index 9698550e..8649620c 100644 --- a/server/routers/__init__.py +++ b/server/routers/__init__.py @@ -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/* diff --git a/server/routers/evaluation_router.py b/server/routers/evaluation_router.py new file mode 100644 index 00000000..913eb7b6 --- /dev/null +++ b/server/routers/evaluation_router.py @@ -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)}") diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index 211d83d3..e579f204 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -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, # 是否包含同名文件标志 } diff --git a/server/services/tasker.py b/server/services/tasker.py index 657d1fa2..1bdfdb35 100644 --- a/server/services/tasker.py +++ b/server/services/tasker.py @@ -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, } diff --git a/src/knowledge/implementations/milvus.py b/src/knowledge/implementations/milvus.py index d2d70dae..85688bd1 100644 --- a/src/knowledge/implementations/milvus.py +++ b/src/knowledge/implementations/milvus.py @@ -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 diff --git a/src/knowledge/indexing.py b/src/knowledge/indexing.py index 49f10be4..bbaafc1a 100644 --- a/src/knowledge/indexing.py +++ b/src/knowledge/indexing.py @@ -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: diff --git a/src/knowledge/manager.py b/src/knowledge/manager.py index 47595e47..f19ed6c0 100644 --- a/src/knowledge/manager.py +++ b/src/knowledge/manager.py @@ -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) diff --git a/src/knowledge/utils/kb_utils.py b/src/knowledge/utils/kb_utils.py index d220d4bf..ac9411ed 100644 --- a/src/knowledge/utils/kb_utils.py +++ b/src/knowledge/utils/kb_utils.py @@ -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] diff --git a/src/services/evaluation_service.py b/src/services/evaluation_service.py new file mode 100644 index 00000000..0ba5eb2b --- /dev/null +++ b/src/services/evaluation_service.py @@ -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 diff --git a/src/utils/evaluation_metrics.py b/src/utils/evaluation_metrics.py new file mode 100644 index 00000000..19292fc4 --- /dev/null +++ b/src/utils/evaluation_metrics.py @@ -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 diff --git a/test/test_manual_eval.py b/test/test_manual_eval.py new file mode 100644 index 00000000..ce6c0d87 --- /dev/null +++ b/test/test_manual_eval.py @@ -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() diff --git a/web/src/apis/base.js b/web/src/apis/base.js index dfd78f35..e280e32f 100644 --- a/web/src/apis/base.js +++ b/web/src/apis/base.js @@ -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, diff --git a/web/src/apis/knowledge_api.js b/web/src/apis/knowledge_api.js index 828c99ad..195fd736 100644 --- a/web/src/apis/knowledge_api.js +++ b/web/src/apis/knowledge_api.js @@ -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') + + // 直接传递 FormData,apiAdminPost 会正确处理 + 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`) + } +} diff --git a/web/src/components/EvaluationBenchmarks.vue b/web/src/components/EvaluationBenchmarks.vue new file mode 100644 index 00000000..0c552580 --- /dev/null +++ b/web/src/components/EvaluationBenchmarks.vue @@ -0,0 +1,597 @@ + + + + + \ No newline at end of file diff --git a/web/src/components/FileTable.vue b/web/src/components/FileTable.vue index 90e86e93..5b6430d4 100644 --- a/web/src/components/FileTable.vue +++ b/web/src/components/FileTable.vue @@ -82,7 +82,7 @@ +
+ +
+
+ +
+ + + + {{ benchmark.name }}({{ benchmark.question_count }} 个问题) + + +
+
+
+ + + 开始评估 + +
+
+ + +
+ +
+ + + + + + + + + + + + + +
+ + + +
+
+ + + +
+ +

正在加载评估结果...

+
+ +
+ + + {{ selectedResult.task_id }} + + + {{ getStatusText(selectedResult.status) }} + + + + + + {{ (selectedResult.overall_score * 100).toFixed(1) }}% + + + - + + {{ selectedResult.total_questions }} + {{ selectedResult.completed_questions }} + + + {{ formatDuration(evaluationStats.totalDuration) }} + + - + + + + +
+

整体评估报告

+ + + +
+
+ {{ getMetricTitle(key) }} + + {{ formatMetricValue(value) }} + +
+
+ - +
+
+ + +
+
+ 正确答案数 + {{ evaluationStats.correctAnswers || 0 }} / {{ evaluationStats.totalQuestions || 0 }} +
+
+ 准确率 + + {{ (evaluationStats.answerAccuracy * 100).toFixed(1) }}% + +
+
+
+
+
+
+ + +

详细评估结果

+ + + +
+ +
+ + 查看基本信息 + +
+
+ + + + + \ No newline at end of file diff --git a/web/src/components/SearchConfigTab.vue b/web/src/components/SearchConfigTab.vue index 3d67047a..23356a52 100644 --- a/web/src/components/SearchConfigTab.vue +++ b/web/src/components/SearchConfigTab.vue @@ -46,11 +46,12 @@ - 启用 - 关闭 + 启用 + 关闭 { + 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('配置已保存'); }; // 组件挂载时加载数据 diff --git a/web/src/components/TaskCenterDrawer.vue b/web/src/components/TaskCenterDrawer.vue index db2f849e..39abcc27 100644 --- a/web/src/components/TaskCenterDrawer.vue +++ b/web/src/components/TaskCenterDrawer.vue @@ -59,7 +59,7 @@ diff --git a/web/src/components/modals/BenchmarkGenerateModal.vue b/web/src/components/modals/BenchmarkGenerateModal.vue new file mode 100644 index 00000000..c3713356 --- /dev/null +++ b/web/src/components/modals/BenchmarkGenerateModal.vue @@ -0,0 +1,318 @@ + + + + + \ No newline at end of file diff --git a/web/src/components/modals/BenchmarkUploadModal.vue b/web/src/components/modals/BenchmarkUploadModal.vue new file mode 100644 index 00000000..d5354ed2 --- /dev/null +++ b/web/src/components/modals/BenchmarkUploadModal.vue @@ -0,0 +1,284 @@ + + + + + \ No newline at end of file diff --git a/web/src/views/DataBaseInfoView.vue b/web/src/views/DataBaseInfoView.vue index 00789d3d..71f50cad 100644 --- a/web/src/views/DataBaseInfoView.vue +++ b/web/src/views/DataBaseInfoView.vue @@ -44,6 +44,30 @@ + + + + +
+
+ +
+
+
@@ -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; +}