From c6cde2500b54f0d9ea41a70f21c779b51779bb43 Mon Sep 17 00:00:00 2001 From: yxx <2412119399@qq.com> Date: Wed, 2 Sep 2026 23:20:34 +0800 Subject: [PATCH] =?UTF-8?q?fix(backend):=20=E8=90=BD=E5=AE=9E=20PR=20#9=20?= =?UTF-8?q?=E8=AF=84=E5=AE=A1=E6=84=8F=E8=A7=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 检索调优参数(rrf_k/rerank/rerank_candidates/score_threshold)透传到引擎实际执行 - Recall 去重,避免同一 Note 多 Block 重复导致 Recall 超 1 - RAG 运行改为后台异步执行:创建即 queued + 202,支持取消与 SSE 实时事件 - 数据集元数据校验,坏文件隔离跳过;citation_required 语义修正 - modes 空/重复校验;配置快照记录模型版本与索引元信息 Co-Authored-By: Claude Code --- backend/app/benchmarks/datasets.py | 40 ++++++-- backend/app/benchmarks/metrics.py | 9 +- backend/app/benchmarks/rag.py | 14 ++- backend/app/benchmarks/service.py | 129 +++++++++++++++++++++----- backend/app/contracts.py | 22 ++++- backend/app/retrieval/embedding.py | 2 + backend/app/retrieval/engine.py | 28 ++++-- backend/app/retrieval/reranker.py | 2 + backend/app/routes.py | 20 +++- backend/tests/test_benchmark.py | 144 ++++++++++++++++++++++++++--- 10 files changed, 350 insertions(+), 60 deletions(-) diff --git a/backend/app/benchmarks/datasets.py b/backend/app/benchmarks/datasets.py index 1b822da..3a3866b 100644 --- a/backend/app/benchmarks/datasets.py +++ b/backend/app/benchmarks/datasets.py @@ -11,7 +11,7 @@ import json from dataclasses import dataclass, field from pathlib import Path -from pydantic import ValidationError +from pydantic import BaseModel, Field, ValidationError from app.config import get_settings from app.contracts import ( @@ -34,6 +34,20 @@ class RAGDataset: content_hash: str = "" +class _DatasetMeta(BaseModel): + """Dataset 元数据的最小校验模型。 + + list_datasets 用它逐文件校验元信息字段结构,把「合法 JSON 但字段类型错误」 + (如 cases: 42)这类损坏文件隔离掉,而不是让 len() 抛 TypeError 拖垮整个列表。 + """ + + dataset_id: str = Field(min_length=1) + kind: str = "" + version: str = "" + description: str = "" + cases: list = Field(default_factory=list) + + def _datasets_dir() -> Path: return get_settings().benchmark_datasets_path @@ -109,6 +123,14 @@ def _dataset_from_raw(raw: dict, raw_bytes: bytes, kind: BenchmarkKind) -> RAGDa f"Dataset case '{parsed.case_id}' must declare expected_note_ids or expected_block_ids.", {"dataset_id": dataset_id, "case_id": parsed.case_id}, ) + # citation_required=true 时必须声明 expected_block_ids,否则无法计算 Citation Hit Rate + if parsed.citation_required and not parsed.expected_block_ids: + raise ApiError( + 422, + "BENCHMARK_DATASET_INVALID", + f"Dataset case '{parsed.case_id}' requires expected_block_ids when citation_required is true.", + {"dataset_id": dataset_id, "case_id": parsed.case_id}, + ) cases.append(parsed) return RAGDataset( @@ -124,23 +146,25 @@ def _dataset_from_raw(raw: dict, raw_bytes: bytes, kind: BenchmarkKind) -> RAGDa def list_datasets(kind: BenchmarkKind) -> list[BenchmarkDatasetInfo]: """枚举受控目录下指定 kind 的数据集元信息(不含 Case 内容)。 - 个别文件损坏时跳过而非整体失败,保证列表接口健壮;损坏细节由 load_dataset 抛出。 + 逐文件用 _DatasetMeta 校验元信息字段结构,单个损坏文件隔离跳过而非整体失败, + 保证列表接口健壮;损坏细节由 load_dataset 抛出。 """ infos: list[BenchmarkDatasetInfo] = [] for path in _dataset_files(): try: raw, raw_bytes = _read_json(path) - except ApiError: + meta = _DatasetMeta.model_validate(raw) + except (ApiError, ValidationError): continue - if raw.get("kind", kind.value) != kind.value: + if meta.kind not in ("", kind.value): continue infos.append( BenchmarkDatasetInfo( - dataset_id=raw.get("dataset_id", path.stem), + dataset_id=meta.dataset_id, kind=kind, - version=str(raw.get("version", "")), - description=str(raw.get("description", "")), - case_count=len(raw.get("cases", [])), + version=meta.version, + description=meta.description, + case_count=len(meta.cases), content_hash=_content_hash(raw_bytes), ) ) diff --git a/backend/app/benchmarks/metrics.py b/backend/app/benchmarks/metrics.py index 4ef66d0..340e0bd 100644 --- a/backend/app/benchmarks/metrics.py +++ b/backend/app/benchmarks/metrics.py @@ -13,11 +13,14 @@ def hit_at_k(retrieved: list[str], expected: set[str], k: int) -> bool: def recall_at_k(retrieved: list[str], expected: set[str], k: int) -> float: - """前 k 个结果召回的期望 id 占比;期望为空时视为 0。""" + """前 k 个结果召回的期望 id 占比;期望为空时视为 0。 + + 结果先去重:检索结果是 Block 级,同一 Note 可能经多个 Block 重复出现, + 直接逐项计数会把同一 Note 算多次、导致 Recall 超过 1。 + """ if not expected: return 0.0 - hits = sum(1 for item in retrieved[:k] if item in expected) - return hits / len(expected) + return len(set(retrieved[:k]) & expected) / len(expected) def reciprocal_rank(retrieved: list[str], expected: set[str]) -> float: diff --git a/backend/app/benchmarks/rag.py b/backend/app/benchmarks/rag.py index 8cfcaf6..2ec4ec3 100644 --- a/backend/app/benchmarks/rag.py +++ b/backend/app/benchmarks/rag.py @@ -23,14 +23,20 @@ from app.contracts import ( from app.retrieval.engine import engine +class BenchmarkCancelled(Exception): + """运行在 Case 之间被取消时抛出,用于中断后台执行并标记 cancelled。""" + + async def run_rag( dataset: RAGDataset, request: RAGRunRequest, on_case: Callable[[RAGCaseResult, int, int], None] | None = None, + should_cancel: Callable[[], bool] | None = None, ) -> tuple[dict[str, RAGMetrics], list[RAGCaseResult]]: """执行 RAG Benchmark,返回 (按 mode 聚合的指标, 全部逐样本结果)。 on_case 在每个样本求值完成后回调 (result, done, total),供上层更新进度与事件。 + should_cancel 在每个样本开始前被检查;返回 True 时抛出 BenchmarkCancelled 中断运行。 """ total = len(request.modes) * len(dataset.cases) * request.repeat done = 0 @@ -39,6 +45,8 @@ async def run_rag( for mode in request.modes: for case in dataset.cases: for repeat in range(request.repeat): + if should_cancel is not None and should_cancel(): + raise BenchmarkCancelled() result = await _evaluate_one(case, mode, request, repeat) results.append(result) done += 1 @@ -57,6 +65,10 @@ async def _evaluate_one( mode=mode, limit=request.retrieval.top_k, include_snippet=False, + rrf_k=request.retrieval.rrf_k, + rerank=request.retrieval.rerank, + rerank_candidates=request.retrieval.rerank_candidates, + score_threshold=request.retrieval.score_threshold, ) start = time.perf_counter() try: @@ -89,7 +101,7 @@ async def _evaluate_one( recall=m.recall_at_k(retrieved_note_ids, expected_notes, k), reciprocal_rank=m.reciprocal_rank(retrieved_note_ids, expected_notes), citation_hit=m.citation_hit(retrieved_block_ids, expected_blocks), - citation_applicable=bool(case.expected_block_ids), + citation_applicable=case.citation_required, ) diff --git a/backend/app/benchmarks/service.py b/backend/app/benchmarks/service.py index d35ac2a..dcad814 100644 --- a/backend/app/benchmarks/service.py +++ b/backend/app/benchmarks/service.py @@ -1,19 +1,22 @@ """Benchmark 服务:运行注册表、配置快照与报告组装。 -MVP 阶段运行是同步的(与 index_service 一致):POST 创建后立即执行完并返回 -completed 的 BenchmarkRun。运行记录、事件与报告暂存内存(_runs/_events/_reports), -不持久化到 SQLite;后续接入异步任务队列时再落库。 +RAG Benchmark 采用「创建即返回 queued、后台 Task 异步执行」的模式(与 index_service +的 rebuild 一致):POST 创建后立即返回 202 queued 的 BenchmarkRun,由受管 asyncio.Task +在后台逐 Case 求值,进度与事件实时写入内存注册表,供 SSE 订阅。运行记录、事件与报告 +暂存内存(_runs/_events/_reports),不持久化到 SQLite;后续接入异步任务队列时再落库。 """ from __future__ import annotations +import asyncio import sys from datetime import datetime, timezone from uuid import uuid4 +from app import repository from app.benchmarks import datasets from app.benchmarks.datasets import RAGDataset -from app.benchmarks.rag import run_rag +from app.benchmarks.rag import BenchmarkCancelled, run_rag from app.config import get_settings from app.contracts import ( BenchmarkEvent, @@ -32,6 +35,9 @@ from app.retrieval.engine import engine _runs: dict[str, BenchmarkRun] = {} _events: dict[str, list[BenchmarkEvent]] = {} _reports: dict[str, BenchmarkReport] = {} +_tasks: dict[str, asyncio.Task] = {} +_subscribers: dict[str, list[asyncio.Queue[BenchmarkEvent]]] = {} +_cancel_flags: dict[str, asyncio.Event] = {} MAX_RUNS = 100 @@ -46,6 +52,9 @@ def _remember(run: BenchmarkRun) -> None: _runs.pop(oldest, None) _events.pop(oldest, None) _reports.pop(oldest, None) + _tasks.pop(oldest, None) + _subscribers.pop(oldest, None) + _cancel_flags.pop(oldest, None) def _config_snapshot(request: RAGRunRequest, dataset: RAGDataset) -> dict: @@ -58,8 +67,16 @@ def _config_snapshot(request: RAGRunRequest, dataset: RAGDataset) -> dict: "modes": [m.value for m in request.modes], "retrieval": request.retrieval.model_dump(), "repeat": request.repeat, - "embedding": {"model_id": engine.embedding.model_id, "dim": engine.embedding.dim}, - "reranker": {"model_id": engine.reranker.model_id}, + "embedding": { + "model_id": engine.embedding.model_id, + "version": engine.embedding.version, + "dim": engine.embedding.dim, + }, + "reranker": { + "model_id": engine.reranker.model_id, + "version": engine.reranker.version, + }, + "index_meta": repository.get_index_meta(), "app": {"version": settings.version, "environment": settings.environment}, "python": sys.version.split()[0], "metadata": request.metadata, @@ -67,7 +84,7 @@ def _config_snapshot(request: RAGRunRequest, dataset: RAGDataset) -> dict: async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun: - """创建并同步执行一次 RAG Benchmark,返回 completed 的 BenchmarkRun。""" + """创建一次 RAG Benchmark,立即返回 queued 的 BenchmarkRun,由后台 Task 执行。""" dataset = datasets.load_dataset(request.dataset_id, BenchmarkKind.rag) run_id = "benchmark_" + uuid4().hex[:12] snapshot = _config_snapshot(request, dataset) @@ -77,24 +94,41 @@ async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun: kind=BenchmarkKind.rag, dataset_id=dataset.dataset_id, dataset_hash=dataset.content_hash, - status=BenchmarkStatus.running, + status=BenchmarkStatus.queued, progress=0.0, config_snapshot=snapshot, created_at=_now(), - started_at=_now(), ) _remember(run) _events[run_id] = [] + _subscribers[run_id] = [] + _cancel_flags[run_id] = asyncio.Event() + _tasks[run_id] = asyncio.create_task(_execute_rag(run_id, request, dataset, snapshot)) + return run + + +async def _execute_rag( + run_id: str, request: RAGRunRequest, dataset: RAGDataset, snapshot: dict +) -> None: + """后台执行 RAG Benchmark,实时更新进度/事件,结束后写入报告并关闭订阅。""" + cancel_event = _cancel_flags[run_id] def emit(event_type: BenchmarkEventType, data: dict) -> None: sequence = len(_events[run_id]) - _events[run_id].append( - BenchmarkEvent( - event=event_type, run_id=run_id, sequence=sequence, - data=data, timestamp=_now(), - ) + event = BenchmarkEvent( + event=event_type, run_id=run_id, sequence=sequence, data=data, timestamp=_now() ) + _events[run_id].append(event) + for queue in _subscribers.get(run_id, []): + queue.put_nowait(event) + def finish() -> None: + _subscribers.pop(run_id, None) + _cancel_flags.pop(run_id, None) + + _runs[run_id] = _runs[run_id].model_copy( + update={"status": BenchmarkStatus.running, "started_at": _now()} + ) emit( BenchmarkEventType.run_started, {"dataset_id": dataset.dataset_id, "modes": [m.value for m in request.modes]}, @@ -107,8 +141,31 @@ async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun: emit(BenchmarkEventType.case_completed, result.model_dump(mode="json")) try: - metrics_by_mode, results = await run_rag(dataset, request, on_case=on_case) - except Exception as exc: + metrics_by_mode, results = await run_rag( + dataset, + request, + on_case=on_case, + should_cancel=cancel_event.is_set, + ) + except BenchmarkCancelled: + _runs[run_id] = _runs[run_id].model_copy( + update={ + "status": BenchmarkStatus.cancelled, + "progress": 1.0, + "completed_at": _now(), + } + ) + _reports[run_id] = BenchmarkReport( + run_id=run_id, + kind=BenchmarkKind.rag, + dataset_id=dataset.dataset_id, + dataset_hash=dataset.content_hash, + status=BenchmarkStatus.cancelled, + config_snapshot=snapshot, + ) + finish() + return + except Exception as exc: # 单次运行失败不拖垮服务,记录错误后结束 _runs[run_id] = _runs[run_id].model_copy( update={ "status": BenchmarkStatus.failed, @@ -127,7 +184,8 @@ async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun: config_snapshot=snapshot, error=str(exc), ) - raise ApiError(500, "BENCHMARK_RUN_FAILED", str(exc), {"run_id": run_id}) from exc + finish() + return metrics = {mode: m.model_dump() for mode, m in metrics_by_mode.items()} _runs[run_id] = _runs[run_id].model_copy( @@ -149,7 +207,7 @@ async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun: metrics=metrics, cases=results, ) - return _runs[run_id] + finish() def list_runs( @@ -181,13 +239,38 @@ def get_events(run_id: str) -> list[BenchmarkEvent]: def cancel_run(run_id: str) -> BenchmarkRun | None: - """取消运行:同步 MVP 下运行通常已结束,仅对仍在排队/运行的记录置为 cancelled。""" + """取消运行:对 queued/running 设置取消标志,后台 Task 在 Case 边界检查后置为 cancelled。""" run = _runs.get(run_id) if run is None: return None if run.status in (BenchmarkStatus.queued, BenchmarkStatus.running): - run = run.model_copy( - update={"status": BenchmarkStatus.cancelled, "completed_at": _now()} - ) - _runs[run_id] = run + _cancel_flags[run_id].set() return run + + +def subscribe(run_id: str) -> asyncio.Queue[BenchmarkEvent] | None: + """订阅运行事件流;运行已结束(completed/failed/cancelled)时返回 None。""" + run = _runs.get(run_id) + if run is None or run.status in ( + BenchmarkStatus.completed, + BenchmarkStatus.failed, + BenchmarkStatus.cancelled, + ): + return None + queue: asyncio.Queue[BenchmarkEvent] = asyncio.Queue() + _subscribers.setdefault(run_id, []).append(queue) + return queue + + +def unsubscribe(run_id: str, queue: asyncio.Queue[BenchmarkEvent]) -> None: + subscribers = _subscribers.get(run_id) + if subscribers and queue in subscribers: + subscribers.remove(queue) + + +async def wait_for_run(run_id: str) -> BenchmarkRun: + """等待后台任务结束(测试/轮询用);无任务时直接返回当前状态。""" + task = _tasks.get(run_id) + if task is not None: + await task + return _runs.get(run_id) diff --git a/backend/app/contracts.py b/backend/app/contracts.py index 49fd6a2..aac441d 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -2,7 +2,7 @@ from datetime import datetime from enum import Enum from typing import Any, Literal -from pydantic import BaseModel, ConfigDict, Field, SecretStr +from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator class Contract(BaseModel): @@ -144,6 +144,12 @@ class SearchRequest(Contract): limit: int = Field(default=20, ge=1, le=100) offset: int = Field(default=0, ge=0) include_snippet: bool = True + # 检索调优参数(Benchmark 与 Skill 共用):控制 RRF / 精排 / 候选池 / 分数阈值。 + # rerank_candidates=None 表示对全部候选精排(保留原有行为),Benchmark 传显式值。 + rrf_k: int = Field(default=60, ge=1) + rerank: bool = True + rerank_candidates: int | None = Field(default=None, ge=1) + score_threshold: float = Field(default=0.0, ge=0.0) class Citation(Contract): @@ -678,8 +684,8 @@ class RAGDatasetCase(Contract): class RAGRetrievalConfig(Contract): - """RAG Benchmark 的检索参数。top_k 映射到 SearchRequest.limit,其余参数当前 - 记录进 config_snapshot,由 RetrievalProfile 共享(§9.9)落地后再接入引擎。""" + """RAG Benchmark 的检索参数。top_k 映射到 SearchRequest.limit, + 其余参数透传到 SearchRequest,由检索引擎实际执行。""" top_k: int = Field(default=10, ge=1, le=100) rrf_k: int = Field(default=60, ge=1) @@ -691,12 +697,20 @@ class RAGRetrievalConfig(Contract): class RAGRunRequest(Contract): dataset_id: str = Field(min_length=1) modes: list[SearchMode] = Field( - default_factory=lambda: [SearchMode.fts, SearchMode.vector, SearchMode.hybrid] + default_factory=lambda: [SearchMode.fts, SearchMode.vector, SearchMode.hybrid], + min_length=1, ) retrieval: RAGRetrievalConfig = Field(default_factory=RAGRetrievalConfig) repeat: int = Field(default=1, ge=1, le=10) metadata: dict[str, Any] = Field(default_factory=dict) + @field_validator("modes") + @classmethod + def _no_duplicate_modes(cls, value: list[SearchMode]) -> list[SearchMode]: + if len(value) != len(set(value)): + raise ValueError("modes must not contain duplicates") + return value + class RAGMetrics(Contract): hit_at_1: float = 0.0 diff --git a/backend/app/retrieval/embedding.py b/backend/app/retrieval/embedding.py index 6449d4e..551eb09 100644 --- a/backend/app/retrieval/embedding.py +++ b/backend/app/retrieval/embedding.py @@ -19,6 +19,7 @@ class EmbeddingProvider(Protocol): """统一 Embedding 接口(与文档一致)。""" model_id: str + version: str dim: int async def embed_documents(self, texts: list[str]) -> list[list[float]]: ... @@ -33,6 +34,7 @@ class HashEmbeddingProvider: """ model_id = "hash-v1" + version = "1" dim = EMBEDDING_DIM async def embed_documents(self, texts: list[str]) -> list[list[float]]: diff --git a/backend/app/retrieval/engine.py b/backend/app/retrieval/engine.py index 432edc4..3ac3c2d 100644 --- a/backend/app/retrieval/engine.py +++ b/backend/app/retrieval/engine.py @@ -84,7 +84,7 @@ class RetrievalEngine: elif request.mode == SearchMode.vector: candidate_scores = vec_scores else: # hybrid:RRF 融合 - candidate_scores = rrf_fuse([fts_ranked, vec_ranked]) + candidate_scores = rrf_fuse([fts_ranked, vec_ranked], k=request.rrf_k) if not candidate_scores: return self._empty(request) @@ -97,14 +97,23 @@ class RetrievalEngine: if not filtered: return self._empty(request) - # 4. 排序 / 精排 + # 4. 排序 / 精排:hybrid 先按融合分预排序,再对前 rerank_candidates 个候选做精排, + # 剩余候选按融合分排在精排结果之后;rerank=False 时跳过精排直接按融合分排序。 if request.mode == SearchMode.hybrid: - candidates = [ - RankedCandidate(block_id=h.block_id, score=candidate_scores[h.block_id], text=h.content) - for h in filtered - ] - ranked = await self.reranker.rerank(request.query, candidates) - ordered = [(c.block_id, c.score) for c in ranked] + pre_sorted = sorted(filtered, key=lambda h: -candidate_scores[h.block_id]) + if request.rerank: + limit = request.rerank_candidates + pool = pre_sorted if limit is None else pre_sorted[:limit] + rest = [] if limit is None else pre_sorted[limit:] + candidates = [ + RankedCandidate(block_id=h.block_id, score=candidate_scores[h.block_id], text=h.content) + for h in pool + ] + ranked = await self.reranker.rerank(request.query, candidates) + ordered = [(c.block_id, c.score) for c in ranked] + ordered += [(h.block_id, candidate_scores[h.block_id]) for h in rest] + else: + ordered = [(h.block_id, candidate_scores[h.block_id]) for h in pre_sorted] else: ordered = sorted( ((h.block_id, candidate_scores[h.block_id]) for h in filtered), @@ -112,6 +121,8 @@ class RetrievalEngine: ) ordered = normalize_scores(ordered) + # score_threshold:归一化后过滤低分结果(默认 0 不过滤) + ordered = [(bid, score) for bid, score in ordered if score >= request.score_threshold] # 5. 分页:total = 过滤后候选集大小。fts 已取全量(≤FTS_FETCH_LIMIT)故为真实命中数; # vector/hybrid 为 KNN 候选集,无全局 total。 @@ -154,6 +165,7 @@ class RetrievalEngine: ordered = normalize_scores( [(hit.block_id, -hit.bm25) for hit in fts_hits if hit.block_id in hits] ) + ordered = [(bid, score) for bid, score in ordered if score >= request.score_threshold] items = [self._build_result(hits[block_id], request, score) for block_id, score in ordered] return SearchResponse( query=request.query, diff --git a/backend/app/retrieval/reranker.py b/backend/app/retrieval/reranker.py index 15548d6..3d45f52 100644 --- a/backend/app/retrieval/reranker.py +++ b/backend/app/retrieval/reranker.py @@ -24,6 +24,7 @@ class RerankerProvider(Protocol): """统一 Reranker 接口:输入候选块,输出按相关性重排后的候选块。""" model_id: str + version: str async def rerank(self, query: str, candidates: list[RankedCandidate]) -> list[RankedCandidate]: ... @@ -32,6 +33,7 @@ class LexicalReranker: """轻量精排:query 与块正文的词面重叠度,与归一化后的原始分数加权求和。""" model_id = "lexical-v1" + version = "1" def __init__(self, lexical_weight: float = 0.5) -> None: self.lexical_weight = lexical_weight diff --git a/backend/app/routes.py b/backend/app/routes.py index 2a83799..b012830 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -13,6 +13,7 @@ from app.contracts import ( AgentTraceResponse, ChatRequest, BenchmarkDatasetListResponse, + BenchmarkEventType, BenchmarkKind, BenchmarkReport, BenchmarkRun, @@ -924,7 +925,7 @@ async def cancel_benchmark_run(run_id: str) -> OperationResponse: 404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id} ) return OperationResponse( - status="completed", + status="accepted", resource_id=run_id, message=f"Benchmark run status: {run.status.value}", ) @@ -951,10 +952,27 @@ async def benchmark_events( ) async def stream() -> AsyncIterator[str]: + # 先订阅(保证订阅之后产生的事件也能收到),再回放历史事件,最后实时输出新事件 + queue = benchmark_service.subscribe(run_id) + last_sequence = after_sequence for event in benchmark_service.get_events(run_id): if event.sequence <= after_sequence: continue yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence) + last_sequence = event.sequence + if queue is None: + return + try: + while True: + event = await queue.get() + if event.sequence <= last_sequence: + continue + yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence) + last_sequence = event.sequence + if event.event in (BenchmarkEventType.run_completed, BenchmarkEventType.run_failed): + break + finally: + benchmark_service.unsubscribe(run_id, queue) return StreamingResponse(stream(), media_type="text/event-stream") diff --git a/backend/tests/test_benchmark.py b/backend/tests/test_benchmark.py index 8f1e9af..72b3507 100644 --- a/backend/tests/test_benchmark.py +++ b/backend/tests/test_benchmark.py @@ -2,6 +2,9 @@ 沿用 conftest 的隔离机制:APP_DATA_DIR / DB / Vault 都指向临时目录,benchmark 数据集也落在临时目录(settings.benchmark_datasets_path),不读写真实数据。 + +运行采用「创建即 queued + 后台 Task 执行」的异步模型,测试通过 _run 在同一事件循环内 +创建并等待后台任务结束,得到终态 BenchmarkRun 后再断言。 """ from __future__ import annotations @@ -10,6 +13,7 @@ import asyncio import json import pytest +from pydantic import ValidationError from app.benchmarks import datasets, metrics as m, service from app.config import get_settings @@ -32,6 +36,25 @@ def _write_dataset(dataset_id: str, cases: list[dict], *, kind: str = "rag") -> ) +def _write_raw(dataset_id: str, raw: dict) -> None: + directory = get_settings().benchmark_datasets_path + directory.mkdir(parents=True, exist_ok=True) + (directory / f"{dataset_id}.json").write_text( + json.dumps(raw, ensure_ascii=False), encoding="utf-8" + ) + + +def _run(request: RAGRunRequest): + """创建运行并在同一事件循环内等待后台任务结束,返回终态 BenchmarkRun。""" + from app.contracts import BenchmarkRun + + async def _execute() -> BenchmarkRun: + run = await service.create_rag_run(request) + return await service.wait_for_run(run.run_id) + + return asyncio.run(_execute()) + + # --------------------------------------------------------------------------- # # 指标纯函数 # --------------------------------------------------------------------------- # @@ -44,6 +67,12 @@ def test_hit_at_k_and_recall() -> None: assert m.recall_at_k(retrieved, expected, 5) == 0.5 # 只召回 b +def test_recall_at_k_dedups_duplicate_notes() -> None: + # 同一 Note 经多个 Block 重复出现,去重后 Recall 不应超过 1 + assert m.recall_at_k(["note-a", "note-a"], {"note-a"}, 2) == 1.0 + assert m.recall_at_k(["note-a", "note-a", "note-b"], {"note-a"}, 3) == 1.0 + + def test_reciprocal_rank_and_citation_hit() -> None: assert m.reciprocal_rank(["x", "a", "b"], {"b"}) == 1 / 3 assert m.reciprocal_rank(["x"], {"b"}) == 0.0 @@ -86,6 +115,41 @@ def test_dataset_kind_mismatch_is_invalid() -> None: assert exc.value.code == "BENCHMARK_DATASET_INVALID" +def test_citation_required_requires_expected_block_ids() -> None: + # citation_required=true 却没有 expected_block_ids,无法计算 Citation Hit Rate,应拒绝 + _write_dataset( + "cit-req-v1", + [{"case_id": "x", "query": "q", "expected_note_ids": ["n"], "citation_required": True}], + ) + with pytest.raises(ApiError) as exc: + datasets.load_dataset("cit-req-v1", BenchmarkKind.rag) + assert exc.value.code == "BENCHMARK_DATASET_INVALID" + + +def test_list_datasets_skips_corrupted_structure() -> None: + # 合法 JSON 但字段结构错误(cases: 42),列表接口应隔离该文件而非整体 500 + _write_raw("bad-structure", {"dataset_id": "bad-structure", "kind": "rag", "cases": 42}) + _write_dataset("good-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}]) + + infos = datasets.list_datasets(BenchmarkKind.rag) + ids = {info.dataset_id for info in infos} + assert "good-v1" in ids + assert "bad-structure" not in ids + + +# --------------------------------------------------------------------------- # +# 请求校验(空 / 重复 modes) +# --------------------------------------------------------------------------- # +def test_empty_modes_rejected() -> None: + with pytest.raises(ValidationError): + RAGRunRequest(dataset_id="x", modes=[]) + + +def test_duplicate_modes_rejected() -> None: + with pytest.raises(ValidationError): + RAGRunRequest(dataset_id="x", modes=[SearchMode.fts, SearchMode.fts]) + + # --------------------------------------------------------------------------- # # RAG Benchmark 端到端 # --------------------------------------------------------------------------- # @@ -115,9 +179,7 @@ def test_rag_benchmark_end_to_end() -> None: _, _, case = _single_note_case() _write_dataset("e2e-v1", [case]) - run = asyncio.run( - service.create_rag_run(RAGRunRequest(dataset_id="e2e-v1", modes=[SearchMode.fts])) - ) + run = _run(RAGRunRequest(dataset_id="e2e-v1", modes=[SearchMode.fts])) assert run.status.value == "completed" assert run.dataset_hash.startswith("sha256:") @@ -136,7 +198,7 @@ def test_rag_benchmark_all_modes_produce_metrics() -> None: _, _, case = _single_note_case() _write_dataset("e2e-modes-v1", [case]) - run = asyncio.run(service.create_rag_run(RAGRunRequest(dataset_id="e2e-modes-v1"))) + run = _run(RAGRunRequest(dataset_id="e2e-modes-v1")) assert run.status.value == "completed" for mode in ("fts", "vector", "hybrid"): @@ -145,11 +207,25 @@ def test_rag_benchmark_all_modes_produce_metrics() -> None: assert 0.0 <= run.metrics[mode][key] <= 1.0 +def test_config_snapshot_records_index_and_models() -> None: + _, _, case = _single_note_case() + _write_dataset("snapshot-v1", [case]) + + run = _run(RAGRunRequest(dataset_id="snapshot-v1", modes=[SearchMode.fts])) + + snapshot = run.config_snapshot + assert snapshot["index_meta"] is not None + assert snapshot["embedding"]["version"] + assert snapshot["embedding"]["dim"] + assert snapshot["reranker"]["version"] + assert snapshot["retrieval"]["rrf_k"] == 60 + + def test_benchmark_report_and_events() -> None: _, _, case = _single_note_case() _write_dataset("report-v1", [case]) - run = asyncio.run(service.create_rag_run(RAGRunRequest(dataset_id="report-v1", modes=[SearchMode.fts]))) + run = _run(RAGRunRequest(dataset_id="report-v1", modes=[SearchMode.fts])) report = service.get_report(run.run_id) events = service.get_events(run.run_id) @@ -168,9 +244,48 @@ def test_cancel_completed_run_keeps_status() -> None: _, _, case = _single_note_case() _write_dataset("cancel-v1", [case]) - run = asyncio.run(service.create_rag_run(RAGRunRequest(dataset_id="cancel-v1", modes=[SearchMode.fts]))) + run = _run(RAGRunRequest(dataset_id="cancel-v1", modes=[SearchMode.fts])) + assert run.status.value == "completed" + cancelled = service.cancel_run(run.run_id) - assert cancelled.status.value == "completed" # 同步运行已结束,不再变 cancelled + assert cancelled.status.value == "completed" # 已结束,不再变 cancelled + + +def test_cancel_queued_run_marks_cancelled() -> None: + _, _, case = _single_note_case() + _write_dataset("cancel-queued-v1", [case]) + + async def _scenario(): + run = await service.create_rag_run( + RAGRunRequest(dataset_id="cancel-queued-v1", modes=[SearchMode.fts]) + ) + service.cancel_run(run.run_id) + return await service.wait_for_run(run.run_id) + + run = asyncio.run(_scenario()) + assert run.status.value == "cancelled" + + +# --------------------------------------------------------------------------- # +# 指标聚合:Citation Hit Rate 只统计 citation_required 样本 +# --------------------------------------------------------------------------- # +def test_citation_hit_rate_only_counts_citation_required() -> None: + from app.benchmarks import rag as rag_module + from app.contracts import RAGCaseResult + + cases = [ + RAGCaseResult( + case_id="a", mode=SearchMode.fts, repeat=0, latency_ms=1.0, + citation_hit=True, citation_applicable=True, + ), + RAGCaseResult( + case_id="b", mode=SearchMode.fts, repeat=0, latency_ms=1.0, + citation_hit=False, citation_applicable=False, + ), + ] + metrics = rag_module._aggregate(cases, SearchMode.fts) + # 只有 citation_applicable(citation_required=true)的样本计入分母 + assert metrics.citation_hit_rate == 1.0 # --------------------------------------------------------------------------- # @@ -182,12 +297,17 @@ def test_benchmark_routes_wired() -> None: _, _, case = _single_note_case() _write_dataset("route-v1", [case]) - listed = asyncio.run(routes.list_benchmark_datasets(BenchmarkKind.rag)) - assert any(item.dataset_id == "route-v1" for item in listed.items) + async def _scenario(): + listed = await routes.list_benchmark_datasets(BenchmarkKind.rag) + assert any(item.dataset_id == "route-v1" for item in listed.items) - run = asyncio.run( - routes.create_rag_benchmark(RAGRunRequest(dataset_id="route-v1", modes=[SearchMode.fts])) - ) + run = await routes.create_rag_benchmark( + RAGRunRequest(dataset_id="route-v1", modes=[SearchMode.fts]) + ) + assert run.status.value == "queued" + return await service.wait_for_run(run.run_id) + + run = asyncio.run(_scenario()) assert run.status.value == "completed" got = asyncio.run(routes.get_benchmark_run(run.run_id))