fix(benchmark): 记录实际Embedding空间与逐样本回退信息

This commit is contained in:
2026-09-04 07:25:12 +08:00
parent 5dd5a46aae
commit a75d81a7d9
12 changed files with 157 additions and 7 deletions
+5
View File
@@ -24,6 +24,7 @@ from app.contracts import (
SearchRequest, SearchRequest,
) )
from app.retrieval.engine import engine from app.retrieval.engine import engine
from app.retrieval.provenance import capture_embedding
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -90,7 +91,9 @@ async def _evaluate_one(
score_threshold=request.retrieval.score_threshold, score_threshold=request.retrieval.score_threshold,
) )
start = time.perf_counter() start = time.perf_counter()
embedding = {}
try: try:
with capture_embedding() as embedding:
response = await engine.search(search_request) response = await engine.search(search_request)
latency_ms = (time.perf_counter() - start) * 1000.0 latency_ms = (time.perf_counter() - start) * 1000.0
except Exception as exc: # 单个样本失败不中断整个 Benchmark except Exception as exc: # 单个样本失败不中断整个 Benchmark
@@ -100,6 +103,7 @@ async def _evaluate_one(
exc_info=exc, exc_info=exc,
) )
return RAGCaseResult( return RAGCaseResult(
embedding=embedding,
case_id=case.case_id, case_id=case.case_id,
mode=mode, mode=mode,
repeat=repeat, repeat=repeat,
@@ -115,6 +119,7 @@ async def _evaluate_one(
k = request.retrieval.top_k k = request.retrieval.top_k
return RAGCaseResult( return RAGCaseResult(
embedding=embedding,
case_id=case.case_id, case_id=case.case_id,
mode=mode, mode=mode,
repeat=repeat, repeat=repeat,
+2 -1
View File
@@ -86,7 +86,8 @@ 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": { "embedding": {"policy": "per_case", "details": "cases[].embedding"},
"local_embedding": {
"model_id": engine.embedding.model_id, "model_id": engine.embedding.model_id,
"version": engine.embedding.version, "version": engine.embedding.version,
"dim": engine.embedding.dim, "dim": engine.embedding.dim,
+1
View File
@@ -1131,6 +1131,7 @@ class BenchmarkEvent(Contract):
class RAGCaseResult(Contract): class RAGCaseResult(Contract):
embedding: dict[str, Any] = Field(default_factory=dict)
case_id: str case_id: str
mode: SearchMode mode: SearchMode
repeat: int repeat: int
+5 -1
View File
@@ -24,6 +24,7 @@ from app.providers.base import ProviderError
from app.providers.credentials import CredentialResolver, CredentialStoreError from app.providers.credentials import CredentialResolver, CredentialStoreError
from app.providers.registry import ProviderNotFoundError, ProviderRegistry from app.providers.registry import ProviderNotFoundError, ProviderRegistry
from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
from app.retrieval.provenance import record_embedding
CAPABILITIES = ("embedding", "transcription", "speaker_matching") CAPABILITIES = ("embedding", "transcription", "speaker_matching")
HTTP_TYPES = {ProviderType.openai_chat, ProviderType.openai_compatible} HTTP_TYPES = {ProviderType.openai_chat, ProviderType.openai_compatible}
@@ -175,7 +176,10 @@ class ModelRoutingService:
return data, url return data, url
async def embed(self, texts: list[str]) -> EmbeddingResult: async def embed(self, texts: list[str]) -> EmbeddingResult:
binding = self.configuration().embedding config = self.configuration()
binding = config.embedding
record_embedding(route_version=config.version,
requested_route=binding.model_dump() if binding else None)
reason = None reason = None
if binding and texts: if binding and texts:
try: try:
+4
View File
@@ -23,6 +23,7 @@ from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
from app.retrieval.hybrid import normalize_scores, rrf_fuse from app.retrieval.hybrid import normalize_scores, rrf_fuse
from app.retrieval.reranker import LexicalReranker, RankedCandidate, RerankerProvider from app.retrieval.reranker import LexicalReranker, RankedCandidate, RerankerProvider
from app.retrieval import routed_vectors from app.retrieval import routed_vectors
from app.retrieval.provenance import record_embedding
from app.retrieval.vectorstore import SqliteVecStore, VectorStore from app.retrieval.vectorstore import SqliteVecStore, VectorStore
from app.textutils import make_snippet, match_query from app.textutils import make_snippet, match_query
@@ -80,6 +81,7 @@ class RetrievalEngine:
fts_scores = {h.block_id: -h.bm25 for h in fts_hits} fts_scores = {h.block_id: -h.bm25 for h in fts_hits}
if request.mode in (SearchMode.vector, SearchMode.hybrid): if request.mode in (SearchMode.vector, SearchMode.hybrid):
record_embedding(source="unavailable")
vec_hits = None vec_hits = None
if ( if (
self._routed_defaults is not None self._routed_defaults is not None
@@ -90,6 +92,8 @@ class RetrievalEngine:
if vec_hits is None: if vec_hits is None:
query_vec = await self.embedding.embed_query(request.query) query_vec = await self.embedding.embed_query(request.query)
vec_hits = await self.vector_store.search(query_vec, top_k=recall) vec_hits = await self.vector_store.search(query_vec, top_k=recall)
record_embedding(source="local", model_id=self.embedding.model_id,
dimensions=self.embedding.dim, version=self.embedding.version)
vec_ranked = [v.id for v in vec_hits] vec_ranked = [v.id for v in vec_hits]
vec_scores = {v.id: v.score for v in vec_hits} vec_scores = {v.id: v.score for v in vec_hits}
+21
View File
@@ -0,0 +1,21 @@
"""Task-local observations of the embedding path actually used by a search."""
from contextlib import contextmanager
from contextvars import ContextVar
_observation: ContextVar[dict | None] = ContextVar("embedding_observation", default=None)
@contextmanager
def capture_embedding():
result = {"source": "not_used"}
token = _observation.set(result)
try:
yield result
finally:
_observation.reset(token)
def record_embedding(**fields) -> None:
result = _observation.get()
if result is not None:
result.update(fields)
+10 -1
View File
@@ -20,6 +20,7 @@ from typing import Protocol
from app.database.db import connect, transaction from app.database.db import connect, transaction
from app.retrieval.vectorstore import VectorHit from app.retrieval.vectorstore import VectorHit
from app.retrieval.provenance import record_embedding
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -80,6 +81,7 @@ async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
return None return None
result = await runtime.embed(texts) result = await runtime.embed(texts)
if result.source != "api": if result.source != "api":
record_embedding(fallback_reason=result.fallback_reason)
return None return None
if not isinstance(result.model_id, str) or not result.model_id or result.model_id == "hash-v1": if not isinstance(result.model_id, str) or not result.model_id or result.model_id == "hash-v1":
raise ValueError("API embedding needs a distinct space ID") raise ValueError("API embedding needs a distinct space ID")
@@ -94,6 +96,7 @@ async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
) )
except Exception as exc: except Exception as exc:
# Avoid logging provider exceptions containing credentials or note text. # Avoid logging provider exceptions containing credentials or note text.
record_embedding(fallback_reason="REMOTE_EMBEDDING_UNAVAILABLE")
logger.warning("Remote embedding unavailable (%s); using local index", type(exc).__name__) logger.warning("Remote embedding unavailable (%s); using local index", type(exc).__name__)
return None return None
@@ -158,6 +161,7 @@ async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
batch = await embed_remote([query]) batch = await embed_remote([query])
if batch is None: if batch is None:
return None return None
record_embedding(attempted_space={"model_id": batch.space_id, "dimensions": batch.dimensions})
try: try:
conn = connect() conn = connect()
try: try:
@@ -166,6 +170,7 @@ async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
"SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'routed_block_vectors'" "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'routed_block_vectors'"
).fetchone() ).fetchone()
if exists is None: if exists is None:
record_embedding(fallback_reason="REMOTE_INDEX_MISSING")
return None return None
rows = conn.execute( rows = conn.execute(
"""SELECT b.block_id, r.vector """SELECT b.block_id, r.vector
@@ -184,9 +189,13 @@ async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
score = math.fsum(a * b for a, b in zip(batch.vectors[0], vector)) score = math.fsum(a * b for a, b in zip(batch.vectors[0], vector))
yield VectorHit(id=row["block_id"], score=max(0.0, min(1.0, score))) yield VectorHit(id=row["block_id"], score=max(0.0, min(1.0, score)))
return heapq.nlargest(top_k, hits(), key=lambda hit: hit.score) result = heapq.nlargest(top_k, hits(), key=lambda hit: hit.score)
record_embedding(source="api", model_id=batch.space_id,
dimensions=batch.dimensions, fallback_reason=None)
return result
finally: finally:
conn.close() conn.close()
except Exception as exc: except Exception as exc:
record_embedding(fallback_reason="REMOTE_INDEX_UNAVAILABLE")
logger.debug("Remote vector search unavailable (%s); using local index", type(exc).__name__) logger.debug("Remote vector search unavailable (%s); using local index", type(exc).__name__)
return None return None
+3 -2
View File
@@ -221,8 +221,9 @@ def test_config_snapshot_records_index_and_models() -> None:
snapshot = run.config_snapshot snapshot = run.config_snapshot
assert snapshot["index_meta"] is not None assert snapshot["index_meta"] is not None
assert snapshot["embedding"]["version"] assert snapshot["embedding"]["policy"] == "per_case"
assert snapshot["embedding"]["dim"] assert snapshot["local_embedding"]["version"]
assert snapshot["local_embedding"]["dim"]
assert snapshot["reranker"]["version"] assert snapshot["reranker"]["version"]
assert snapshot["retrieval"]["rrf_k"] == 60 assert snapshot["retrieval"]["rrf_k"] == 60
+21
View File
@@ -150,6 +150,27 @@ def assert_local(rig, result, texts, reason):
assert rig.embedding.calls == [texts] assert rig.embedding.calls == [texts]
def test_embedding_observation_keeps_request_binding_when_config_changes(rig):
from app.retrieval.provenance import capture_embedding
initial = bind(rig, model="original-model")
def handler(request):
assert json.loads(request.content)["model"] == "original-model"
bind(rig, model="next-model")
return response({"data": [{"index": 0, "embedding": [1, 0, 0]}]})
rig.http.handler = handler
with capture_embedding() as observation:
result = run(rig.service.embed(["query"]))
assert result.source == "api"
assert observation["route_version"] == initial.config.version
assert observation["requested_route"]["model"] == "original-model"
assert observation["requested_route"]["provider_id"] == "test-provider"
assert rig.service.configuration().embedding.model == "next-model"
assert rig.credentials.value not in json.dumps(observation)
assert "credential_id" not in json.dumps(observation)
@pytest.fixture @pytest.fixture
def audio(tmp_path): def audio(tmp_path):
source, reference = tmp_path / "audio.wav", tmp_path / "reference.wav" source, reference = tmp_path / "audio.wav", tmp_path / "reference.wav"
+77
View File
@@ -66,6 +66,83 @@ async def seed():
return apple, banana return apple, banana
@pytest.mark.parametrize("outcome", ["api", "api_failure", "missing_space"])
def test_benchmark_reports_actual_embedding_and_fallback(runtime, outcome):
from app.benchmarks import service
from app.contracts import RAGRunRequest
async def scenario():
apple, banana = await seed()
if outcome == "api_failure":
runtime.result_override = SimpleNamespace(source="local", fallback_reason="PROVIDER_TIMEOUT")
elif outcome == "missing_space":
runtime.model_id = "space-without-index"
directory = get_settings().benchmark_datasets_path
directory.mkdir(parents=True, exist_ok=True)
(directory / "routing.json").write_text(json.dumps({
"dataset_id": "routing", "kind": "rag", "version": "1",
"cases": [{"case_id": "query", "query": "apple", "expected_note_ids": [banana.note_id]}],
}), encoding="utf-8")
run = await service.create_rag_run(RAGRunRequest(
dataset_id="routing", modes=[SearchMode.fts, SearchMode.vector],
))
await service.wait_for_run(run.run_id)
report = service.get_report(run.run_id)
assert report.config_snapshot["embedding"]["policy"] == "per_case"
fts, vector = report.cases
assert fts.embedding == {"source": "not_used"}
if outcome == "api":
assert vector.embedding["source"] == "api"
assert vector.embedding["model_id"] == "space-a"
assert vector.embedding["dimensions"] == 3
assert vector.retrieved_note_ids[0] == banana.note_id
else:
assert vector.embedding["source"] == "local"
assert vector.embedding["model_id"] == "hash-v1"
assert vector.embedding["dimensions"] == 128
assert vector.retrieved_note_ids[0] == apple.note_id
if outcome == "api_failure":
assert vector.embedding["fallback_reason"] == "PROVIDER_TIMEOUT"
if outcome == "missing_space":
assert vector.embedding["fallback_reason"] == "REMOTE_INDEX_UNAVAILABLE"
assert vector.embedding["attempted_space"]["model_id"] == "space-without-index"
events = service.get_events(run.run_id)
case_events = [e for e in events if e.event.value == "CaseCompleted"]
assert case_events[-1].data["embedding"] == vector.embedding
asyncio.run(scenario())
def test_embedding_observations_are_isolated_between_concurrent_searches(runtime, monkeypatch):
from app.retrieval.provenance import capture_embedding
async def scenario():
await seed()
original = runtime.embed
async def embed(texts):
await asyncio.sleep(0)
if texts == ["offline"]:
raise RuntimeError("private upstream details")
return await original(texts)
monkeypatch.setattr(runtime, "embed", embed)
async def query(text):
with capture_embedding() as observation:
await engine.search(SearchRequest(query=text, mode=SearchMode.vector))
return observation
remote, local, another = await asyncio.gather(query("apple"), query("offline"), query("apple"))
assert remote["source"] == another["source"] == "api"
assert local["source"] == "local"
assert local["fallback_reason"] == "REMOTE_EMBEDDING_UNAVAILABLE"
assert "fallback_reason" not in remote or remote["fallback_reason"] is None
assert "private upstream" not in json.dumps(local)
asyncio.run(scenario())
@pytest.mark.parametrize("failure", ["cancel", "write"]) @pytest.mark.parametrize("failure", ["cancel", "write"])
def test_rebuild_failure_preserves_concurrent_configuration_and_all_indexes(runtime, monkeypatch, failure): def test_rebuild_failure_preserves_concurrent_configuration_and_all_indexes(runtime, monkeypatch, failure):
from app.container import container from app.container import container
@@ -1567,3 +1567,7 @@ frontend/src/
``` ```
目录调整应按实际代码规模渐进进行。Router 只做参数接收和错误映射,状态机、第三方 SDK 与文件处理继续放在 Service/Adapter 层。 目录调整应按实际代码规模渐进进行。Router 只做参数接收和错误映射,状态机、第三方 SDK 与文件处理继续放在 Service/Adapter 层。
### Benchmark Embedding 运行归属(阶段 E 集成修复)
`config_snapshot.local_embedding` 仅表示本地基线;`config_snapshot.embedding``{ "policy": "per_case", "details": "cases[].embedding" }`。报告与 CaseCompleted 事件的逐样本 `embedding` 包含实际 sourceapi/local/not_used/unavailable)、model_id、dimensions,以及可选 version、fallback_reason、requested_route、route_version、attempted_space。requested_route 仅含提供商引用、模型、相对端点和维度,不包含 API Key 或凭据引用。FTS 不使用 Embedding,标记 not_used;远程失败或索引不完整回退时记录实际本地模型及原因。
+3 -1
View File
@@ -64,7 +64,9 @@ BENCHMARK_CASE_EVALUATION_FAILED
## 配置快照 ## 配置快照
报告与运行记录保存 `config_snapshot`dataset hash/version、modes、retrieval 参数、Embedding model/version/dim、Reranker、索引元数据、App 版本与环境、Python 版本,保证不同实验结果可复现 报告与运行记录保存 `config_snapshot`dataset hash/version、modes、retrieval 参数、Reranker、索引元数据、App 版本与环境、Python 版本`local_embedding` 记录本地基线 model/version/dim`embedding.policy = per_case` 表示实际来源以逐样本结果为准,不能把本地基线当作本次使用的模型
每个 `RAGCaseResult.embedding`(同时出现在报告 cases 和 CaseCompleted SSE 中)记录 `source`api/local/not_used/unavailable)、实际 `model_id` 空间标识、`dimensions`、本地 `version``fallback_reason`。远程路由还记录请求时的 `route_version``requested_route`provider_id/model/endpoint/dimensions,不含凭据)、成功生成查询向量后的 `attempted_space`。FTS 标记 not_used;调用失败而未完成向量检索时标记 unavailable。API 不可用或远程索引缺失时,实际模型仍记录最终使用的本地基线。配置允许在样本间改变,逐样本记录对应实际调用;汇总指标可能包含多种空间,比较实验时需检查 cases。记录使用任务局部上下文隔离,并发评测不会相互覆盖。
## 测试 ## 测试