diff --git a/backend/app/benchmarks/rag.py b/backend/app/benchmarks/rag.py index 54cec1d..9d298d0 100644 --- a/backend/app/benchmarks/rag.py +++ b/backend/app/benchmarks/rag.py @@ -7,10 +7,12 @@ from __future__ import annotations +import asyncio import logging import time from collections.abc import Callable +from app import repository from app.benchmarks import metrics as m from app.benchmarks.datasets import RAGDataset from app.contracts import ( @@ -47,10 +49,13 @@ async def run_rag( for mode in request.modes: for case in dataset.cases: + expected_notes = _expected_notes(case) for repeat in range(request.repeat): + # 让出事件循环:使运行中取消、SSE 进度与并发 API 请求能及时得到调度 + await asyncio.sleep(0) 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, expected_notes) results.append(result) done += 1 if on_case is not None: @@ -60,8 +65,19 @@ async def run_rag( return metrics_by_mode, results +def _expected_notes(case: RAGDatasetCase) -> set[str]: + """返回笔记级期望 id;仅标注块 ID 时从块反查所属笔记,避免把标注缺失误判为检索失败。""" + if case.expected_note_ids: + return set(case.expected_note_ids) + return {hit.note_id for hit in repository.get_block_hits(case.expected_block_ids)} + + async def _evaluate_one( - case: RAGDatasetCase, mode: SearchMode, request: RAGRunRequest, repeat: int + case: RAGDatasetCase, + mode: SearchMode, + request: RAGRunRequest, + repeat: int, + expected_notes: set[str], ) -> RAGCaseResult: search_request = SearchRequest( query=case.query, @@ -95,7 +111,6 @@ async def _evaluate_one( retrieved_note_ids = [item.note_id for item in response.items] retrieved_block_ids = [item.block_id for item in response.items] - expected_notes = set(case.expected_note_ids) expected_blocks = set(case.expected_block_ids) k = request.retrieval.top_k diff --git a/backend/app/repository.py b/backend/app/repository.py index 29a082c..b0217dc 100644 --- a/backend/app/repository.py +++ b/backend/app/repository.py @@ -273,11 +273,15 @@ def update_note_location( raise LookupError(note_id) -def fts_search_page( - *, +_FTS_FROM = """ + FROM blocks_fts + JOIN blocks AS b ON b.block_id = blocks_fts.block_id + JOIN notes AS n ON n.note_id = b.note_id + """ + + +def _fts_where( match: str, - limit: int, - offset: int, folders: list[str], note_ids: list[str], tags: list[str], @@ -285,8 +289,11 @@ def fts_search_page( created_to: datetime | None, updated_from: datetime | None, updated_to: datetime | None, -) -> tuple[list[FtsHit], int]: - """执行带元数据过滤的 FTS 精确分页,并返回过滤后的完整命中数。""" +) -> tuple[str, list[object]]: + """构建 FTS 过滤 WHERE 子句(不含 WHERE 关键字),返回 (where_sql, params)。 + + fts_search_page 与 fts_score_bounds 共用,保证计数与取数口径一致。 + """ where = ["blocks_fts MATCH ?"] params: list[object] = [match] @@ -317,22 +324,44 @@ def fts_search_page( where.append(f"julianday({column}) <= julianday(?)") params.append(_iso(upper)) - from_sql = """ - FROM blocks_fts - JOIN blocks AS b ON b.block_id = blocks_fts.block_id - JOIN notes AS n ON n.note_id = b.note_id + return " AND ".join(where), params + + +def fts_search_page( + *, + match: str, + limit: int, + offset: int, + folders: list[str], + note_ids: list[str], + tags: list[str], + created_from: datetime | None, + created_to: datetime | None, + updated_from: datetime | None, + updated_to: datetime | None, + bm25_max: float | None = None, +) -> tuple[list[FtsHit], int]: + """执行带元数据过滤的 FTS 精确分页,并返回过滤后的完整命中数。 + + bm25_max 非空时按 bm25 截止值过滤(用于阈值过滤的精确分页),计数与取数同口径。 """ - where_sql = " AND ".join(where) + where_sql, params = _fts_where( + match, folders, note_ids, tags, + created_from, created_to, updated_from, updated_to, + ) + if bm25_max is not None: + where_sql += " AND bm25(blocks_fts) <= ?" + params.append(bm25_max) conn = connect() try: total = conn.execute( - f"SELECT COUNT(*) {from_sql} WHERE {where_sql}", params + f"SELECT COUNT(*) {_FTS_FROM} WHERE {where_sql}", params ).fetchone()[0] rows = conn.execute( f""" SELECT blocks_fts.block_id, blocks_fts.note_id, bm25(blocks_fts) AS rank - {from_sql} + {_FTS_FROM} WHERE {where_sql} ORDER BY rank LIMIT ? OFFSET ? @@ -348,6 +377,45 @@ def fts_search_page( conn.close() +def fts_score_bounds( + *, + match: str, + folders: list[str], + note_ids: list[str], + tags: list[str], + created_from: datetime | None, + created_to: datetime | None, + updated_from: datetime | None, + updated_to: datetime | None, +) -> tuple[float, float] | None: + """返回 metadata 过滤后的 FTS 命中集里 bm25 的 (min, max),无命中时返回 None。 + + 用于阈值过滤:min-max 归一化是 bm25 的线性函数,据此可把阈值换算为 bm25 截止值。 + """ + where_sql, params = _fts_where( + match, folders, note_ids, tags, + created_from, created_to, updated_from, updated_to, + ) + conn = connect() + try: + # bm25() 不能作为聚合函数参数,也不能用在被聚合的子查询里;改用 ORDER BY 取首尾两行 + lo_row = conn.execute( + f"SELECT bm25(blocks_fts) AS rank {_FTS_FROM} WHERE {where_sql}" + " ORDER BY rank ASC LIMIT 1", + params, + ).fetchone() + if lo_row is None or lo_row["rank"] is None: + return None + hi_row = conn.execute( + f"SELECT bm25(blocks_fts) AS rank {_FTS_FROM} WHERE {where_sql}" + " ORDER BY rank DESC LIMIT 1", + params, + ).fetchone() + return (float(lo_row["rank"]), float(hi_row["rank"])) + finally: + conn.close() + + def get_block_hits(block_ids: list[str]) -> list[BlockHit]: if not block_ids: return [] diff --git a/backend/app/retrieval/engine.py b/backend/app/retrieval/engine.py index 15d4ba6..0688f34 100644 --- a/backend/app/retrieval/engine.py +++ b/backend/app/retrieval/engine.py @@ -29,8 +29,6 @@ from app.textutils import make_snippet, match_query CANDIDATE_POOL = 50 # 分页窗口上限:候选池至少覆盖 offset+limit,但设上限防止超大 offset 撑爆内存 MAX_CANDIDATE_POOL = 200 -# FTS 全量取回上限:统一归一化 + 阈值过滤后再分页,保证阈值语义跨页一致 -FTS_FETCH_LIMIT = 5000 # 带 metadata 过滤时放大召回倍数,缓解「先截断候选池再过滤」造成的漏召回 OVERSCAN_FACTOR = 4 @@ -58,7 +56,7 @@ class RetrievalEngine: # 候选池至少覆盖本次请求的 offset+limit,保证分页能取到目标页;设上限防内存失控 window = min(request.offset + request.limit, MAX_CANDIDATE_POOL) pool_size = max(CANDIDATE_POOL, window) - # 带过滤时放大召回;FTS 则一次性取全量命中(≤FTS_FETCH_LIMIT)避免截断漏召回 + # 带过滤时放大召回,缓解「先截断候选池再过滤」造成的漏召回 recall = min(pool_size * OVERSCAN_FACTOR, MAX_CANDIDATE_POOL) if has_filters else pool_size # 1. 按模式收集候选(FTS 与 Vector 各产出「按相关性降序」的 block_id 列表) @@ -126,7 +124,7 @@ class RetrievalEngine: # 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 走数据库精确分页,total 为真实命中数; # vector/hybrid 为 KNN 候选集,无全局 total。 total = len(ordered) page = ordered[request.offset : request.offset + request.limit] @@ -139,18 +137,18 @@ class RetrievalEngine: ) def _search_fts(self, request: SearchRequest) -> SearchResponse: - """FTS 专用路径:先取全量命中(≤FTS_FETCH_LIMIT),统一归一化 + 阈值过滤后再分页。 + """FTS 专用路径:在数据库侧完成过滤、计数与分页,不取全量后再截断。 - 阈值过滤必须在计数与分页之前完成,否则 score_threshold 只作用于当前页, - 且返回的 total 与 items 数量不一致(如 items 为空但 total 非零)。""" + 阈值过滤时,min-max 归一化是 bm25 的线性函数,据此把 score_threshold 换算为 + bm25 截止值(bm25_max),使过滤、计数与分页口径一致;无阈值时走数据库原生分页, + total 始终为过滤后的真实命中数,不再受固定截断影响。 + """ match = match_query(request.query) if not match: return self._empty(request) - fts_hits, _ = repository.fts_search_page( + bounds = repository.fts_score_bounds( match=match, - limit=FTS_FETCH_LIMIT, - offset=0, folders=request.folders, note_ids=request.note_ids, tags=request.tags, @@ -159,17 +157,47 @@ class RetrievalEngine: updated_from=request.updated_from, updated_to=request.updated_to, ) - if not fts_hits: + if bounds is None: return self._empty(request) - ordered = normalize_scores([(hit.block_id, -hit.bm25) for hit in fts_hits]) - ordered = [(bid, score) for bid, score in ordered if score >= request.score_threshold] - total = len(ordered) - page = ordered[request.offset : request.offset + request.limit] - hits = {h.block_id: h for h in repository.get_block_hits([bid for bid, _ in page])} + lo, hi = bounds + bm25_max: float | None = None + if request.score_threshold > 0: + # norm = (hi - bm25) / (hi - lo);norm >= threshold ⟺ bm25 <= hi - threshold*(hi - lo) + bm25_max = hi - request.score_threshold * (hi - lo) + + fts_hits, total = repository.fts_search_page( + match=match, + limit=request.limit, + offset=request.offset, + folders=request.folders, + note_ids=request.note_ids, + tags=request.tags, + created_from=request.created_from, + created_to=request.created_to, + updated_from=request.updated_from, + updated_to=request.updated_to, + bm25_max=bm25_max, + ) + if not fts_hits: + # 本页无结果:offset 越过末页时 total 仍为真实命中数(>0),需保留而非归零 + return SearchResponse( + query=request.query, + mode=request.mode, + items=[], + page=PageMeta(total=total, limit=request.limit, offset=request.offset), + ) + + # 分数按全局 bm25 上下界归一化(与取全量后 normalize_scores 等价),保证跨页一致 + span = hi - lo + if span == 0: + ordered = [(hit.block_id, 1.0) for hit in fts_hits] + else: + ordered = [(hit.block_id, round((hi - hit.bm25) / span, 6)) for hit in fts_hits] + hits = {h.block_id: h for h in repository.get_block_hits([bid for bid, _ in ordered])} items = [ self._build_result(hits[block_id], request, score) - for block_id, score in page + for block_id, score in ordered if block_id in hits ] return SearchResponse( diff --git a/backend/app/routes.py b/backend/app/routes.py index efce047..0c14ce2 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -961,30 +961,38 @@ async def benchmark_events( async def stream() -> AsyncIterator[str]: # 先订阅(保证订阅之后产生的事件也能收到),再回放历史事件,最后实时输出新事件 + terminal = ( + BenchmarkEventType.run_completed, + BenchmarkEventType.run_failed, + BenchmarkEventType.run_cancelled, + ) queue = benchmark_service.subscribe(run_id) - last_sequence = cursor - for event in benchmark_service.get_events(run_id): - if event.sequence <= cursor: - 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: + last_sequence = cursor + # 回放按订阅时刻的快照长度遍历,避免列表在回放期间被追加;终止事件同样要结束流, + # 防止回放完成后进入实时队列却因序号去重跳过同一终止事件而永久等待。 + history = benchmark_service.get_events(run_id) + for index in range(len(history)): + event = history[index] + if event.sequence <= cursor: + continue + yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence) + last_sequence = event.sequence + if event.event in terminal: + return + if queue is None: + return 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, - BenchmarkEventType.run_cancelled, - ): - break + if event.event in terminal: + return finally: - benchmark_service.unsubscribe(run_id, queue) + if queue is not None: + 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 b2fea65..0b9e0d2 100644 --- a/backend/tests/test_benchmark.py +++ b/backend/tests/test_benchmark.py @@ -439,3 +439,128 @@ def test_load_dataset_top_level_must_be_object() -> None: with pytest.raises(ApiError) as exc: datasets.load_dataset("array-top", BenchmarkKind.rag) assert exc.value.code == "BENCHMARK_DATASET_INVALID" + + +# --------------------------------------------------------------------------- # +# 审阅回归:运行中取消 / 仅块标注 / SSE 终止事件 +# --------------------------------------------------------------------------- # +def test_cancel_running_benchmark_stops_early() -> None: + """运行中取消应在样本边界及时生效,而非跑完全部样本(审阅 P1)。""" + from app.benchmarks import service + from app.services import note_service + + note = asyncio.run( + note_service.create_note( + title="取消回归", markdown="向量数据库用于存储高维向量。", folder="", tags=["向量"] + ) + ) + cases = [ + { + "case_id": f"c{i}", + "query": "向量数据库", + "expected_note_ids": [note.note_id], + "expected_block_ids": [note.blocks[0].block_id], + "citation_required": True, + } + for i in range(50) + ] + _write_dataset("cancel-running-v1", cases) + + async def _scenario(): + run = await service.create_rag_run( + RAGRunRequest(dataset_id="cancel-running-v1", modes=[SearchMode.fts]) + ) + + async def _cancel_after_start(): + # 取消通过事件循环调度(独立 Task),而非同步直调,才能复现事件循环饥饿 + while service.get_run(run.run_id).status == BenchmarkStatus.queued: + await asyncio.sleep(0) + service.cancel_run(run.run_id) + + cancel_task = asyncio.create_task(_cancel_after_start()) + finished = await service.wait_for_run(run.run_id) + await cancel_task + return finished + + run = asyncio.run(_scenario()) + assert run.status.value == "cancelled" + completed = sum( + 1 for e in service.get_events(run.run_id) if e.event.value == "CaseCompleted" + ) + assert completed < 50 # 未跑完全部样本,证明取消在样本边界生效 + + +def test_block_only_annotation_resolves_note_and_scores() -> None: + """仅标注 expected_block_ids 的样本应按块反查笔记评分,而非零分(审阅 P2)。""" + from app.services import note_service + + note = asyncio.run( + note_service.create_note( + title="仅块标注", markdown="向量数据库存储高维向量。", folder="", tags=["向量"] + ) + ) + _write_dataset("block-only-v1", [{ + "case_id": "c1", + "query": "向量数据库", + "expected_block_ids": [note.blocks[0].block_id], + "citation_required": False, + }]) + + run = _run(RAGRunRequest(dataset_id="block-only-v1", modes=[SearchMode.fts])) + + assert run.status.value == "completed" + fts = run.metrics["fts"] + assert fts["hit_at_1"] == 1.0 + assert fts["recall_at_k"] == 1.0 + assert fts["mrr"] == 1.0 + + +def test_sse_stream_ends_on_terminal_event_in_replay() -> None: + """历史回放期间遇到终止事件时流应立即结束,而非进入实时队列永久等待(审阅 P2)。""" + from app import routes + from app.benchmarks import service + from app.contracts import BenchmarkEvent, BenchmarkEventType + + run_id = "benchmark_sse_replay" + now = service._now() + # 模拟「回放期间运行完成」:run 仍为 running(subscribe 返回非空队列), + # 但历史事件里已含 RunCompleted 终止事件。 + service._runs[run_id] = BenchmarkRun( + run_id=run_id, + kind=BenchmarkKind.rag, + dataset_id="d", + dataset_hash="sha256:x", + status=BenchmarkStatus.running, + created_at=now, + ) + service._events[run_id] = [ + BenchmarkEvent( + event=BenchmarkEventType.run_started, run_id=run_id, sequence=0, + data={}, timestamp=now, + ), + BenchmarkEvent( + event=BenchmarkEventType.run_completed, run_id=run_id, sequence=1, + data={}, timestamp=now, + ), + ] + try: + # 直调路由函数时 FastAPI 不解析 Query/Header 默认值,需显式传 None 覆盖 Header 哨兵 + response = asyncio.run( + routes.benchmark_events(run_id, after_sequence=-1, last_event_id=None) + ) + + async def _collect() -> list[str]: + out: list[str] = [] + async for chunk in response.body_iterator: + out.append(chunk) + return out + + # 加超时防止回归(旧实现会永久挂起) + chunks = asyncio.run(asyncio.wait_for(_collect(), timeout=5)) + finally: + service._forget(run_id) + + events = [ + line for chunk in chunks for line in chunk.splitlines() if line.startswith("event: ") + ] + assert events == ["event: RunStarted", "event: RunCompleted"] diff --git a/backend/tests/test_retrieval.py b/backend/tests/test_retrieval.py index e9d4e62..d8cffe3 100644 --- a/backend/tests/test_retrieval.py +++ b/backend/tests/test_retrieval.py @@ -474,6 +474,45 @@ def test_fts_score_threshold_filters_before_total(vault) -> None: assert none.items == [] +def test_fts_offset_beyond_end_reports_real_total(vault) -> None: + """offset 越过末页时 items 为空,但 total 仍为真实命中数而非归零。""" + from app.retrieval.engine import engine + from app.services import note_service + + asyncio.run( + note_service.create_note(title="越界分页", markdown="检索 检索 检索 检索", folder="", tags=[]) + ) + + resp = asyncio.run( + engine.search(SearchRequest(query="检索", mode=SearchMode.fts, limit=10, offset=100)) + ) + assert resp.page.total >= 1 + assert resp.items == [] + + +def test_fts_not_truncated_at_five_thousand(vault) -> None: + """FTS 结果不再被 5000 条上限截断:>5000 命中时 total 为真实计数,末页仍可访问。""" + from app.retrieval.engine import engine + from app.services import note_service + + markdown = "\n\n".join(f"共同词 q{i}" for i in range(5010)) + asyncio.run( + note_service.create_note(title="五千条分页", markdown=markdown, folder="", tags=[]) + ) + + first = asyncio.run( + engine.search(SearchRequest(query="共同词", mode=SearchMode.fts, limit=10, offset=0)) + ) + assert first.page.total == 5010 + assert len(first.items) == 10 + + last = asyncio.run( + engine.search(SearchRequest(query="共同词", mode=SearchMode.fts, limit=10, offset=5005)) + ) + assert last.page.total == 5010 + assert len(last.items) == 5 + + # --------------------------------------------------------------------------- # # 审阅回归:PATCH tags 语义 / 向量-块一致性 / 过滤漏召回 / rebuild 语义与回滚 # --------------------------------------------------------------------------- #