diff --git a/backend/package/yuxi/knowledge/base.py b/backend/package/yuxi/knowledge/base.py index 2bb7d03d..ea40243b 100644 --- a/backend/package/yuxi/knowledge/base.py +++ b/backend/package/yuxi/knowledge/base.py @@ -1428,13 +1428,11 @@ class KnowledgeBase(ABC): return retrievers async def _load_metadata(self) -> None: - from yuxi.repositories.evaluation_repository import EvaluationRepository from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository kb_repo = KnowledgeBaseRepository() file_repo = KnowledgeFileRepository() - eval_repo = EvaluationRepository() databases = [kb for kb in await kb_repo.get_all() if kb.kb_type == self.kb_type] self.databases_meta = { @@ -1484,26 +1482,6 @@ class KnowledgeBase(ABC): } self.benchmarks_meta = {} - for kb in databases: - benchmarks = await eval_repo.list_benchmarks(kb.db_id) - if not benchmarks: - continue - self.benchmarks_meta[kb.db_id] = {} - for bench in benchmarks: - self.benchmarks_meta[kb.db_id][bench.benchmark_id] = { - "id": bench.benchmark_id, - "benchmark_id": bench.benchmark_id, - "name": bench.name, - "description": bench.description, - "db_id": bench.db_id, - "question_count": bench.question_count, - "has_gold_chunks": bench.has_gold_chunks, - "has_gold_answers": bench.has_gold_answers, - "benchmark_file": bench.data_file_path, - "created_by": bench.created_by, - "created_at": utc_isoformat(bench.created_at) if bench.created_at else None, - "updated_at": utc_isoformat(bench.updated_at) if bench.updated_at else None, - } logger.info(f"Loaded {self.kb_type} metadata from database for {len(self.databases_meta)} databases") await self._fill_missing_file_sizes() @@ -1556,13 +1534,11 @@ class KnowledgeBase(ABC): await self._persist_file(file_id) async def _save_metadata(self) -> None: - from yuxi.repositories.evaluation_repository import EvaluationRepository from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository kb_repo = KnowledgeBaseRepository() file_repo = KnowledgeFileRepository() - eval_repo = EvaluationRepository() self._normalize_metadata_state() @@ -1608,23 +1584,6 @@ class KnowledgeBase(ABC): }, ) - for db_id, benchmarks in self.benchmarks_meta.items(): - for benchmark_id, meta in benchmarks.items(): - existing = await eval_repo.get_benchmark(benchmark_id) - payload = { - "benchmark_id": benchmark_id, - "db_id": db_id, - "name": meta.get("name") or benchmark_id, - "description": meta.get("description"), - "question_count": int(meta.get("question_count") or 0), - "has_gold_chunks": bool(meta.get("has_gold_chunks")), - "has_gold_answers": bool(meta.get("has_gold_answers")), - "data_file_path": meta.get("benchmark_file"), - "created_by": str(meta.get("created_by")) if meta.get("created_by") else None, - } - if existing is None: - await eval_repo.create_benchmark(payload) - async def _persist_file(self, file_id: str) -> None: """只保存单个文件到数据库,避免全量遍历""" from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository diff --git a/backend/package/yuxi/knowledge/eval/benchmark_generation.py b/backend/package/yuxi/knowledge/eval/benchmark_generation.py index 88f694f5..6abadfb8 100644 --- a/backend/package/yuxi/knowledge/eval/benchmark_generation.py +++ b/backend/package/yuxi/knowledge/eval/benchmark_generation.py @@ -1,3 +1,4 @@ +import asyncio import json import random from collections.abc import AsyncIterator, Callable @@ -8,6 +9,9 @@ import json_repair from yuxi.models import select_model from yuxi.utils import logger +DEFAULT_BENCHMARK_GENERATION_CONCURRENCY = 10 +MAX_BENCHMARK_GENERATION_CONCURRENCY = 20 + async def collect_kb_chunks(kb_instance: Any, db_id: str) -> list[dict[str, Any]]: chunks = [] @@ -34,6 +38,12 @@ def clamp_neighbors_count(neighbors_count: int) -> int: return min(max(neighbors_count, 0), 10) +def normalize_generation_concurrency_count(value: Any) -> int: + if value in (None, ""): + return DEFAULT_BENCHMARK_GENERATION_CONCURRENCY + return min(max(1, int(value)), MAX_BENCHMARK_GENERATION_CONCURRENCY) + + def _is_anchor_chunk(candidate: dict[str, Any], anchor_chunk: dict[str, Any]) -> bool: metadata = candidate.get("metadata") or {} candidate_id = metadata.get("chunk_id") @@ -99,6 +109,46 @@ def build_benchmark_generation_prompt(ctx_items: list[tuple[str, str]]) -> str: ) +async def _generate_benchmark_item_once( + *, + kb_instance: Any, + db_id: str, + all_chunks: list[dict[str, Any]], + llm: Any, + context_count: int, +) -> dict[str, Any] | None: + anchor_chunk = all_chunks[random.randrange(len(all_chunks))] + neighbor_chunks = await select_neighbor_chunks_by_kb_query( + kb_instance=kb_instance, + db_id=db_id, + anchor_chunk=anchor_chunk, + neighbors_count=context_count - 1, + ) + ctx_chunks = [anchor_chunk] + neighbor_chunks + ctx_items = [(chunk["id"], chunk["content"]) for chunk in ctx_chunks] + allowed_ids = {cid for cid, _ in ctx_items} + + try: + resp = await llm.call(build_benchmark_generation_prompt(ctx_items), False) + obj = json_repair.loads(resp.content if resp else "") + query = obj.get("query") + answer = obj.get("gold_answer") + gold_ids = obj.get("gold_chunk_ids") + if not query or not answer or not isinstance(gold_ids, list): + logger.warning(f"Generated JSON missing fields or invalid format: {obj}") + return None + + gold_ids = [str(item) for item in gold_ids if str(item) in allowed_ids] + if not gold_ids: + logger.warning("Generated gold_chunk_ids not found in allowed context") + return None + + return {"query": query, "gold_chunk_ids": gold_ids, "gold_answer": answer} + except Exception as e: + logger.warning(f"Benchmark generation failed for one item: {e}") + return None + + async def iter_generated_benchmark_items( *, kb_instance: Any, @@ -106,7 +156,9 @@ async def iter_generated_benchmark_items( count: int, neighbors_count: int, llm_model_spec: str | None, + concurrency_count: int = DEFAULT_BENCHMARK_GENERATION_CONCURRENCY, progress_cb: Callable[[int, str], Any] | None = None, + cancel_cb: Callable[[], Any] | None = None, ) -> AsyncIterator[dict[str, Any]]: if progress_cb: await progress_cb(5, "加载chunks") @@ -122,46 +174,71 @@ async def iter_generated_benchmark_items( raise ValueError("llm_model_spec 不能为空") llm = select_model(model_spec=llm_model_spec) context_count = max(clamp_neighbors_count(neighbors_count), 1) - generated = 0 - attempts = 0 max_attempts = max(count * 5, 50) + worker_count = normalize_generation_concurrency_count(concurrency_count) + actual_worker_count = min(worker_count, max(count, 1), max_attempts) + generated = 0 + results: list[tuple[int, dict[str, Any]]] = [] + state_lock = asyncio.Lock() + queue: asyncio.Queue[int] = asyncio.Queue() - while generated < count and attempts < max_attempts: - attempts += 1 - anchor_chunk = all_chunks[random.randrange(len(all_chunks))] - neighbor_chunks = await select_neighbor_chunks_by_kb_query( - kb_instance=kb_instance, - db_id=db_id, - anchor_chunk=anchor_chunk, - neighbors_count=context_count - 1, - ) - ctx_chunks = [anchor_chunk] + neighbor_chunks - ctx_items = [(chunk["id"], chunk["content"]) for chunk in ctx_chunks] - allowed_ids = {cid for cid, _ in ctx_items} + for attempt_no in range(max_attempts): + queue.put_nowait(attempt_no) - try: - resp = await llm.call(build_benchmark_generation_prompt(ctx_items), False) - obj = json_repair.loads(resp.content if resp else "") - query = obj.get("query") - answer = obj.get("gold_answer") - gold_ids = obj.get("gold_chunk_ids") - if not query or not answer or not isinstance(gold_ids, list): - logger.warning(f"Generated JSON missing fields or invalid format: {obj}") - continue + async def worker() -> None: + nonlocal generated + while True: + if cancel_cb: + await cancel_cb() + async with state_lock: + if generated >= count: + return + try: + attempt_no = queue.get_nowait() + except asyncio.QueueEmpty: + return + try: + item = await _generate_benchmark_item_once( + kb_instance=kb_instance, + db_id=db_id, + all_chunks=all_chunks, + llm=llm, + context_count=context_count, + ) + if item is None: + continue + progress = None + message = None + async with state_lock: + if generated >= count: + continue + generated += 1 + results.append((attempt_no, item)) + if progress_cb: + progress = int(99 * generated / max(count, 1)) + message = f"已生成 {generated}/{count}" + if progress_cb: + await progress_cb(progress, message) + finally: + queue.task_done() - gold_ids = [str(item) for item in gold_ids if str(item) in allowed_ids] - if not gold_ids: - logger.warning("Generated gold_chunk_ids not found in allowed context") - continue + workers = [asyncio.create_task(worker()) for _ in range(actual_worker_count)] + try: + await asyncio.gather(*workers) + except asyncio.CancelledError: + for task in workers: + task.cancel() + await asyncio.gather(*workers, return_exceptions=True) + raise + except Exception: + for task in workers: + task.cancel() + await asyncio.gather(*workers, return_exceptions=True) + raise - generated += 1 - if progress_cb: - await progress_cb(0 + int(99 * generated / max(count, 1)), f"已生成 {generated}/{count}") - yield {"query": query, "gold_chunk_ids": gold_ids, "gold_answer": answer} - except Exception as e: - logger.warning(f"Benchmark generation failed for one item: {e}") - continue + for _, item in sorted(results, key=lambda pair: pair[0]): + yield item def dump_benchmark_item(item: dict[str, Any]) -> str: - return json.dumps(item, ensure_ascii=False) + "\n" + return json.dumps(item, ensure_ascii=False, separators=(",", ":")) + "\n" diff --git a/backend/package/yuxi/repositories/evaluation_repository.py b/backend/package/yuxi/repositories/evaluation_repository.py index 973870e8..f9c74d65 100644 --- a/backend/package/yuxi/repositories/evaluation_repository.py +++ b/backend/package/yuxi/repositories/evaluation_repository.py @@ -2,73 +2,37 @@ from __future__ import annotations from typing import Any -from sqlalchemy import delete, select +from sqlalchemy import delete, func, select from yuxi.storage.postgres.manager import pg_manager -from yuxi.storage.postgres.models_knowledge import EvaluationBenchmark, EvaluationResult, EvaluationResultDetail +from yuxi.storage.postgres.models_knowledge import ( + EvaluationDataset, + EvaluationDatasetItem, + EvaluationRun, + EvaluationRunItem, +) class EvaluationRepository: - async def get_all_benchmarks(self) -> list[EvaluationBenchmark]: - """获取所有评估基准""" + async def create_dataset(self, dataset_data: dict[str, Any]) -> EvaluationDataset: + dataset = EvaluationDataset(**dataset_data) async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(EvaluationBenchmark)) - return list(result.scalars().all()) + session.add(dataset) + return dataset - async def create_benchmark(self, data: dict[str, Any]) -> EvaluationBenchmark: - benchmark = EvaluationBenchmark(**data) + async def create_dataset_with_items( + self, dataset_data: dict[str, Any], items_data: list[dict[str, Any]] + ) -> EvaluationDataset: + dataset = EvaluationDataset(**dataset_data) + items = [EvaluationDatasetItem(**item) for item in items_data] async with pg_manager.get_async_session_context() as session: - session.add(benchmark) - return benchmark + session.add(dataset) + session.add_all(items) + return dataset - async def get_benchmark(self, benchmark_id: str) -> EvaluationBenchmark | None: + async def update_dataset(self, dataset_id: str, data: dict[str, Any]) -> EvaluationDataset | None: async with pg_manager.get_async_session_context() as session: - result = await session.execute( - select(EvaluationBenchmark).where(EvaluationBenchmark.benchmark_id == benchmark_id) - ) - return result.scalar_one_or_none() - - async def list_benchmarks(self, db_id: str) -> list[EvaluationBenchmark]: - async with pg_manager.get_async_session_context() as session: - result = await session.execute( - select(EvaluationBenchmark) - .where(EvaluationBenchmark.db_id == db_id) - .order_by(EvaluationBenchmark.created_at.desc()) - ) - return list(result.scalars().all()) - - async def delete_benchmark(self, benchmark_id: str) -> None: - async with pg_manager.get_async_session_context() as session: - result = await session.execute( - select(EvaluationBenchmark).where(EvaluationBenchmark.benchmark_id == benchmark_id) - ) - record = result.scalar_one_or_none() - if record is not None: - await session.delete(record) - - async def create_result(self, data: dict[str, Any]) -> EvaluationResult: - result_row = EvaluationResult(**data) - async with pg_manager.get_async_session_context() as session: - session.add(result_row) - return result_row - - async def get_result(self, task_id: str) -> EvaluationResult | None: - async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(EvaluationResult).where(EvaluationResult.task_id == task_id)) - return result.scalar_one_or_none() - - async def list_results(self, db_id: str) -> list[EvaluationResult]: - async with pg_manager.get_async_session_context() as session: - result = await session.execute( - select(EvaluationResult) - .where(EvaluationResult.db_id == db_id) - .order_by(EvaluationResult.started_at.desc()) - ) - return list(result.scalars().all()) - - async def update_result(self, task_id: str, data: dict[str, Any]) -> EvaluationResult | None: - async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(EvaluationResult).where(EvaluationResult.task_id == task_id)) + result = await session.execute(select(EvaluationDataset).where(EvaluationDataset.dataset_id == dataset_id)) record = result.scalar_one_or_none() if record is None: return None @@ -76,43 +40,134 @@ class EvaluationRepository: setattr(record, key, value) return record - async def delete_result(self, task_id: str) -> None: + async def add_dataset_items(self, items_data: list[dict[str, Any]]) -> None: + items = [EvaluationDatasetItem(**item) for item in items_data] async with pg_manager.get_async_session_context() as session: - await session.execute(delete(EvaluationResultDetail).where(EvaluationResultDetail.task_id == task_id)) - result = await session.execute(select(EvaluationResult).where(EvaluationResult.task_id == task_id)) + session.add_all(items) + + async def get_dataset(self, dataset_id: str) -> EvaluationDataset | None: + async with pg_manager.get_async_session_context() as session: + result = await session.execute(select(EvaluationDataset).where(EvaluationDataset.dataset_id == dataset_id)) + return result.scalar_one_or_none() + + async def list_datasets(self, db_id: str) -> list[EvaluationDataset]: + async with pg_manager.get_async_session_context() as session: + result = await session.execute( + select(EvaluationDataset) + .where(EvaluationDataset.db_id == db_id) + .order_by(EvaluationDataset.created_at.desc()) + ) + return list(result.scalars().all()) + + async def list_dataset_items( + self, dataset_id: str, offset: int = 0, limit: int = 100 + ) -> list[EvaluationDatasetItem]: + async with pg_manager.get_async_session_context() as session: + result = await session.execute( + select(EvaluationDatasetItem) + .where(EvaluationDatasetItem.dataset_id == dataset_id) + .order_by(EvaluationDatasetItem.item_index.asc()) + .offset(offset) + .limit(limit) + ) + return list(result.scalars().all()) + + async def count_dataset_items(self, dataset_id: str) -> int: + async with pg_manager.get_async_session_context() as session: + result = await session.execute( + select(func.count(EvaluationDatasetItem.id)).where(EvaluationDatasetItem.dataset_id == dataset_id) + ) + return int(result.scalar() or 0) + + async def list_all_dataset_items(self, dataset_id: str) -> list[EvaluationDatasetItem]: + async with pg_manager.get_async_session_context() as session: + result = await session.execute( + select(EvaluationDatasetItem) + .where(EvaluationDatasetItem.dataset_id == dataset_id) + .order_by(EvaluationDatasetItem.item_index.asc()) + ) + return list(result.scalars().all()) + + async def delete_dataset(self, dataset_id: str) -> None: + async with pg_manager.get_async_session_context() as session: + result = await session.execute(select(EvaluationDataset).where(EvaluationDataset.dataset_id == dataset_id)) record = result.scalar_one_or_none() if record is not None: await session.delete(record) - async def upsert_result_detail( - self, task_id: str, query_index: int, data: dict[str, Any] - ) -> EvaluationResultDetail: + async def create_run(self, data: dict[str, Any]) -> EvaluationRun: + run = EvaluationRun(**data) + async with pg_manager.get_async_session_context() as session: + session.add(run) + return run + + async def get_run(self, run_id: str) -> EvaluationRun | None: + async with pg_manager.get_async_session_context() as session: + result = await session.execute(select(EvaluationRun).where(EvaluationRun.run_id == run_id)) + return result.scalar_one_or_none() + + async def list_runs(self, db_id: str) -> list[EvaluationRun]: async with pg_manager.get_async_session_context() as session: result = await session.execute( - select(EvaluationResultDetail).where( - (EvaluationResultDetail.task_id == task_id) & (EvaluationResultDetail.query_index == query_index) + select(EvaluationRun).where(EvaluationRun.db_id == db_id).order_by(EvaluationRun.started_at.desc()) + ) + return list(result.scalars().all()) + + async def update_run(self, run_id: str, data: dict[str, Any]) -> EvaluationRun | None: + async with pg_manager.get_async_session_context() as session: + result = await session.execute(select(EvaluationRun).where(EvaluationRun.run_id == run_id)) + record = result.scalar_one_or_none() + if record is None: + return None + for key, value in data.items(): + setattr(record, key, value) + return record + + async def delete_run(self, run_id: str) -> None: + async with pg_manager.get_async_session_context() as session: + await session.execute(delete(EvaluationRunItem).where(EvaluationRunItem.run_id == run_id)) + result = await session.execute(select(EvaluationRun).where(EvaluationRun.run_id == run_id)) + record = result.scalar_one_or_none() + if record is not None: + await session.delete(record) + + async def upsert_run_item(self, run_id: str, item_index: int, data: dict[str, Any]) -> EvaluationRunItem: + async with pg_manager.get_async_session_context() as session: + result = await session.execute( + select(EvaluationRunItem).where( + (EvaluationRunItem.run_id == run_id) & (EvaluationRunItem.item_index == item_index) ) ) record = result.scalar_one_or_none() if record is None: - record = EvaluationResultDetail(task_id=task_id, query_index=query_index, **data) + record = EvaluationRunItem(run_id=run_id, item_index=item_index, **data) session.add(record) return record for key, value in data.items(): setattr(record, key, value) return record - async def list_result_details(self, task_id: str) -> list[EvaluationResultDetail]: + async def list_run_items(self, run_id: str, offset: int = 0, limit: int = 100) -> list[EvaluationRunItem]: async with pg_manager.get_async_session_context() as session: result = await session.execute( - select(EvaluationResultDetail) - .where(EvaluationResultDetail.task_id == task_id) - .order_by(EvaluationResultDetail.query_index.asc()) + select(EvaluationRunItem) + .where(EvaluationRunItem.run_id == run_id) + .order_by(EvaluationRunItem.item_index.asc()) + .offset(offset) + .limit(limit) ) return list(result.scalars().all()) + async def count_run_items(self, run_id: str) -> int: + async with pg_manager.get_async_session_context() as session: + result = await session.execute( + select(func.count(EvaluationRunItem.id)).where(EvaluationRunItem.run_id == run_id) + ) + return int(result.scalar() or 0) + async def delete_all(self) -> None: async with pg_manager.get_async_session_context() as session: - await session.execute(delete(EvaluationResultDetail)) - await session.execute(delete(EvaluationResult)) - await session.execute(delete(EvaluationBenchmark)) + await session.execute(delete(EvaluationRunItem)) + await session.execute(delete(EvaluationRun)) + await session.execute(delete(EvaluationDatasetItem)) + await session.execute(delete(EvaluationDataset)) diff --git a/backend/package/yuxi/services/evaluation_service.py b/backend/package/yuxi/services/evaluation_service.py index 0dd0c2c1..7a4b9f00 100644 --- a/backend/package/yuxi/services/evaluation_service.py +++ b/backend/package/yuxi/services/evaluation_service.py @@ -1,18 +1,22 @@ import json -import os import re import uuid -from datetime import datetime from typing import Any from yuxi.knowledge import knowledge_base -from yuxi.knowledge.eval.benchmark_generation import dump_benchmark_item, iter_generated_benchmark_items +from yuxi.knowledge.eval.benchmark_generation import ( + dump_benchmark_item, + iter_generated_benchmark_items, + normalize_generation_concurrency_count, +) from yuxi.knowledge.eval.evaluator import aggregate_metrics, evaluate_question from yuxi.models import select_model from yuxi.repositories.evaluation_repository import EvaluationRepository from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository +from yuxi.repositories.task_repository import TaskRepository from yuxi.services.task_service import TaskContext, tasker from yuxi.utils import logger +from yuxi.utils.datetime_utils import format_utc_datetime, utc_now_naive class EvaluationService: @@ -21,343 +25,389 @@ class EvaluationService: def __init__(self): self.eval_repo = EvaluationRepository() self.kb_repo = KnowledgeBaseRepository() + self.task_repo = TaskRepository() - async def _get_benchmark_dir(self, db_id: str) -> str: - """获取评估基准目录""" - kb_instance = await knowledge_base.aget_kb(db_id) - base_dir = os.path.join(kb_instance.work_dir, db_id) - path = os.path.join(base_dir, "benchmarks") - os.makedirs(path, exist_ok=True) - return path + def _dataset_to_dict(self, row) -> dict[str, Any]: + return { + "id": row.dataset_id, + "dataset_id": row.dataset_id, + "name": row.name, + "description": row.description, + "db_id": row.db_id, + "item_count": row.item_count, + "has_gold_chunks": row.has_gold_chunks, + "has_gold_answers": row.has_gold_answers, + "build_metadata": row.build_metadata or {}, + "created_by": row.created_by, + "created_at": format_utc_datetime(row.created_at), + "updated_at": format_utc_datetime(row.updated_at), + } - async def _get_result_dir(self, db_id: str) -> str: - """获取评估结果目录""" - kb_instance = await knowledge_base.aget_kb(db_id) - base_dir = os.path.join(kb_instance.work_dir, db_id) - path = os.path.join(base_dir, "results") - os.makedirs(path, exist_ok=True) - return path + def _dataset_item_to_dict(self, item) -> dict[str, Any]: + return { + "item_id": item.item_id, + "item_index": item.item_index, + "query": item.query_text, + "gold_chunk_ids": item.gold_chunk_ids or [], + "gold_answer": item.gold_answer, + } - # 已移除基准回退逻辑,统一使用集中元数据 + def _run_item_to_dict(self, item) -> dict[str, Any]: + return { + "query": item.query_text, + "gold_chunk_ids": item.gold_chunk_ids, + "gold_answer": item.gold_answer, + "generated_answer": item.generated_answer, + "retrieved_chunks": item.retrieved_chunks, + "metrics": item.metrics or {}, + } - # 已移除结果回退逻辑,统一通过 db_id 定位 + def _is_error_run_item(self, item) -> bool: + metrics = item.metrics or {} + return metrics.get("score", 1.0) <= 0.5 or any( + metrics.get(key, 1.0) < 0.3 for key in metrics if key.startswith("recall@") + ) - 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 = await 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 = { - "id": benchmark_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, - "benchmark_file": data_file_path, - "created_by": created_by, - "created_at": datetime.utcnow().isoformat(), - "updated_at": datetime.utcnow().isoformat(), - } - await self.eval_repo.create_benchmark( - { - "benchmark_id": benchmark_id, - "db_id": db_id, - "name": name, - "description": description, - "question_count": len(questions), - "has_gold_chunks": has_gold_chunks, - "has_gold_answers": has_gold_answers, - "data_file_path": data_file_path, - "created_by": created_by, - } - ) - return meta - - except Exception as e: - logger.error(f"上传评估基准失败: {e}") - raise - - async def get_benchmarks(self, db_id: str) -> list[dict[str, Any]]: - """获取知识库的评估基准列表""" - try: - rows = await self.eval_repo.list_benchmarks(db_id) - return [ - { - "id": row.benchmark_id, - "benchmark_id": row.benchmark_id, - "name": row.name, - "description": row.description, - "db_id": row.db_id, - "question_count": row.question_count, - "has_gold_chunks": row.has_gold_chunks, - "has_gold_answers": row.has_gold_answers, - "benchmark_file": row.data_file_path, - "created_by": row.created_by, - "created_at": row.created_at.isoformat() if row.created_at else None, - "updated_at": row.updated_at.isoformat() if row.updated_at else None, - } - for row in rows - ] - - except Exception as e: - logger.error(f"获取评估基准列表失败: {e}") - raise - - async def get_benchmark_detail(self, benchmark_id: str) -> dict[str, Any]: - """获取评估基准详情 (包含问题列表)""" - try: - row = await self.eval_repo.get_benchmark(benchmark_id) - if row is None: - raise ValueError("Benchmark not found") - questions = [] - if row.data_file_path and os.path.exists(row.data_file_path): - with open(row.data_file_path, encoding="utf-8") as f: - for line in f: - if line.strip(): - questions.append(json.loads(line)) - return { - "id": row.benchmark_id, - "benchmark_id": row.benchmark_id, - "name": row.name, - "description": row.description, - "db_id": row.db_id, - "question_count": row.question_count, - "has_gold_chunks": row.has_gold_chunks, - "has_gold_answers": row.has_gold_answers, - "benchmark_file": row.data_file_path, - "created_by": row.created_by, - "created_at": row.created_at.isoformat() if row.created_at else None, - "updated_at": row.updated_at.isoformat() if row.updated_at else None, - "questions": questions, - } - - except Exception as e: - logger.error(f"获取评估基准详情失败: {e}") - raise - - async def get_benchmark_detail_by_db( - self, db_id: str, benchmark_id: str, page: int = 1, page_size: int = 10 - ) -> dict[str, Any]: - """根据 db_id 获取评估基准详情(支持分页)""" - try: - row = await self.eval_repo.get_benchmark(benchmark_id) - if row is None or row.db_id != db_id: - raise ValueError("Benchmark not found") - data_file_path = row.data_file_path - total_questions = row.question_count or 0 - questions = [] - - if data_file_path and os.path.exists(data_file_path): - # 计算分页范围 - start_index = (page - 1) * page_size - end_index = start_index + page_size - - # 读取指定范围的问题 - with open(data_file_path, encoding="utf-8") as f: - current_index = 0 - for line in f: - if not line.strip(): - continue - - # 只处理指定范围内的问题 - if current_index >= start_index and current_index < end_index: - questions.append(json.loads(line)) - elif current_index >= end_index: - break # 已经读取到足够的问题,停止读取 - - current_index += 1 - - # 计算分页信息 - total_pages = (total_questions + page_size - 1) // page_size - - return { - "id": row.benchmark_id, - "benchmark_id": row.benchmark_id, - "name": row.name, - "description": row.description, - "db_id": row.db_id, - "question_count": row.question_count, - "has_gold_chunks": row.has_gold_chunks, - "has_gold_answers": row.has_gold_answers, - "benchmark_file": data_file_path, - "created_by": row.created_by, - "created_at": row.created_at.isoformat() if row.created_at else None, - "updated_at": row.updated_at.isoformat() if row.updated_at else None, - "questions": questions, - "pagination": { - "current_page": page, - "page_size": page_size, - "total_questions": total_questions, - "total_pages": total_pages, - "has_next": page < total_pages, - "has_prev": page > 1, - }, - } - except Exception as e: - logger.error(f"获取评估基准详情失败: {e}") - raise - - async def get_benchmark_download_info(self, benchmark_id: str) -> dict[str, str]: - """获取评估基准下载信息""" - row = await self.eval_repo.get_benchmark(benchmark_id) - if row is None: - raise ValueError("Benchmark not found") - - data_file_path = row.data_file_path or "" - if not data_file_path or not os.path.exists(data_file_path): - raise ValueError("Benchmark file not found") - - filename_base = (row.name or "").strip() - if not filename_base: - filename_base = row.benchmark_id - - filename_base = re.sub(r"[\\/:*?\"<>|]+", "_", filename_base).strip() - if not filename_base or filename_base in {".", ".."}: - filename_base = row.benchmark_id - - if not filename_base.endswith(".jsonl"): - filename_base = f"{filename_base}.jsonl" - - return {"file_path": data_file_path, "filename": filename_base} - - async def delete_benchmark(self, benchmark_id: str) -> None: - """删除评估基准""" - try: - row = await self.eval_repo.get_benchmark(benchmark_id) - if row is None: - raise ValueError("Benchmark not found") - if row.data_file_path and os.path.exists(row.data_file_path): - os.remove(row.data_file_path) - await self.eval_repo.delete_benchmark(benchmark_id) - logger.info(f"成功删除评估基准: {benchmark_id}") + async def _sync_dataset_build_metadata(self, row) -> None: + metadata = dict(row.build_metadata or {}) + if metadata.get("source") != "generated" or metadata.get("status") not in {"pending", "running"}: return + task_id = metadata.get("task_id") + task = await self.task_repo.get_by_id(task_id) if task_id else None + if task is None: + metadata.pop("progress", None) + metadata.update(status="failed", message="生成任务不存在") + elif task.status == "success": + metadata.update(status="completed", progress=100, message=task.message or "完成") + elif task.status in {"failed", "cancelled"}: + metadata.pop("progress", None) + metadata.update(status="failed", message=task.error or task.message or "生成任务失败") + else: + metadata.update(status=task.status, progress=task.progress, message=task.message) + + if metadata != (row.build_metadata or {}): + await self.eval_repo.update_dataset(row.dataset_id, {"build_metadata": metadata}) + row.build_metadata = metadata + + def _build_dataset_items( + self, dataset_id: str, db_id: str, questions: list[dict[str, Any]] + ) -> list[dict[str, Any]]: + return [ + { + "item_id": f"dataset_item_{uuid.uuid4().hex[:12]}", + "dataset_id": dataset_id, + "db_id": db_id, + "item_index": index, + "query_text": item["query"], + "gold_chunk_ids": item.get("gold_chunk_ids") or [], + "gold_answer": item.get("gold_answer"), + } + for index, item in enumerate(questions) + ] + + def _build_jsonl_content(self, items: list[Any]) -> str: + lines = [] + for item in items: + payload = {"query": item.query_text} + if item.gold_chunk_ids: + payload["gold_chunk_ids"] = item.gold_chunk_ids + if item.gold_answer: + payload["gold_answer"] = item.gold_answer + lines.append(dump_benchmark_item(payload).rstrip("\n")) + return "\n".join(lines) + ("\n" if lines else "") + + def _safe_jsonl_filename(self, name: str | None, fallback: str) -> str: + filename = (name or "").strip() or fallback + filename = re.sub(r"[\\/:*?\"<>|]+", "_", filename).strip() + if not filename or filename in {".", ".."}: + filename = fallback + return filename if filename.endswith(".jsonl") else f"{filename}.jsonl" + + def _parse_jsonl_questions(self, file_content: bytes) -> tuple[list[dict[str, Any]], bool, bool]: + questions = [] + has_gold_chunks = False + has_gold_answers = False + content = file_content.decode("utf-8") + + for line_num, line in enumerate(content.strip().split("\n"), 1): + if not line.strip(): + continue + try: + item = json.loads(line) + except json.JSONDecodeError as e: + raise ValueError(f"第{line_num}行JSON格式错误: {str(e)}") + 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) + + if not questions: + raise ValueError("文件中没有有效的问题数据") + return questions, has_gold_chunks, has_gold_answers + + async def upload_dataset( + self, db_id: str, file_content: bytes, filename: str, name: str, description: str, created_by: str + ) -> dict[str, Any]: + try: + questions, has_gold_chunks, has_gold_answers = self._parse_jsonl_questions(file_content) + dataset_id = f"dataset_{uuid.uuid4().hex[:8]}" + dataset_name = name.strip() or filename or dataset_id + + row = await self.eval_repo.create_dataset_with_items( + { + "dataset_id": dataset_id, + "db_id": db_id, + "name": dataset_name, + "description": description, + "item_count": len(questions), + "has_gold_chunks": has_gold_chunks, + "has_gold_answers": has_gold_answers, + "build_metadata": { + "source": "upload", + "status": "completed", + "progress": 100, + "filename": filename, + }, + "created_by": created_by, + }, + self._build_dataset_items(dataset_id, db_id, questions), + ) + return self._dataset_to_dict(row) except Exception as e: - logger.error(f"删除评估基准失败: {e}") + logger.error(f"上传评估数据集失败: {e}") raise - async def delete_evaluation_result(self, task_id: str, db_id: str) -> None: - """删除评估结果""" - if not task_id: - raise ValueError("task_id is required") - await self.delete_evaluation_result_by_db(db_id, task_id) + async def list_datasets(self, db_id: str) -> list[dict[str, Any]]: + try: + rows = await self.eval_repo.list_datasets(db_id) + for row in rows: + await self._sync_dataset_build_metadata(row) + return [self._dataset_to_dict(row) for row in rows] + 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]: - task_id = f"gen_benchmark_{uuid.uuid4().hex[:8]}" - await tasker.enqueue( - name="生成评估基准", - task_type="benchmark_generation", - payload={"task_id": task_id, "db_id": db_id, "created_by": created_by, **params}, - coroutine=self._generate_benchmark_task, + async def get_dataset_detail( + self, db_id: str, dataset_id: str, page: int = 1, page_size: int = 10 + ) -> dict[str, Any]: + try: + row = await self.eval_repo.get_dataset(dataset_id) + if row is None or row.db_id != db_id: + raise ValueError("Dataset not found") + if (row.build_metadata or {}).get("status", "completed") != "completed": + raise ValueError("Dataset is not ready") + + total_items = await self.eval_repo.count_dataset_items(dataset_id) + items = await self.eval_repo.list_dataset_items(dataset_id, (page - 1) * page_size, page_size) + total_pages = (total_items + page_size - 1) // page_size + data = self._dataset_to_dict(row) + data.update( + { + "items": [self._dataset_item_to_dict(item) for item in items], + "pagination": { + "current_page": page, + "page_size": page_size, + "total_items": total_items, + "total_pages": total_pages, + "has_next": page < total_pages, + "has_prev": page > 1, + }, + } + ) + return data + except Exception as e: + logger.error(f"获取评估数据集详情失败: {e}") + raise + + async def export_dataset_jsonl(self, dataset_id: str) -> dict[str, str]: + row = await self.eval_repo.get_dataset(dataset_id) + if row is None: + raise ValueError("Dataset not found") + if (row.build_metadata or {}).get("status", "completed") != "completed": + raise ValueError("Dataset is not ready") + items = await self.eval_repo.list_all_dataset_items(dataset_id) + return { + "filename": self._safe_jsonl_filename(row.name, row.dataset_id), + "content": self._build_jsonl_content(items), + } + + async def delete_dataset(self, dataset_id: str) -> None: + try: + row = await self.eval_repo.get_dataset(dataset_id) + if row is None: + raise ValueError("Dataset not found") + await self.eval_repo.delete_dataset(dataset_id) + logger.info(f"成功删除评估数据集: {dataset_id}") + except Exception as e: + logger.error(f"删除评估数据集失败: {e}") + raise + + async def generate_dataset( + self, + db_id: str, + name: str, + description: str, + count: int, + neighbors_count: int, + concurrency_count: int, + llm_model_spec: str, + created_by: str, + ) -> dict[str, Any]: + dataset_id = f"dataset_{uuid.uuid4().hex[:8]}" + count = int(count) + neighbors_count = int(neighbors_count) + concurrency_count = normalize_generation_concurrency_count(concurrency_count) + build_metadata = { + "source": "generated", + "status": "pending", + "progress": 0, + "params": { + "count": count, + "neighbors_count": neighbors_count, + "concurrency_count": concurrency_count, + "llm_model_spec": llm_model_spec, + }, + } + await self.eval_repo.create_dataset( + { + "dataset_id": dataset_id, + "db_id": db_id, + "name": name, + "description": description, + "item_count": 0, + "has_gold_chunks": True, + "has_gold_answers": True, + "build_metadata": build_metadata, + "created_by": created_by, + } ) - return {"task_id": task_id, "message": "基准生成任务已提交"} + task = await tasker.enqueue( + name="生成评估数据集", + task_type="dataset_generation", + payload={ + "dataset_id": dataset_id, + "db_id": db_id, + "created_by": created_by, + "name": name, + "description": description, + "count": count, + "neighbors_count": neighbors_count, + "concurrency_count": concurrency_count, + "llm_model_spec": llm_model_spec, + }, + coroutine=self._generate_dataset_task, + ) + build_metadata["task_id"] = task.id + await self.eval_repo.update_dataset(dataset_id, {"build_metadata": build_metadata}) + return {"dataset_id": dataset_id, "task_id": task.id, "message": "评估数据集生成任务已提交"} - async def _generate_benchmark_task(self, context: TaskContext): + async def _update_dataset_build_metadata( + self, dataset_id: str, metadata: dict[str, Any], **updates + ) -> dict[str, Any]: + metadata.update(updates) + await self.eval_repo.update_dataset(dataset_id, {"build_metadata": metadata}) + return metadata + + async def _generate_dataset_task(self, context: TaskContext): await context.set_progress(0, "初始化") - task = context._tasker._tasks.get(context.task_id) payload = task.payload if task else {} + dataset_id = payload.get("dataset_id") db_id = payload.get("db_id") - name = payload.get("name", "自动生成评估基准") - description = payload.get("description", "") count = int(payload.get("count", 10)) neighbors_count = int(payload.get("neighbors_count", 1)) + concurrency_count = normalize_generation_concurrency_count(payload.get("concurrency_count")) llm_model_spec = payload.get("llm_model_spec") + build_metadata = { + "source": "generated", + "status": "running", + "progress": 0, + "task_id": context.task_id, + "params": { + "count": count, + "neighbors_count": neighbors_count, + "concurrency_count": concurrency_count, + "llm_model_spec": llm_model_spec, + }, + } + await self._update_dataset_build_metadata(dataset_id, build_metadata) - kb_instance = await knowledge_base.aget_kb(db_id) - if not kb_instance: - await context.set_message("知识库不存在") - raise ValueError("Knowledge Base not found") - if kb_instance.kb_type != "milvus": - await context.set_message("仅支持 commonrag/Milvus 类型知识库生成评估基准") - raise ValueError("Unsupported KB type for benchmark generation") - - benchmark_id = f"benchmark_{uuid.uuid4().hex[:8]}" - bench_dir = await self._get_benchmark_dir(db_id) - data_file_path = os.path.join(bench_dir, f"{benchmark_id}.jsonl") - generated = 0 + async def report_progress(progress: float, message: str | None = None) -> None: + await context.set_progress(progress, message) + await self._update_dataset_build_metadata( + dataset_id, + build_metadata, + progress=max(0, min(round(progress), 100)), + message=message or build_metadata.get("message", ""), + ) try: - with open(data_file_path, "w", encoding="utf-8") as f: + kb_instance = await knowledge_base.aget_kb(db_id) + if not kb_instance: + await report_progress(100, "知识库不存在") + raise ValueError("Knowledge Base not found") + if kb_instance.kb_type != "milvus": + await report_progress(100, "仅支持 commonrag/Milvus 类型知识库生成评估数据集") + raise ValueError("Unsupported KB type for dataset generation") + + questions = [] + try: async for item in iter_generated_benchmark_items( kb_instance=kb_instance, db_id=db_id, count=count, neighbors_count=neighbors_count, llm_model_spec=llm_model_spec, - progress_cb=context.set_progress, + concurrency_count=concurrency_count, + progress_cb=report_progress, + cancel_cb=context.raise_if_cancelled, ): - f.write(dump_benchmark_item(item)) - generated += 1 - except ValueError as e: - if str(e) == "No chunks found in knowledge base": - await context.set_message("知识库为空或未解析到chunks") + questions.append(item) + except ValueError as e: + if str(e) == "No chunks found in knowledge base": + await report_progress(100, "知识库为空或未解析到chunks") + raise + + if not questions: + raise ValueError("未生成有效评估题目") + + await self.eval_repo.add_dataset_items(self._build_dataset_items(dataset_id, db_id, questions)) + await self.eval_repo.update_dataset(dataset_id, {"item_count": len(questions)}) + await self._update_dataset_build_metadata( + dataset_id, + build_metadata, + status="completed", + progress=100, + message="完成", + ) + await context.set_progress(100, "完成") + except Exception as e: + await self._update_dataset_build_metadata( + dataset_id, + build_metadata, + status="failed", + progress=100, + error_message=str(e), + message=str(e), + ) raise - await self.eval_repo.create_benchmark( - { - "benchmark_id": benchmark_id, - "db_id": db_id, - "name": name, - "description": description, - "question_count": generated, - "has_gold_chunks": True, - "has_gold_answers": True, - "data_file_path": data_file_path, - "created_by": payload.get("created_by"), - } - ) - - await context.set_progress(100, "完成") - async def run_evaluation( - self, db_id: str, benchmark_id: str, model_config: dict[str, Any] = None, created_by: str = "system" + self, db_id: str, dataset_id: str, model_config: dict[str, Any] = None, created_by: str = "system" ) -> str: - """运行RAG评估""" try: - task_id = f"eval_{uuid.uuid4().hex[:8]}" + run_id = f"run_{uuid.uuid4().hex[:8]}" + dataset_row = await self.eval_repo.get_dataset(dataset_id) + if dataset_row is None or dataset_row.db_id != db_id: + raise ValueError("Dataset not found") + if (dataset_row.build_metadata or {}).get("status", "completed") != "completed": + raise ValueError("Dataset is not ready") - benchmark_row = await self.eval_repo.get_benchmark(benchmark_id) - if benchmark_row is None or benchmark_row.db_id != db_id: - raise ValueError("Benchmark not found") - - # 从知识库元数据中获取检索配置 retrieval_config = {} try: kb_row = await self.kb_repo.get_by_id(db_id) @@ -370,85 +420,70 @@ class EvaluationService: logger.info(f"从知识库 {db_id} 加载检索配置: {list(retrieval_config.keys())}") except Exception as e: logger.error(f"获取知识库检索配置失败: {e}") - # 使用空配置作为默认值 - # 合并前端传递的模型配置 if model_config: retrieval_config.update(model_config) - await self.eval_repo.create_result( + await self.eval_repo.create_run( { - "task_id": task_id, + "run_id": run_id, "db_id": db_id, - "benchmark_id": benchmark_id, + "dataset_id": dataset_id, "status": "running", "retrieval_config": retrieval_config, "metrics": {}, "overall_score": None, - "total_questions": benchmark_row.question_count or 0, - "completed_questions": 0, - "started_at": datetime.utcnow(), + "total_items": dataset_row.item_count or 0, + "completed_items": 0, + "started_at": utc_now_naive(), "completed_at": None, "created_by": created_by, } ) await tasker.enqueue( - name=f"RAG评估({benchmark_row.name})", + name=f"RAG评估({dataset_row.name})", task_type="rag_evaluation", payload={ - "task_id": task_id, + "run_id": run_id, "db_id": db_id, - "benchmark_id": benchmark_id, + "dataset_id": dataset_id, "retrieval_config": retrieval_config, "created_by": created_by, }, coroutine=self._run_evaluation_task, ) - - return task_id - + return run_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"] + run_id = payload["run_id"] db_id = payload["db_id"] - benchmark_id = payload["benchmark_id"] + dataset_id = payload["dataset_id"] retrieval_config = payload["retrieval_config"] - # 加载基准数据 - await context.set_progress(5, "加载基准数据") - benchmark_row = await self.eval_repo.get_benchmark(benchmark_id) - if benchmark_row is None or benchmark_row.db_id != db_id: - raise ValueError("Benchmark not found") - data_path = benchmark_row.data_file_path - if not data_path or not os.path.exists(data_path): - raise ValueError("Benchmark file not found") + await context.set_progress(5, "加载评估数据集") + dataset_row = await self.eval_repo.get_dataset(dataset_id) + if dataset_row is None or dataset_row.db_id != db_id: + raise ValueError("Dataset not found") + dataset_items = await self.eval_repo.list_all_dataset_items(dataset_id) + if not dataset_items: + raise ValueError("Dataset has no items") - 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 = await knowledge_base.aget_kb(db_id) if not kb_instance: raise ValueError(f"Knowledge Base {db_id} not found") - # 初始化 Judge LLM judge_llm = None - if benchmark_row.has_gold_answers: - # 优先使用配置中的 judge_llm,否则回退到 answer_llm,或者默认 + if dataset_row.has_gold_answers: judge_model_spec = retrieval_config.get("judge_llm") or retrieval_config.get("answer_llm") if judge_model_spec: try: @@ -457,176 +492,153 @@ class EvaluationService: except Exception as e: logger.error(f"Failed to load judge LLM: {e}") - total_questions = len(benchmark_data) all_retrieval_metrics = [] all_answer_metrics = [] + total_items = len(dataset_items) - async def update_result_db( - status: str | None = None, completed: int | None = None, metrics=None, final_score=None - ): - payload = {} + async def update_run_db(status=None, completed=None, metrics=None, final_score=None): + data = {} if status is not None: - payload["status"] = status + data["status"] = status if status in ["completed", "failed"]: - payload["completed_at"] = datetime.utcnow() + data["completed_at"] = utc_now_naive() if completed is not None: - payload["completed_questions"] = completed + data["completed_items"] = completed if metrics is not None: - payload["metrics"] = metrics + data["metrics"] = metrics if final_score is not None: - payload["overall_score"] = final_score - if payload: - await self.eval_repo.update_result(task_id, payload) + data["overall_score"] = final_score + if data: + await self.eval_repo.update_run(run_id, data) - for i, question_data in enumerate(benchmark_data): + for index, item in enumerate(dataset_items): await context.raise_if_cancelled() - progress = 10 + (i / total_questions) * 80 - await context.set_progress(progress, f"评估 {i + 1}/{total_questions}") + progress = 10 + (index / total_items) * 80 + await context.set_progress(progress, f"评估 {index + 1}/{total_items}") + question_data = { + "query": item.query_text, + "gold_chunk_ids": item.gold_chunk_ids or [], + "gold_answer": item.gold_answer, + } question_result = await evaluate_question( kb_instance=kb_instance, db_id=db_id, question_data=question_data, retrieval_config=retrieval_config, - has_gold_chunks=benchmark_row.has_gold_chunks, - has_gold_answers=benchmark_row.has_gold_answers, + has_gold_chunks=dataset_row.has_gold_chunks, + has_gold_answers=dataset_row.has_gold_answers, judge_llm=judge_llm, select_model_fn=select_model, ) - if benchmark_row.has_gold_chunks and question_data.get("gold_chunk_ids"): + if dataset_row.has_gold_chunks and question_data.get("gold_chunk_ids"): all_retrieval_metrics.append(question_result["retrieval_scores"]) - if benchmark_row.has_gold_answers and question_data.get("gold_answer") and judge_llm: + if dataset_row.has_gold_answers and question_data.get("gold_answer") and judge_llm: all_answer_metrics.append(question_result["answer_scores"]) - await self.eval_repo.upsert_result_detail( - task_id=task_id, - query_index=i, - data=question_result["detail"], + await self.eval_repo.upsert_run_item( + run_id=run_id, + item_index=index, + data={"dataset_item_id": item.item_id, **question_result["detail"]}, ) - current_overall_metrics, _ = aggregate_metrics(all_retrieval_metrics, all_answer_metrics) - await context.set_result( - { - "current_metrics": current_overall_metrics, - "completed_questions": i + 1, - "total_questions": total_questions, - } - ) - - if (i + 1) % 5 == 0 or (i + 1) == total_questions: - await update_result_db(completed=i + 1) + if (index + 1) % 5 == 0 or (index + 1) == total_items: + current_metrics, _ = aggregate_metrics(all_retrieval_metrics, all_answer_metrics) + await context.set_result( + {"current_metrics": current_metrics, "completed_items": index + 1, "total_items": total_items} + ) + await update_run_db(completed=index + 1) await context.set_progress(95, "计算最终指标") overall_metrics, overall_score = aggregate_metrics( all_retrieval_metrics, all_answer_metrics, include_overall_score=True ) - - await update_result_db( + await update_run_db( status="completed", - completed=total_questions, + completed=total_items, metrics=overall_metrics, final_score=overall_score, ) await context.set_progress(100, "完成") - except Exception as e: logger.error(f"Task failed: {e}") try: if "payload" in locals(): - await self.eval_repo.update_result( - payload["task_id"], - {"status": "failed", "metrics": {"error": str(e)}, "completed_at": datetime.utcnow()}, + await self.eval_repo.update_run( + payload["run_id"], + {"status": "failed", "metrics": {"error": str(e)}, "completed_at": utc_now_naive()}, ) except Exception as exc: - logger.error(f"Error updating result record: {exc}") - + logger.error(f"Error updating run record: {exc}") await context.set_message(f"Error: {str(e)}") raise - async def get_evaluation_results(self, task_id: str, db_id: str) -> dict[str, Any]: - """获取评估结果""" - if not task_id: - raise ValueError("task_id is required") - return await self.get_evaluation_results_by_db(db_id, task_id) - - async def get_evaluation_history(self, db_id: str) -> list[dict[str, Any]]: - """获取知识库的评估历史记录""" + async def list_runs(self, db_id: str) -> list[dict[str, Any]]: try: - rows = await self.eval_repo.list_results(db_id) + rows = await self.eval_repo.list_runs(db_id) return [ { - "task_id": row.task_id, - "benchmark_id": row.benchmark_id, + "run_id": row.run_id, + "dataset_id": row.dataset_id, "status": row.status, - "started_at": row.started_at.isoformat() if row.started_at else None, - "completed_at": row.completed_at.isoformat() if row.completed_at else None, - "total_questions": row.total_questions, - "completed_questions": row.completed_questions, + "started_at": format_utc_datetime(row.started_at), + "completed_at": format_utc_datetime(row.completed_at), + "total_items": row.total_items, + "completed_items": row.completed_items, "overall_score": row.overall_score, "retrieval_config": row.retrieval_config or {}, "metrics": row.metrics or {}, } for row in rows ] - except Exception as e: - logger.error(f"获取评估历史失败: {e}") + logger.error(f"获取评估运行历史失败: {e}") raise - # 索引与回退逻辑已移除,统一通过 db_id 定位 - async def get_evaluation_results_by_db( - self, db_id: str, task_id: str, page: int = 1, page_size: int = 20, error_only: bool = False + async def get_run_results( + self, db_id: str, run_id: str, page: int = 1, page_size: int = 20, error_only: bool = False ) -> dict[str, Any]: - if not re.match(r"^eval_[a-f0-9]{8}$", task_id): - raise ValueError("Invalid task_id format") - row = await self.eval_repo.get_result(task_id) + if not re.match(r"^run_[a-f0-9]{8}$", run_id): + raise ValueError("Invalid run_id format") + row = await self.eval_repo.get_run(run_id) if row is None or row.db_id != db_id: - task = await tasker.get_task(task_id) + task = await tasker.get_task(run_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}") + return {"run_id": run_id, "status": task.status, "progress": task.progress, "message": task.message} + raise ValueError(f"Run not found for {run_id}") - details = await self.eval_repo.list_result_details(task_id) - all_results = [ - { - "query": d.query_text, - "gold_chunk_ids": d.gold_chunk_ids, - "gold_answer": d.gold_answer, - "generated_answer": d.generated_answer, - "retrieved_chunks": d.retrieved_chunks, - "metrics": d.metrics or {}, - } - for d in details - ] - - if error_only: - filtered_results = [] - for item in all_results: - if item.get("metrics", {}).get("score", 1.0) <= 0.5: - filtered_results.append(item) - continue - metrics = item.get("metrics", {}) - has_low_recall = any(metrics.get(k, 1.0) < 0.3 for k in metrics if k.startswith("recall@")) - if has_low_recall: - filtered_results.append(item) - all_results = filtered_results - - total = len(all_results) start_idx = (page - 1) * page_size - end_idx = start_idx + page_size - paged_results = all_results[start_idx:end_idx] - + if error_only: + total = 0 + paged_items = [] + offset = 0 + batch_size = 200 + while True: + batch = await self.eval_repo.list_run_items(run_id, offset, batch_size) + if not batch: + break + for item in batch: + if not self._is_error_run_item(item): + continue + if start_idx <= total < start_idx + page_size: + paged_items.append(self._run_item_to_dict(item)) + total += 1 + offset += batch_size + else: + total = await self.eval_repo.count_run_items(run_id) + details = await self.eval_repo.list_run_items(run_id, start_idx, page_size) + paged_items = [self._run_item_to_dict(item) for item in details] return { - "task_id": row.task_id, + "run_id": row.run_id, "status": row.status, - "started_at": row.started_at.isoformat() if row.started_at else None, - "completed_at": row.completed_at.isoformat() if row.completed_at else None, - "total_questions": row.total_questions or 0, - "completed_questions": row.completed_questions or 0, + "started_at": format_utc_datetime(row.started_at), + "completed_at": format_utc_datetime(row.completed_at), + "total_items": row.total_items or 0, + "completed_items": row.completed_items or 0, "overall_score": row.overall_score, "retrieval_config": row.retrieval_config or {}, - "interim_results": paged_results, + "items": paged_items, "pagination": { "current_page": page, "page_size": page_size, @@ -636,12 +648,11 @@ class EvaluationService: }, } - async def delete_evaluation_result_by_db(self, db_id: str, task_id: str) -> None: - if not re.match(r"^eval_[a-f0-9]{8}$", task_id): - raise ValueError("Invalid task_id format") - row = await self.eval_repo.get_result(task_id) + async def delete_run(self, db_id: str, run_id: str) -> None: + if not re.match(r"^run_[a-f0-9]{8}$", run_id): + raise ValueError("Invalid run_id format") + row = await self.eval_repo.get_run(run_id) if row is None or row.db_id != db_id: - raise ValueError("Result not found") - await self.eval_repo.delete_result(task_id) - logger.info(f"成功删除评估结果: {task_id}") - return + raise ValueError("Run not found") + await self.eval_repo.delete_run(run_id) + logger.info(f"成功删除评估运行: {run_id}") diff --git a/backend/package/yuxi/storage/postgres/manager.py b/backend/package/yuxi/storage/postgres/manager.py index adabeab8..31d0badf 100644 --- a/backend/package/yuxi/storage/postgres/manager.py +++ b/backend/package/yuxi/storage/postgres/manager.py @@ -148,21 +148,85 @@ class PostgresManager(metaclass=SingletonMeta): "ALTER TABLE IF EXISTS knowledge_files ADD COLUMN IF NOT EXISTS created_by VARCHAR(64)", "ALTER TABLE IF EXISTS knowledge_files ADD COLUMN IF NOT EXISTS updated_by VARCHAR(64)", "ALTER TABLE IF EXISTS knowledge_files ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ", - "ALTER TABLE IF EXISTS evaluation_benchmarks ADD COLUMN IF NOT EXISTS data_file_path VARCHAR(1024)", - "ALTER TABLE IF EXISTS evaluation_benchmarks ADD COLUMN IF NOT EXISTS created_by VARCHAR(64)", - "ALTER TABLE IF EXISTS evaluation_benchmarks ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ", - "ALTER TABLE IF EXISTS evaluation_results ADD COLUMN IF NOT EXISTS metrics JSONB", - "ALTER TABLE IF EXISTS evaluation_results ADD COLUMN IF NOT EXISTS overall_score DOUBLE PRECISION", - "ALTER TABLE IF EXISTS evaluation_results ADD COLUMN IF NOT EXISTS total_questions INTEGER", - "ALTER TABLE IF EXISTS evaluation_results ADD COLUMN IF NOT EXISTS completed_questions INTEGER", - "ALTER TABLE IF EXISTS evaluation_results ADD COLUMN IF NOT EXISTS started_at TIMESTAMPTZ", - "ALTER TABLE IF EXISTS evaluation_results ADD COLUMN IF NOT EXISTS completed_at TIMESTAMPTZ", - "ALTER TABLE IF EXISTS evaluation_results ADD COLUMN IF NOT EXISTS created_by VARCHAR(64)", - "ALTER TABLE IF EXISTS evaluation_result_details ADD COLUMN IF NOT EXISTS gold_chunk_ids JSONB", - "ALTER TABLE IF EXISTS evaluation_result_details ADD COLUMN IF NOT EXISTS gold_answer TEXT", - "ALTER TABLE IF EXISTS evaluation_result_details ADD COLUMN IF NOT EXISTS generated_answer TEXT", - "ALTER TABLE IF EXISTS evaluation_result_details ADD COLUMN IF NOT EXISTS retrieved_chunks JSONB", - "ALTER TABLE IF EXISTS evaluation_result_details ADD COLUMN IF NOT EXISTS metrics JSONB", + "ALTER TABLE IF EXISTS evaluation_datasets ADD COLUMN IF NOT EXISTS created_by VARCHAR(64)", + "ALTER TABLE IF EXISTS evaluation_datasets ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ", + "ALTER TABLE IF EXISTS evaluation_datasets ADD COLUMN IF NOT EXISTS build_metadata JSONB", + "ALTER TABLE IF EXISTS evaluation_runs ADD COLUMN IF NOT EXISTS metrics JSONB", + "ALTER TABLE IF EXISTS evaluation_runs ADD COLUMN IF NOT EXISTS overall_score DOUBLE PRECISION", + "ALTER TABLE IF EXISTS evaluation_runs ADD COLUMN IF NOT EXISTS total_items INTEGER", + "ALTER TABLE IF EXISTS evaluation_runs ADD COLUMN IF NOT EXISTS completed_items INTEGER", + "ALTER TABLE IF EXISTS evaluation_runs ADD COLUMN IF NOT EXISTS started_at TIMESTAMPTZ", + "ALTER TABLE IF EXISTS evaluation_runs ADD COLUMN IF NOT EXISTS completed_at TIMESTAMPTZ", + "ALTER TABLE IF EXISTS evaluation_runs ADD COLUMN IF NOT EXISTS created_by VARCHAR(64)", + "ALTER TABLE IF EXISTS evaluation_run_items ADD COLUMN IF NOT EXISTS gold_chunk_ids JSONB", + "ALTER TABLE IF EXISTS evaluation_run_items ADD COLUMN IF NOT EXISTS gold_answer TEXT", + "ALTER TABLE IF EXISTS evaluation_run_items ADD COLUMN IF NOT EXISTS generated_answer TEXT", + "ALTER TABLE IF EXISTS evaluation_run_items ADD COLUMN IF NOT EXISTS retrieved_chunks JSONB", + "ALTER TABLE IF EXISTS evaluation_run_items ADD COLUMN IF NOT EXISTS metrics JSONB", + "ALTER TABLE IF EXISTS evaluation_run_items ADD COLUMN IF NOT EXISTS created_at TIMESTAMPTZ", + """ + CREATE TABLE IF NOT EXISTS evaluation_datasets ( + id SERIAL PRIMARY KEY, + dataset_id VARCHAR(64) NOT NULL UNIQUE, + db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + name VARCHAR(255) NOT NULL, + description TEXT, + item_count INTEGER DEFAULT 0, + has_gold_chunks BOOLEAN DEFAULT FALSE, + has_gold_answers BOOLEAN DEFAULT FALSE, + build_metadata JSONB, + created_by VARCHAR(64), + created_at TIMESTAMPTZ DEFAULT NOW(), + updated_at TIMESTAMPTZ DEFAULT NOW() + ) + """, + """ + CREATE TABLE IF NOT EXISTS evaluation_dataset_items ( + id SERIAL PRIMARY KEY, + item_id VARCHAR(64) NOT NULL UNIQUE, + dataset_id VARCHAR(64) NOT NULL REFERENCES evaluation_datasets(dataset_id) ON DELETE CASCADE, + db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + item_index INTEGER NOT NULL, + query_text TEXT NOT NULL, + gold_chunk_ids JSONB, + gold_answer TEXT, + created_at TIMESTAMPTZ DEFAULT NOW(), + CONSTRAINT uq_evaluation_dataset_items_dataset_index UNIQUE (dataset_id, item_index) + ) + """, + """ + CREATE TABLE IF NOT EXISTS evaluation_runs ( + id SERIAL PRIMARY KEY, + run_id VARCHAR(64) NOT NULL UNIQUE, + db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + dataset_id VARCHAR(64) REFERENCES evaluation_datasets(dataset_id) ON DELETE SET NULL, + status VARCHAR(32) DEFAULT 'running', + retrieval_config JSONB, + metrics JSONB, + overall_score DOUBLE PRECISION, + total_items INTEGER DEFAULT 0, + completed_items INTEGER DEFAULT 0, + started_at TIMESTAMPTZ DEFAULT NOW(), + completed_at TIMESTAMPTZ, + created_by VARCHAR(64) + ) + """, + """ + CREATE TABLE IF NOT EXISTS evaluation_run_items ( + id SERIAL PRIMARY KEY, + run_id VARCHAR(64) NOT NULL REFERENCES evaluation_runs(run_id) ON DELETE CASCADE, + dataset_item_id VARCHAR(64) REFERENCES evaluation_dataset_items(item_id) ON DELETE SET NULL, + item_index INTEGER NOT NULL, + query_text TEXT NOT NULL, + gold_chunk_ids JSONB, + gold_answer TEXT, + generated_answer TEXT, + retrieved_chunks JSONB, + metrics JSONB, + created_at TIMESTAMPTZ DEFAULT NOW(), + CONSTRAINT uq_evaluation_run_items_run_index UNIQUE (run_id, item_index) + ) + """, """ CREATE TABLE IF NOT EXISTS knowledge_chunks ( id SERIAL PRIMARY KEY, @@ -238,19 +302,25 @@ class PostgresManager(metaclass=SingletonMeta): # 扩展 db_id 字段长度以支持最长 75 字符的 ID(kb_private_ + 64字符hash) "ALTER TABLE IF EXISTS knowledge_bases ALTER COLUMN db_id TYPE VARCHAR(80)", "ALTER TABLE IF EXISTS knowledge_files ALTER COLUMN db_id TYPE VARCHAR(80)", - "ALTER TABLE IF EXISTS evaluation_benchmarks ALTER COLUMN db_id TYPE VARCHAR(80)", - "ALTER TABLE IF EXISTS evaluation_results ALTER COLUMN db_id TYPE VARCHAR(80)", + "ALTER TABLE IF EXISTS evaluation_datasets ALTER COLUMN db_id TYPE VARCHAR(80)", + "ALTER TABLE IF EXISTS evaluation_dataset_items ALTER COLUMN db_id TYPE VARCHAR(80)", + "ALTER TABLE IF EXISTS evaluation_runs ALTER COLUMN db_id TYPE VARCHAR(80)", "CREATE INDEX IF NOT EXISTS idx_kb_type ON knowledge_bases(kb_type)", "CREATE INDEX IF NOT EXISTS idx_kb_name ON knowledge_bases(name)", "CREATE INDEX IF NOT EXISTS idx_kf_db_id ON knowledge_files(db_id)", "CREATE INDEX IF NOT EXISTS idx_kf_parent ON knowledge_files(parent_id)", "CREATE INDEX IF NOT EXISTS idx_kf_status ON knowledge_files(status)", "CREATE INDEX IF NOT EXISTS idx_kf_hash ON knowledge_files(content_hash)", - "CREATE INDEX IF NOT EXISTS idx_eb_db_id ON evaluation_benchmarks(db_id)", - "CREATE INDEX IF NOT EXISTS idx_er_db_id ON evaluation_results(db_id)", - "CREATE INDEX IF NOT EXISTS idx_er_status ON evaluation_results(status)", - "CREATE INDEX IF NOT EXISTS idx_er_started ON evaluation_results(started_at DESC)", - "CREATE INDEX IF NOT EXISTS idx_erd_task ON evaluation_result_details(task_id)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_datasets_db_id ON evaluation_datasets(db_id)", + ( + "CREATE INDEX IF NOT EXISTS ix_evaluation_dataset_items_dataset_index " + "ON evaluation_dataset_items(dataset_id, item_index)" + ), + "CREATE INDEX IF NOT EXISTS ix_evaluation_dataset_items_db_id ON evaluation_dataset_items(db_id)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_runs_db_id ON evaluation_runs(db_id)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_runs_status ON evaluation_runs(status)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_runs_started ON evaluation_runs(started_at DESC)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_run_items_run_index ON evaluation_run_items(run_id, item_index)", "CREATE UNIQUE INDEX IF NOT EXISTS uq_knowledge_chunks_chunk_id ON knowledge_chunks(chunk_id)", "CREATE INDEX IF NOT EXISTS ix_knowledge_chunks_file_id ON knowledge_chunks(file_id)", "CREATE INDEX IF NOT EXISTS ix_knowledge_chunks_db_id ON knowledge_chunks(db_id)", diff --git a/backend/package/yuxi/storage/postgres/models_knowledge.py b/backend/package/yuxi/storage/postgres/models_knowledge.py index 29a86824..60c9c07e 100644 --- a/backend/package/yuxi/storage/postgres/models_knowledge.py +++ b/backend/package/yuxi/storage/postgres/models_knowledge.py @@ -187,68 +187,101 @@ class KnowledgeGraphTripleMention(Base): created_at = Column(DateTime(timezone=True), default=utc_now_naive) -class EvaluationBenchmark(Base): - """评估基准模型""" +class EvaluationDataset(Base): + """评估数据集模型""" - __tablename__ = "evaluation_benchmarks" - __table_args__ = (UniqueConstraint("benchmark_id", name="uq_evaluation_benchmarks_benchmark_id"),) + __tablename__ = "evaluation_datasets" + __table_args__ = (UniqueConstraint("dataset_id", name="uq_evaluation_datasets_dataset_id"),) id = Column(Integer, primary_key=True, autoincrement=True) - benchmark_id = Column(String(64), unique=True, nullable=False, index=True) + dataset_id = Column(String(64), unique=True, nullable=False, index=True) db_id = Column(String(80), ForeignKey("knowledge_bases.db_id", ondelete="CASCADE"), nullable=False, index=True) name = Column(String(255), nullable=False) description = Column(Text) - question_count = Column(Integer, default=0) + item_count = Column(Integer, default=0) has_gold_chunks = Column(Boolean, default=False) has_gold_answers = Column(Boolean, default=False) - data_file_path = Column(String(1024)) + build_metadata = Column(JSON_VALUE) created_by = Column(String(64)) created_at = Column(DateTime(timezone=True), default=utc_now_naive) updated_at = Column(DateTime(timezone=True), default=utc_now_naive, onupdate=utc_now_naive) -class EvaluationResult(Base): - """评估结果模型""" +class EvaluationDatasetItem(Base): + """评估数据集题目模型""" - __tablename__ = "evaluation_results" - __table_args__ = (UniqueConstraint("task_id", name="uq_evaluation_results_task_id"),) + __tablename__ = "evaluation_dataset_items" + __table_args__ = ( + UniqueConstraint("item_id", name="uq_evaluation_dataset_items_item_id"), + UniqueConstraint("dataset_id", "item_index", name="uq_evaluation_dataset_items_dataset_index"), + Index("ix_evaluation_dataset_items_dataset_index", "dataset_id", "item_index"), + ) id = Column(Integer, primary_key=True, autoincrement=True) - task_id = Column(String(64), unique=True, nullable=False, index=True) - db_id = Column(String(80), ForeignKey("knowledge_bases.db_id", ondelete="CASCADE"), nullable=False, index=True) - benchmark_id = Column( + item_id = Column(String(64), unique=True, nullable=False, index=True) + dataset_id = Column( String(64), - ForeignKey("evaluation_benchmarks.benchmark_id", ondelete="SET NULL"), + ForeignKey("evaluation_datasets.dataset_id", ondelete="CASCADE"), + nullable=False, + index=True, + ) + db_id = Column(String(80), ForeignKey("knowledge_bases.db_id", ondelete="CASCADE"), nullable=False, index=True) + item_index = Column(Integer, nullable=False) + query_text = Column(Text, nullable=False) + gold_chunk_ids = Column(JSON_VALUE) + gold_answer = Column(Text) + created_at = Column(DateTime(timezone=True), default=utc_now_naive) + + +class EvaluationRun(Base): + """评估运行模型""" + + __tablename__ = "evaluation_runs" + __table_args__ = (UniqueConstraint("run_id", name="uq_evaluation_runs_run_id"),) + + id = Column(Integer, primary_key=True, autoincrement=True) + run_id = Column(String(64), unique=True, nullable=False, index=True) + db_id = Column(String(80), ForeignKey("knowledge_bases.db_id", ondelete="CASCADE"), nullable=False, index=True) + dataset_id = Column( + String(64), + ForeignKey("evaluation_datasets.dataset_id", ondelete="SET NULL"), index=True, ) status = Column(String(32), default="running", index=True) retrieval_config = Column(JSON_VALUE) metrics = Column(JSON_VALUE) overall_score = Column(Float) - total_questions = Column(Integer, default=0) - completed_questions = Column(Integer, default=0) + total_items = Column(Integer, default=0) + completed_items = Column(Integer, default=0) started_at = Column(DateTime(timezone=True), default=utc_now_naive, index=True) completed_at = Column(DateTime(timezone=True)) created_by = Column(String(64)) -class EvaluationResultDetail(Base): - """评估结果详情模型""" +class EvaluationRunItem(Base): + """评估逐题结果模型""" - __tablename__ = "evaluation_result_details" - __table_args__ = (UniqueConstraint("task_id", "query_index", name="uq_evaluation_result_details_task_query"),) + __tablename__ = "evaluation_run_items" + __table_args__ = ( + UniqueConstraint("run_id", "item_index", name="uq_evaluation_run_items_run_index"), + Index("ix_evaluation_run_items_run_index", "run_id", "item_index"), + ) id = Column(Integer, primary_key=True, autoincrement=True) - task_id = Column( + run_id = Column( String(64), - ForeignKey("evaluation_results.task_id", ondelete="CASCADE"), + ForeignKey("evaluation_runs.run_id", ondelete="CASCADE"), nullable=False, index=True, ) - query_index = Column(Integer, nullable=False) + dataset_item_id = Column( + String(64), ForeignKey("evaluation_dataset_items.item_id", ondelete="SET NULL"), index=True + ) + item_index = Column(Integer, nullable=False) query_text = Column(Text, nullable=False) gold_chunk_ids = Column(JSON_VALUE) gold_answer = Column(Text) generated_answer = Column(Text) retrieved_chunks = Column(JSON_VALUE) metrics = Column(JSON_VALUE) + created_at = Column(DateTime(timezone=True), default=utc_now_naive) diff --git a/backend/server/routers/knowledge_eval_router.py b/backend/server/routers/knowledge_eval_router.py index 8821d932..4e558528 100644 --- a/backend/server/routers/knowledge_eval_router.py +++ b/backend/server/routers/knowledge_eval_router.py @@ -1,226 +1,261 @@ import traceback +from typing import Any +from urllib.parse import quote -from fastapi import APIRouter, HTTPException, Depends, File, Form, Body, UploadFile -from fastapi.responses import FileResponse -from yuxi.storage.postgres.models_business import User +from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile +from fastapi.responses import Response +from pydantic import BaseModel, Field from server.utils.auth_middleware import get_admin_user +from yuxi.knowledge.eval.benchmark_generation import ( + DEFAULT_BENCHMARK_GENERATION_CONCURRENCY, + MAX_BENCHMARK_GENERATION_CONCURRENCY, +) +from yuxi.storage.postgres.models_business import User from yuxi.utils import logger -# 创建路由器 + evaluation = APIRouter(prefix="/evaluation", tags=["evaluation"]) -# 移除旧详情接口,统一使用带 db_id 的接口 -# ============================================================================ -# 评估基准 -# ============================================================================ +class GenerateDatasetRequest(BaseModel): + name: str = Field(default="自动生成评估数据集", min_length=1, max_length=100) + description: str = "" + count: int = Field(default=10, ge=1, le=100) + neighbors_count: int = Field(default=1, ge=0, le=10) + concurrency_count: int = Field( + default=DEFAULT_BENCHMARK_GENERATION_CONCURRENCY, + ge=1, + le=MAX_BENCHMARK_GENERATION_CONCURRENCY, + ) + llm_model_spec: str = Field(..., min_length=1) -@evaluation.get("/databases/{db_id}/benchmarks/{benchmark_id}") -async def get_evaluation_benchmark_by_db( - db_id: str, benchmark_id: str, page: int = 1, page_size: int = 10, current_user: User = Depends(get_admin_user) -): - """根据 db_id 获取评估基准详情(支持分页)""" - from yuxi.services.evaluation_service import EvaluationService - - try: - # 验证分页参数 - if page < 1: - raise HTTPException(status_code=400, detail="页码必须大于0") - if page_size < 1 or page_size > 100: - raise HTTPException(status_code=400, detail="每页大小必须在1-100之间") - - service = EvaluationService() - benchmark = await service.get_benchmark_detail_by_db(db_id, benchmark_id, page, page_size) - return {"message": "success", "data": benchmark} - except HTTPException: - raise - except Exception as e: - logger.error(f"获取评估基准详情失败: {e}, {traceback.format_exc()}") - raise HTTPException(status_code=500, detail=f"获取评估基准详情失败: {str(e)}") +class RunEvaluationRequest(BaseModel): + dataset_id: str = Field(..., min_length=1) + retrieval_config: dict[str, Any] = Field(default_factory=dict, alias="model_config") -@evaluation.delete("/benchmarks/{benchmark_id}") -async def delete_evaluation_benchmark(benchmark_id: str, current_user: User = Depends(get_admin_user)): - """删除评估基准""" - from yuxi.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("/benchmarks/{benchmark_id}/download") -async def download_evaluation_benchmark(benchmark_id: str, current_user: User = Depends(get_admin_user)): - """下载评估基准文件""" - from yuxi.services.evaluation_service import EvaluationService - - try: - service = EvaluationService() - download_info = await service.get_benchmark_download_info(benchmark_id) - return FileResponse( - path=download_info["file_path"], - filename=download_info["filename"], - media_type="application/x-ndjson", - ) - except ValueError as e: - if "not found" in str(e).lower(): - raise HTTPException(status_code=404, detail=str(e)) - logger.error(f"下载评估基准失败: {e}, {traceback.format_exc()}") - raise HTTPException(status_code=500, detail=f"下载评估基准失败: {str(e)}") - 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}/results/{task_id}") -async def get_evaluation_results_by_db( - db_id: str, - task_id: str, - page: int = 1, - page_size: int = 20, - error_only: bool = False, - current_user: User = Depends(get_admin_user), -): - """获取评估结果(带 db_id,支持分页)""" - from yuxi.services.evaluation_service import EvaluationService - - try: - # 验证分页参数 - if page < 1: - raise HTTPException(status_code=400, detail="页码必须大于0") - if page_size < 1 or page_size > 100: - raise HTTPException(status_code=400, detail="每页大小必须在1-100之间") - - service = EvaluationService() - results = await service.get_evaluation_results_by_db( - db_id, task_id, page=page, page_size=page_size, error_only=error_only - ) - 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("/databases/{db_id}/results/{task_id}") -async def delete_evaluation_result_by_db(db_id: str, task_id: str, current_user: User = Depends(get_admin_user)): - """删除评估结果(带 db_id)""" - from yuxi.services.evaluation_service import EvaluationService - - try: - service = EvaluationService() - await service.delete_evaluation_result_by_db(db_id, 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)}") - - -# ============================================================================ -# RAG评估 -# ============================================================================ - - -@evaluation.post("/databases/{db_id}/benchmarks/upload") -async def upload_evaluation_benchmark( +@evaluation.post("/databases/{db_id}/datasets/upload") +async def upload_evaluation_dataset( db_id: str, file: UploadFile = File(...), name: str = Form(...), description: str = Form(""), current_user: User = Depends(get_admin_user), ): - """上传评估基准文件""" + """上传评估数据集""" from yuxi.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( + result = await service.upload_dataset( db_id=db_id, - file_content=content, + file_content=await file.read(), filename=file.filename, name=name, description=description, created_by=current_user.uid, ) - 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)}") + 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)): - """获取知识库的评估基准列表""" +@evaluation.get("/databases/{db_id}/datasets") +async def list_evaluation_datasets(db_id: str, current_user: User = Depends(get_admin_user)): + """获取知识库的评估数据集列表""" from yuxi.services.evaluation_service import EvaluationService try: service = EvaluationService() - benchmarks = await service.get_benchmarks(db_id) - return {"message": "success", "data": benchmarks} + datasets = await service.list_datasets(db_id) + return {"message": "success", "data": datasets} except Exception as e: - logger.error(f"获取评估基准列表失败: {e}, {traceback.format_exc()}") - raise HTTPException(status_code=500, detail=f"获取评估基准列表失败: {str(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) +@evaluation.get("/databases/{db_id}/datasets/{dataset_id}") +async def get_evaluation_dataset( + db_id: str, dataset_id: str, page: int = 1, page_size: int = 10, current_user: User = Depends(get_admin_user) ): - """自动生成评估基准""" + """获取评估数据集详情""" + from yuxi.services.evaluation_service import EvaluationService + + try: + if page < 1: + raise HTTPException(status_code=400, detail="页码必须大于0") + if page_size < 1 or page_size > 100: + raise HTTPException(status_code=400, detail="每页大小必须在1-100之间") + + service = EvaluationService() + dataset = await service.get_dataset_detail(db_id, dataset_id, page, page_size) + return {"message": "success", "data": dataset} + except HTTPException: + raise + except ValueError as e: + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + logger.error(f"获取评估数据集详情失败: {e}, {traceback.format_exc()}") + raise HTTPException(status_code=500, detail=f"获取评估数据集详情失败: {str(e)}") + + +@evaluation.get("/datasets/{dataset_id}/download") +async def download_evaluation_dataset(dataset_id: str, current_user: User = Depends(get_admin_user)): + """导出评估数据集 JSONL""" from yuxi.services.evaluation_service import EvaluationService try: service = EvaluationService() - result = await service.generate_benchmark(db_id=db_id, params=params, created_by=current_user.uid) + export_info = await service.export_dataset_jsonl(dataset_id) + filename = export_info["filename"] + return Response( + content=export_info["content"].encode("utf-8"), + media_type="application/x-ndjson", + headers={"Content-Disposition": f"attachment; filename*=UTF-8''{quote(filename)}"}, + ) + except ValueError as e: + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + logger.error(f"导出评估数据集失败: {e}, {traceback.format_exc()}") + raise HTTPException(status_code=500, detail=f"导出评估数据集失败: {str(e)}") + + +@evaluation.delete("/datasets/{dataset_id}") +async def delete_evaluation_dataset(dataset_id: str, current_user: User = Depends(get_admin_user)): + """删除评估数据集""" + from yuxi.services.evaluation_service import EvaluationService + + try: + service = EvaluationService() + await service.delete_dataset(dataset_id) + return {"message": "success", "data": None} + except ValueError as e: + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=400, detail=str(e)) + 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}/datasets/generate") +async def generate_evaluation_dataset( + db_id: str, request: GenerateDatasetRequest, current_user: User = Depends(get_admin_user) +): + """自动生成评估数据集""" + from yuxi.services.evaluation_service import EvaluationService + + try: + service = EvaluationService() + result = await service.generate_dataset( + db_id=db_id, + name=request.name, + description=request.description, + count=request.count, + neighbors_count=request.neighbors_count, + concurrency_count=request.concurrency_count, + llm_model_spec=request.llm_model_spec, + created_by=current_user.uid, + ) 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)}") + 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)): +@evaluation.post("/databases/{db_id}/runs") +async def run_evaluation(db_id: str, request: RunEvaluationRequest, current_user: User = Depends(get_admin_user)): """运行RAG评估""" from yuxi.services.evaluation_service import EvaluationService try: service = EvaluationService() - task_id = await service.run_evaluation( + run_id = await service.run_evaluation( db_id=db_id, - benchmark_id=params.get("benchmark_id"), - model_config=params.get("model_config", {}), + dataset_id=request.dataset_id, + model_config=request.retrieval_config, created_by=current_user.uid, ) - return {"message": "success", "data": {"task_id": task_id}} + return {"message": "success", "data": {"run_id": run_id}} + except ValueError as e: + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=400, detail=str(e)) 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)): - """获取知识库的评估历史记录""" +@evaluation.get("/databases/{db_id}/runs") +async def list_evaluation_runs(db_id: str, current_user: User = Depends(get_admin_user)): + """获取知识库评估运行历史""" from yuxi.services.evaluation_service import EvaluationService try: service = EvaluationService() - history = await service.get_evaluation_history(db_id) - return {"message": "success", "data": history} + runs = await service.list_runs(db_id) + return {"message": "success", "data": runs} except Exception as e: - logger.error(f"获取评估历史失败: {e}, {traceback.format_exc()}") - raise HTTPException(status_code=500, detail=f"获取评估历史失败: {str(e)}") + logger.error(f"获取评估运行历史失败: {e}, {traceback.format_exc()}") + raise HTTPException(status_code=500, detail=f"获取评估运行历史失败: {str(e)}") + + +@evaluation.get("/databases/{db_id}/runs/{run_id}") +async def get_evaluation_run_results( + db_id: str, + run_id: str, + page: int = 1, + page_size: int = 20, + error_only: bool = False, + current_user: User = Depends(get_admin_user), +): + """获取评估运行结果""" + from yuxi.services.evaluation_service import EvaluationService + + try: + if page < 1: + raise HTTPException(status_code=400, detail="页码必须大于0") + if page_size < 1 or page_size > 100: + raise HTTPException(status_code=400, detail="每页大小必须在1-100之间") + + service = EvaluationService() + results = await service.get_run_results(db_id, run_id, page=page, page_size=page_size, error_only=error_only) + return {"message": "success", "data": results} + except HTTPException: + raise + except ValueError as e: + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + logger.error(f"获取评估运行结果失败: {e}, {traceback.format_exc()}") + raise HTTPException(status_code=500, detail=f"获取评估运行结果失败: {str(e)}") + + +@evaluation.delete("/databases/{db_id}/runs/{run_id}") +async def delete_evaluation_run(db_id: str, run_id: str, current_user: User = Depends(get_admin_user)): + """删除评估运行""" + from yuxi.services.evaluation_service import EvaluationService + + try: + service = EvaluationService() + await service.delete_run(db_id, run_id) + return {"message": "success", "data": None} + except ValueError as e: + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + logger.error(f"删除评估运行失败: {e}, {traceback.format_exc()}") + raise HTTPException(status_code=500, detail=f"删除评估运行失败: {str(e)}") diff --git a/backend/server/utils/lifespan.py b/backend/server/utils/lifespan.py index 10380f25..52e2041a 100644 --- a/backend/server/utils/lifespan.py +++ b/backend/server/utils/lifespan.py @@ -22,7 +22,7 @@ async def lifespan(app: FastAPI): # 初始化数据库连接 try: pg_manager.initialize() - await pg_manager.create_business_tables() + await pg_manager.create_tables() await pg_manager.ensure_business_schema() await pg_manager.ensure_knowledge_schema() except Exception as e: