fix(backend): 落实 PR #9 评审意见
- 检索调优参数(rrf_k/rerank/rerank_candidates/score_threshold)透传到引擎实际执行 - Recall 去重,避免同一 Note 多 Block 重复导致 Recall 超 1 - RAG 运行改为后台异步执行:创建即 queued + 202,支持取消与 SSE 实时事件 - 数据集元数据校验,坏文件隔离跳过;citation_required 语义修正 - modes 空/重复校验;配置快照记录模型版本与索引元信息 Co-Authored-By: Claude Code <noreply@anthropic.com>
This commit is contained in:
@@ -11,7 +11,7 @@ import json
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from pydantic import ValidationError
|
from pydantic import BaseModel, Field, ValidationError
|
||||||
|
|
||||||
from app.config import get_settings
|
from app.config import get_settings
|
||||||
from app.contracts import (
|
from app.contracts import (
|
||||||
@@ -34,6 +34,20 @@ class RAGDataset:
|
|||||||
content_hash: str = ""
|
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:
|
def _datasets_dir() -> Path:
|
||||||
return get_settings().benchmark_datasets_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.",
|
f"Dataset case '{parsed.case_id}' must declare expected_note_ids or expected_block_ids.",
|
||||||
{"dataset_id": dataset_id, "case_id": parsed.case_id},
|
{"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)
|
cases.append(parsed)
|
||||||
|
|
||||||
return RAGDataset(
|
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]:
|
def list_datasets(kind: BenchmarkKind) -> list[BenchmarkDatasetInfo]:
|
||||||
"""枚举受控目录下指定 kind 的数据集元信息(不含 Case 内容)。
|
"""枚举受控目录下指定 kind 的数据集元信息(不含 Case 内容)。
|
||||||
|
|
||||||
个别文件损坏时跳过而非整体失败,保证列表接口健壮;损坏细节由 load_dataset 抛出。
|
逐文件用 _DatasetMeta 校验元信息字段结构,单个损坏文件隔离跳过而非整体失败,
|
||||||
|
保证列表接口健壮;损坏细节由 load_dataset 抛出。
|
||||||
"""
|
"""
|
||||||
infos: list[BenchmarkDatasetInfo] = []
|
infos: list[BenchmarkDatasetInfo] = []
|
||||||
for path in _dataset_files():
|
for path in _dataset_files():
|
||||||
try:
|
try:
|
||||||
raw, raw_bytes = _read_json(path)
|
raw, raw_bytes = _read_json(path)
|
||||||
except ApiError:
|
meta = _DatasetMeta.model_validate(raw)
|
||||||
|
except (ApiError, ValidationError):
|
||||||
continue
|
continue
|
||||||
if raw.get("kind", kind.value) != kind.value:
|
if meta.kind not in ("", kind.value):
|
||||||
continue
|
continue
|
||||||
infos.append(
|
infos.append(
|
||||||
BenchmarkDatasetInfo(
|
BenchmarkDatasetInfo(
|
||||||
dataset_id=raw.get("dataset_id", path.stem),
|
dataset_id=meta.dataset_id,
|
||||||
kind=kind,
|
kind=kind,
|
||||||
version=str(raw.get("version", "")),
|
version=meta.version,
|
||||||
description=str(raw.get("description", "")),
|
description=meta.description,
|
||||||
case_count=len(raw.get("cases", [])),
|
case_count=len(meta.cases),
|
||||||
content_hash=_content_hash(raw_bytes),
|
content_hash=_content_hash(raw_bytes),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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:
|
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:
|
if not expected:
|
||||||
return 0.0
|
return 0.0
|
||||||
hits = sum(1 for item in retrieved[:k] if item in expected)
|
return len(set(retrieved[:k]) & expected) / len(expected)
|
||||||
return hits / len(expected)
|
|
||||||
|
|
||||||
|
|
||||||
def reciprocal_rank(retrieved: list[str], expected: set[str]) -> float:
|
def reciprocal_rank(retrieved: list[str], expected: set[str]) -> float:
|
||||||
|
|||||||
@@ -23,14 +23,20 @@ from app.contracts import (
|
|||||||
from app.retrieval.engine import engine
|
from app.retrieval.engine import engine
|
||||||
|
|
||||||
|
|
||||||
|
class BenchmarkCancelled(Exception):
|
||||||
|
"""运行在 Case 之间被取消时抛出,用于中断后台执行并标记 cancelled。"""
|
||||||
|
|
||||||
|
|
||||||
async def run_rag(
|
async def run_rag(
|
||||||
dataset: RAGDataset,
|
dataset: RAGDataset,
|
||||||
request: RAGRunRequest,
|
request: RAGRunRequest,
|
||||||
on_case: Callable[[RAGCaseResult, int, int], None] | None = None,
|
on_case: Callable[[RAGCaseResult, int, int], None] | None = None,
|
||||||
|
should_cancel: Callable[[], bool] | None = None,
|
||||||
) -> tuple[dict[str, RAGMetrics], list[RAGCaseResult]]:
|
) -> tuple[dict[str, RAGMetrics], list[RAGCaseResult]]:
|
||||||
"""执行 RAG Benchmark,返回 (按 mode 聚合的指标, 全部逐样本结果)。
|
"""执行 RAG Benchmark,返回 (按 mode 聚合的指标, 全部逐样本结果)。
|
||||||
|
|
||||||
on_case 在每个样本求值完成后回调 (result, done, total),供上层更新进度与事件。
|
on_case 在每个样本求值完成后回调 (result, done, total),供上层更新进度与事件。
|
||||||
|
should_cancel 在每个样本开始前被检查;返回 True 时抛出 BenchmarkCancelled 中断运行。
|
||||||
"""
|
"""
|
||||||
total = len(request.modes) * len(dataset.cases) * request.repeat
|
total = len(request.modes) * len(dataset.cases) * request.repeat
|
||||||
done = 0
|
done = 0
|
||||||
@@ -39,6 +45,8 @@ async def run_rag(
|
|||||||
for mode in request.modes:
|
for mode in request.modes:
|
||||||
for case in dataset.cases:
|
for case in dataset.cases:
|
||||||
for repeat in range(request.repeat):
|
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)
|
result = await _evaluate_one(case, mode, request, repeat)
|
||||||
results.append(result)
|
results.append(result)
|
||||||
done += 1
|
done += 1
|
||||||
@@ -57,6 +65,10 @@ async def _evaluate_one(
|
|||||||
mode=mode,
|
mode=mode,
|
||||||
limit=request.retrieval.top_k,
|
limit=request.retrieval.top_k,
|
||||||
include_snippet=False,
|
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()
|
start = time.perf_counter()
|
||||||
try:
|
try:
|
||||||
@@ -89,7 +101,7 @@ async def _evaluate_one(
|
|||||||
recall=m.recall_at_k(retrieved_note_ids, expected_notes, k),
|
recall=m.recall_at_k(retrieved_note_ids, expected_notes, k),
|
||||||
reciprocal_rank=m.reciprocal_rank(retrieved_note_ids, expected_notes),
|
reciprocal_rank=m.reciprocal_rank(retrieved_note_ids, expected_notes),
|
||||||
citation_hit=m.citation_hit(retrieved_block_ids, expected_blocks),
|
citation_hit=m.citation_hit(retrieved_block_ids, expected_blocks),
|
||||||
citation_applicable=bool(case.expected_block_ids),
|
citation_applicable=case.citation_required,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,19 +1,22 @@
|
|||||||
"""Benchmark 服务:运行注册表、配置快照与报告组装。
|
"""Benchmark 服务:运行注册表、配置快照与报告组装。
|
||||||
|
|
||||||
MVP 阶段运行是同步的(与 index_service 一致):POST 创建后立即执行完并返回
|
RAG Benchmark 采用「创建即返回 queued、后台 Task 异步执行」的模式(与 index_service
|
||||||
completed 的 BenchmarkRun。运行记录、事件与报告暂存内存(_runs/_events/_reports),
|
的 rebuild 一致):POST 创建后立即返回 202 queued 的 BenchmarkRun,由受管 asyncio.Task
|
||||||
不持久化到 SQLite;后续接入异步任务队列时再落库。
|
在后台逐 Case 求值,进度与事件实时写入内存注册表,供 SSE 订阅。运行记录、事件与报告
|
||||||
|
暂存内存(_runs/_events/_reports),不持久化到 SQLite;后续接入异步任务队列时再落库。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import sys
|
import sys
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
|
from app import repository
|
||||||
from app.benchmarks import datasets
|
from app.benchmarks import datasets
|
||||||
from app.benchmarks.datasets import RAGDataset
|
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.config import get_settings
|
||||||
from app.contracts import (
|
from app.contracts import (
|
||||||
BenchmarkEvent,
|
BenchmarkEvent,
|
||||||
@@ -32,6 +35,9 @@ from app.retrieval.engine import engine
|
|||||||
_runs: dict[str, BenchmarkRun] = {}
|
_runs: dict[str, BenchmarkRun] = {}
|
||||||
_events: dict[str, list[BenchmarkEvent]] = {}
|
_events: dict[str, list[BenchmarkEvent]] = {}
|
||||||
_reports: dict[str, BenchmarkReport] = {}
|
_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
|
MAX_RUNS = 100
|
||||||
|
|
||||||
|
|
||||||
@@ -46,6 +52,9 @@ def _remember(run: BenchmarkRun) -> None:
|
|||||||
_runs.pop(oldest, None)
|
_runs.pop(oldest, None)
|
||||||
_events.pop(oldest, None)
|
_events.pop(oldest, None)
|
||||||
_reports.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:
|
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],
|
"modes": [m.value for m in request.modes],
|
||||||
"retrieval": request.retrieval.model_dump(),
|
"retrieval": request.retrieval.model_dump(),
|
||||||
"repeat": request.repeat,
|
"repeat": request.repeat,
|
||||||
"embedding": {"model_id": engine.embedding.model_id, "dim": engine.embedding.dim},
|
"embedding": {
|
||||||
"reranker": {"model_id": engine.reranker.model_id},
|
"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},
|
"app": {"version": settings.version, "environment": settings.environment},
|
||||||
"python": sys.version.split()[0],
|
"python": sys.version.split()[0],
|
||||||
"metadata": request.metadata,
|
"metadata": request.metadata,
|
||||||
@@ -67,7 +84,7 @@ def _config_snapshot(request: RAGRunRequest, dataset: RAGDataset) -> dict:
|
|||||||
|
|
||||||
|
|
||||||
async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun:
|
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)
|
dataset = datasets.load_dataset(request.dataset_id, BenchmarkKind.rag)
|
||||||
run_id = "benchmark_" + uuid4().hex[:12]
|
run_id = "benchmark_" + uuid4().hex[:12]
|
||||||
snapshot = _config_snapshot(request, dataset)
|
snapshot = _config_snapshot(request, dataset)
|
||||||
@@ -77,24 +94,41 @@ async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun:
|
|||||||
kind=BenchmarkKind.rag,
|
kind=BenchmarkKind.rag,
|
||||||
dataset_id=dataset.dataset_id,
|
dataset_id=dataset.dataset_id,
|
||||||
dataset_hash=dataset.content_hash,
|
dataset_hash=dataset.content_hash,
|
||||||
status=BenchmarkStatus.running,
|
status=BenchmarkStatus.queued,
|
||||||
progress=0.0,
|
progress=0.0,
|
||||||
config_snapshot=snapshot,
|
config_snapshot=snapshot,
|
||||||
created_at=_now(),
|
created_at=_now(),
|
||||||
started_at=_now(),
|
|
||||||
)
|
)
|
||||||
_remember(run)
|
_remember(run)
|
||||||
_events[run_id] = []
|
_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:
|
def emit(event_type: BenchmarkEventType, data: dict) -> None:
|
||||||
sequence = len(_events[run_id])
|
sequence = len(_events[run_id])
|
||||||
_events[run_id].append(
|
event = BenchmarkEvent(
|
||||||
BenchmarkEvent(
|
event=event_type, run_id=run_id, sequence=sequence, data=data, timestamp=_now()
|
||||||
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(
|
emit(
|
||||||
BenchmarkEventType.run_started,
|
BenchmarkEventType.run_started,
|
||||||
{"dataset_id": dataset.dataset_id, "modes": [m.value for m in request.modes]},
|
{"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"))
|
emit(BenchmarkEventType.case_completed, result.model_dump(mode="json"))
|
||||||
|
|
||||||
try:
|
try:
|
||||||
metrics_by_mode, results = await run_rag(dataset, request, on_case=on_case)
|
metrics_by_mode, results = await run_rag(
|
||||||
except Exception as exc:
|
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(
|
_runs[run_id] = _runs[run_id].model_copy(
|
||||||
update={
|
update={
|
||||||
"status": BenchmarkStatus.failed,
|
"status": BenchmarkStatus.failed,
|
||||||
@@ -127,7 +184,8 @@ async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun:
|
|||||||
config_snapshot=snapshot,
|
config_snapshot=snapshot,
|
||||||
error=str(exc),
|
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()}
|
metrics = {mode: m.model_dump() for mode, m in metrics_by_mode.items()}
|
||||||
_runs[run_id] = _runs[run_id].model_copy(
|
_runs[run_id] = _runs[run_id].model_copy(
|
||||||
@@ -149,7 +207,7 @@ async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun:
|
|||||||
metrics=metrics,
|
metrics=metrics,
|
||||||
cases=results,
|
cases=results,
|
||||||
)
|
)
|
||||||
return _runs[run_id]
|
finish()
|
||||||
|
|
||||||
|
|
||||||
def list_runs(
|
def list_runs(
|
||||||
@@ -181,13 +239,38 @@ def get_events(run_id: str) -> list[BenchmarkEvent]:
|
|||||||
|
|
||||||
|
|
||||||
def cancel_run(run_id: str) -> BenchmarkRun | None:
|
def cancel_run(run_id: str) -> BenchmarkRun | None:
|
||||||
"""取消运行:同步 MVP 下运行通常已结束,仅对仍在排队/运行的记录置为 cancelled。"""
|
"""取消运行:对 queued/running 设置取消标志,后台 Task 在 Case 边界检查后置为 cancelled。"""
|
||||||
run = _runs.get(run_id)
|
run = _runs.get(run_id)
|
||||||
if run is None:
|
if run is None:
|
||||||
return None
|
return None
|
||||||
if run.status in (BenchmarkStatus.queued, BenchmarkStatus.running):
|
if run.status in (BenchmarkStatus.queued, BenchmarkStatus.running):
|
||||||
run = run.model_copy(
|
_cancel_flags[run_id].set()
|
||||||
update={"status": BenchmarkStatus.cancelled, "completed_at": _now()}
|
|
||||||
)
|
|
||||||
_runs[run_id] = run
|
|
||||||
return run
|
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)
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from datetime import datetime
|
|||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Literal
|
from typing import Any, Literal
|
||||||
|
|
||||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr
|
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
|
||||||
|
|
||||||
|
|
||||||
class Contract(BaseModel):
|
class Contract(BaseModel):
|
||||||
@@ -144,6 +144,12 @@ class SearchRequest(Contract):
|
|||||||
limit: int = Field(default=20, ge=1, le=100)
|
limit: int = Field(default=20, ge=1, le=100)
|
||||||
offset: int = Field(default=0, ge=0)
|
offset: int = Field(default=0, ge=0)
|
||||||
include_snippet: bool = True
|
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):
|
class Citation(Contract):
|
||||||
@@ -678,8 +684,8 @@ class RAGDatasetCase(Contract):
|
|||||||
|
|
||||||
|
|
||||||
class RAGRetrievalConfig(Contract):
|
class RAGRetrievalConfig(Contract):
|
||||||
"""RAG Benchmark 的检索参数。top_k 映射到 SearchRequest.limit,其余参数当前
|
"""RAG Benchmark 的检索参数。top_k 映射到 SearchRequest.limit,
|
||||||
记录进 config_snapshot,由 RetrievalProfile 共享(§9.9)落地后再接入引擎。"""
|
其余参数透传到 SearchRequest,由检索引擎实际执行。"""
|
||||||
|
|
||||||
top_k: int = Field(default=10, ge=1, le=100)
|
top_k: int = Field(default=10, ge=1, le=100)
|
||||||
rrf_k: int = Field(default=60, ge=1)
|
rrf_k: int = Field(default=60, ge=1)
|
||||||
@@ -691,12 +697,20 @@ class RAGRetrievalConfig(Contract):
|
|||||||
class RAGRunRequest(Contract):
|
class RAGRunRequest(Contract):
|
||||||
dataset_id: str = Field(min_length=1)
|
dataset_id: str = Field(min_length=1)
|
||||||
modes: list[SearchMode] = Field(
|
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)
|
retrieval: RAGRetrievalConfig = Field(default_factory=RAGRetrievalConfig)
|
||||||
repeat: int = Field(default=1, ge=1, le=10)
|
repeat: int = Field(default=1, ge=1, le=10)
|
||||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
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):
|
class RAGMetrics(Contract):
|
||||||
hit_at_1: float = 0.0
|
hit_at_1: float = 0.0
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ class EmbeddingProvider(Protocol):
|
|||||||
"""统一 Embedding 接口(与文档一致)。"""
|
"""统一 Embedding 接口(与文档一致)。"""
|
||||||
|
|
||||||
model_id: str
|
model_id: str
|
||||||
|
version: str
|
||||||
dim: int
|
dim: int
|
||||||
|
|
||||||
async def embed_documents(self, texts: list[str]) -> list[list[float]]: ...
|
async def embed_documents(self, texts: list[str]) -> list[list[float]]: ...
|
||||||
@@ -33,6 +34,7 @@ class HashEmbeddingProvider:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
model_id = "hash-v1"
|
model_id = "hash-v1"
|
||||||
|
version = "1"
|
||||||
dim = EMBEDDING_DIM
|
dim = EMBEDDING_DIM
|
||||||
|
|
||||||
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
|||||||
@@ -84,7 +84,7 @@ class RetrievalEngine:
|
|||||||
elif request.mode == SearchMode.vector:
|
elif request.mode == SearchMode.vector:
|
||||||
candidate_scores = vec_scores
|
candidate_scores = vec_scores
|
||||||
else: # hybrid:RRF 融合
|
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:
|
if not candidate_scores:
|
||||||
return self._empty(request)
|
return self._empty(request)
|
||||||
@@ -97,14 +97,23 @@ class RetrievalEngine:
|
|||||||
if not filtered:
|
if not filtered:
|
||||||
return self._empty(request)
|
return self._empty(request)
|
||||||
|
|
||||||
# 4. 排序 / 精排
|
# 4. 排序 / 精排:hybrid 先按融合分预排序,再对前 rerank_candidates 个候选做精排,
|
||||||
|
# 剩余候选按融合分排在精排结果之后;rerank=False 时跳过精排直接按融合分排序。
|
||||||
if request.mode == SearchMode.hybrid:
|
if request.mode == SearchMode.hybrid:
|
||||||
candidates = [
|
pre_sorted = sorted(filtered, key=lambda h: -candidate_scores[h.block_id])
|
||||||
RankedCandidate(block_id=h.block_id, score=candidate_scores[h.block_id], text=h.content)
|
if request.rerank:
|
||||||
for h in filtered
|
limit = request.rerank_candidates
|
||||||
]
|
pool = pre_sorted if limit is None else pre_sorted[:limit]
|
||||||
ranked = await self.reranker.rerank(request.query, candidates)
|
rest = [] if limit is None else pre_sorted[limit:]
|
||||||
ordered = [(c.block_id, c.score) for c in ranked]
|
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:
|
else:
|
||||||
ordered = sorted(
|
ordered = sorted(
|
||||||
((h.block_id, candidate_scores[h.block_id]) for h in filtered),
|
((h.block_id, candidate_scores[h.block_id]) for h in filtered),
|
||||||
@@ -112,6 +121,8 @@ class RetrievalEngine:
|
|||||||
)
|
)
|
||||||
|
|
||||||
ordered = normalize_scores(ordered)
|
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)故为真实命中数;
|
# 5. 分页:total = 过滤后候选集大小。fts 已取全量(≤FTS_FETCH_LIMIT)故为真实命中数;
|
||||||
# vector/hybrid 为 KNN 候选集,无全局 total。
|
# vector/hybrid 为 KNN 候选集,无全局 total。
|
||||||
@@ -154,6 +165,7 @@ class RetrievalEngine:
|
|||||||
ordered = normalize_scores(
|
ordered = normalize_scores(
|
||||||
[(hit.block_id, -hit.bm25) for hit in fts_hits if hit.block_id in hits]
|
[(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]
|
items = [self._build_result(hits[block_id], request, score) for block_id, score in ordered]
|
||||||
return SearchResponse(
|
return SearchResponse(
|
||||||
query=request.query,
|
query=request.query,
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ class RerankerProvider(Protocol):
|
|||||||
"""统一 Reranker 接口:输入候选块,输出按相关性重排后的候选块。"""
|
"""统一 Reranker 接口:输入候选块,输出按相关性重排后的候选块。"""
|
||||||
|
|
||||||
model_id: str
|
model_id: str
|
||||||
|
version: str
|
||||||
|
|
||||||
async def rerank(self, query: str, candidates: list[RankedCandidate]) -> list[RankedCandidate]: ...
|
async def rerank(self, query: str, candidates: list[RankedCandidate]) -> list[RankedCandidate]: ...
|
||||||
|
|
||||||
@@ -32,6 +33,7 @@ class LexicalReranker:
|
|||||||
"""轻量精排:query 与块正文的词面重叠度,与归一化后的原始分数加权求和。"""
|
"""轻量精排:query 与块正文的词面重叠度,与归一化后的原始分数加权求和。"""
|
||||||
|
|
||||||
model_id = "lexical-v1"
|
model_id = "lexical-v1"
|
||||||
|
version = "1"
|
||||||
|
|
||||||
def __init__(self, lexical_weight: float = 0.5) -> None:
|
def __init__(self, lexical_weight: float = 0.5) -> None:
|
||||||
self.lexical_weight = lexical_weight
|
self.lexical_weight = lexical_weight
|
||||||
|
|||||||
+19
-1
@@ -13,6 +13,7 @@ from app.contracts import (
|
|||||||
AgentTraceResponse,
|
AgentTraceResponse,
|
||||||
ChatRequest,
|
ChatRequest,
|
||||||
BenchmarkDatasetListResponse,
|
BenchmarkDatasetListResponse,
|
||||||
|
BenchmarkEventType,
|
||||||
BenchmarkKind,
|
BenchmarkKind,
|
||||||
BenchmarkReport,
|
BenchmarkReport,
|
||||||
BenchmarkRun,
|
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}
|
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id}
|
||||||
)
|
)
|
||||||
return OperationResponse(
|
return OperationResponse(
|
||||||
status="completed",
|
status="accepted",
|
||||||
resource_id=run_id,
|
resource_id=run_id,
|
||||||
message=f"Benchmark run status: {run.status.value}",
|
message=f"Benchmark run status: {run.status.value}",
|
||||||
)
|
)
|
||||||
@@ -951,10 +952,27 @@ async def benchmark_events(
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def stream() -> AsyncIterator[str]:
|
async def stream() -> AsyncIterator[str]:
|
||||||
|
# 先订阅(保证订阅之后产生的事件也能收到),再回放历史事件,最后实时输出新事件
|
||||||
|
queue = benchmark_service.subscribe(run_id)
|
||||||
|
last_sequence = after_sequence
|
||||||
for event in benchmark_service.get_events(run_id):
|
for event in benchmark_service.get_events(run_id):
|
||||||
if event.sequence <= after_sequence:
|
if event.sequence <= after_sequence:
|
||||||
continue
|
continue
|
||||||
yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence)
|
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")
|
return StreamingResponse(stream(), media_type="text/event-stream")
|
||||||
|
|
||||||
|
|||||||
+132
-12
@@ -2,6 +2,9 @@
|
|||||||
|
|
||||||
沿用 conftest 的隔离机制:APP_DATA_DIR / DB / Vault 都指向临时目录,benchmark
|
沿用 conftest 的隔离机制:APP_DATA_DIR / DB / Vault 都指向临时目录,benchmark
|
||||||
数据集也落在临时目录(settings.benchmark_datasets_path),不读写真实数据。
|
数据集也落在临时目录(settings.benchmark_datasets_path),不读写真实数据。
|
||||||
|
|
||||||
|
运行采用「创建即 queued + 后台 Task 执行」的异步模型,测试通过 _run 在同一事件循环内
|
||||||
|
创建并等待后台任务结束,得到终态 BenchmarkRun 后再断言。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -10,6 +13,7 @@ import asyncio
|
|||||||
import json
|
import json
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
from app.benchmarks import datasets, metrics as m, service
|
from app.benchmarks import datasets, metrics as m, service
|
||||||
from app.config import get_settings
|
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
|
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:
|
def test_reciprocal_rank_and_citation_hit() -> None:
|
||||||
assert m.reciprocal_rank(["x", "a", "b"], {"b"}) == 1 / 3
|
assert m.reciprocal_rank(["x", "a", "b"], {"b"}) == 1 / 3
|
||||||
assert m.reciprocal_rank(["x"], {"b"}) == 0.0
|
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"
|
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 端到端
|
# RAG Benchmark 端到端
|
||||||
# --------------------------------------------------------------------------- #
|
# --------------------------------------------------------------------------- #
|
||||||
@@ -115,9 +179,7 @@ def test_rag_benchmark_end_to_end() -> None:
|
|||||||
_, _, case = _single_note_case()
|
_, _, case = _single_note_case()
|
||||||
_write_dataset("e2e-v1", [case])
|
_write_dataset("e2e-v1", [case])
|
||||||
|
|
||||||
run = asyncio.run(
|
run = _run(RAGRunRequest(dataset_id="e2e-v1", modes=[SearchMode.fts]))
|
||||||
service.create_rag_run(RAGRunRequest(dataset_id="e2e-v1", modes=[SearchMode.fts]))
|
|
||||||
)
|
|
||||||
|
|
||||||
assert run.status.value == "completed"
|
assert run.status.value == "completed"
|
||||||
assert run.dataset_hash.startswith("sha256:")
|
assert run.dataset_hash.startswith("sha256:")
|
||||||
@@ -136,7 +198,7 @@ def test_rag_benchmark_all_modes_produce_metrics() -> None:
|
|||||||
_, _, case = _single_note_case()
|
_, _, case = _single_note_case()
|
||||||
_write_dataset("e2e-modes-v1", [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"
|
assert run.status.value == "completed"
|
||||||
|
|
||||||
for mode in ("fts", "vector", "hybrid"):
|
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
|
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:
|
def test_benchmark_report_and_events() -> None:
|
||||||
_, _, case = _single_note_case()
|
_, _, case = _single_note_case()
|
||||||
_write_dataset("report-v1", [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)
|
report = service.get_report(run.run_id)
|
||||||
events = service.get_events(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()
|
_, _, case = _single_note_case()
|
||||||
_write_dataset("cancel-v1", [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)
|
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()
|
_, _, case = _single_note_case()
|
||||||
_write_dataset("route-v1", [case])
|
_write_dataset("route-v1", [case])
|
||||||
|
|
||||||
listed = asyncio.run(routes.list_benchmark_datasets(BenchmarkKind.rag))
|
async def _scenario():
|
||||||
assert any(item.dataset_id == "route-v1" for item in listed.items)
|
listed = await routes.list_benchmark_datasets(BenchmarkKind.rag)
|
||||||
|
assert any(item.dataset_id == "route-v1" for item in listed.items)
|
||||||
|
|
||||||
run = asyncio.run(
|
run = await routes.create_rag_benchmark(
|
||||||
routes.create_rag_benchmark(RAGRunRequest(dataset_id="route-v1", modes=[SearchMode.fts]))
|
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"
|
assert run.status.value == "completed"
|
||||||
|
|
||||||
got = asyncio.run(routes.get_benchmark_run(run.run_id))
|
got = asyncio.run(routes.get_benchmark_run(run.run_id))
|
||||||
|
|||||||
Reference in New Issue
Block a user