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

This commit is contained in:
2026-09-04 07:25:12 +08:00
parent 96e109f95e
commit d32e9c6efa
10 changed files with 150 additions and 6 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