fix(benchmark): 记录实际Embedding空间与逐样本回退信息
This commit is contained in:
@@ -24,6 +24,7 @@ from app.contracts import (
|
||||
SearchRequest,
|
||||
)
|
||||
from app.retrieval.engine import engine
|
||||
from app.retrieval.provenance import capture_embedding
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -90,8 +91,10 @@ async def _evaluate_one(
|
||||
score_threshold=request.retrieval.score_threshold,
|
||||
)
|
||||
start = time.perf_counter()
|
||||
embedding = {}
|
||||
try:
|
||||
response = await engine.search(search_request)
|
||||
with capture_embedding() as embedding:
|
||||
response = await engine.search(search_request)
|
||||
latency_ms = (time.perf_counter() - start) * 1000.0
|
||||
except Exception as exc: # 单个样本失败不中断整个 Benchmark
|
||||
# 详细异常只进日志,公开响应只带项目错误码与安全消息,避免泄露路径/SQL 等敏感信息
|
||||
@@ -100,6 +103,7 @@ async def _evaluate_one(
|
||||
exc_info=exc,
|
||||
)
|
||||
return RAGCaseResult(
|
||||
embedding=embedding,
|
||||
case_id=case.case_id,
|
||||
mode=mode,
|
||||
repeat=repeat,
|
||||
@@ -115,6 +119,7 @@ async def _evaluate_one(
|
||||
k = request.retrieval.top_k
|
||||
|
||||
return RAGCaseResult(
|
||||
embedding=embedding,
|
||||
case_id=case.case_id,
|
||||
mode=mode,
|
||||
repeat=repeat,
|
||||
|
||||
@@ -86,7 +86,8 @@ 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": {
|
||||
"embedding": {"policy": "per_case", "details": "cases[].embedding"},
|
||||
"local_embedding": {
|
||||
"model_id": engine.embedding.model_id,
|
||||
"version": engine.embedding.version,
|
||||
"dim": engine.embedding.dim,
|
||||
|
||||
@@ -1131,6 +1131,7 @@ class BenchmarkEvent(Contract):
|
||||
|
||||
|
||||
class RAGCaseResult(Contract):
|
||||
embedding: dict[str, Any] = Field(default_factory=dict)
|
||||
case_id: str
|
||||
mode: SearchMode
|
||||
repeat: int
|
||||
|
||||
@@ -24,6 +24,7 @@ from app.providers.base import ProviderError
|
||||
from app.providers.credentials import CredentialResolver, CredentialStoreError
|
||||
from app.providers.registry import ProviderNotFoundError, ProviderRegistry
|
||||
from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
|
||||
from app.retrieval.provenance import record_embedding
|
||||
|
||||
CAPABILITIES = ("embedding", "transcription", "speaker_matching")
|
||||
HTTP_TYPES = {ProviderType.openai_chat, ProviderType.openai_compatible}
|
||||
@@ -175,7 +176,10 @@ class ModelRoutingService:
|
||||
return data, url
|
||||
|
||||
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
|
||||
if binding and texts:
|
||||
try:
|
||||
|
||||
@@ -23,6 +23,7 @@ from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
|
||||
from app.retrieval.hybrid import normalize_scores, rrf_fuse
|
||||
from app.retrieval.reranker import LexicalReranker, RankedCandidate, RerankerProvider
|
||||
from app.retrieval import routed_vectors
|
||||
from app.retrieval.provenance import record_embedding
|
||||
from app.retrieval.vectorstore import SqliteVecStore, VectorStore
|
||||
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}
|
||||
|
||||
if request.mode in (SearchMode.vector, SearchMode.hybrid):
|
||||
record_embedding(source="unavailable")
|
||||
vec_hits = None
|
||||
if (
|
||||
self._routed_defaults is not None
|
||||
@@ -90,6 +92,8 @@ class RetrievalEngine:
|
||||
if vec_hits is None:
|
||||
query_vec = await self.embedding.embed_query(request.query)
|
||||
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_scores = {v.id: v.score for v in vec_hits}
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -20,6 +20,7 @@ from typing import Protocol
|
||||
|
||||
from app.database.db import connect, transaction
|
||||
from app.retrieval.vectorstore import VectorHit
|
||||
from app.retrieval.provenance import record_embedding
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -80,6 +81,7 @@ async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
|
||||
return None
|
||||
result = await runtime.embed(texts)
|
||||
if result.source != "api":
|
||||
record_embedding(fallback_reason=result.fallback_reason)
|
||||
return None
|
||||
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")
|
||||
@@ -94,6 +96,7 @@ async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
|
||||
)
|
||||
except Exception as exc:
|
||||
# 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__)
|
||||
return None
|
||||
|
||||
@@ -158,6 +161,7 @@ async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
|
||||
batch = await embed_remote([query])
|
||||
if batch is None:
|
||||
return None
|
||||
record_embedding(attempted_space={"model_id": batch.space_id, "dimensions": batch.dimensions})
|
||||
try:
|
||||
conn = connect()
|
||||
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'"
|
||||
).fetchone()
|
||||
if exists is None:
|
||||
record_embedding(fallback_reason="REMOTE_INDEX_MISSING")
|
||||
return None
|
||||
rows = conn.execute(
|
||||
"""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))
|
||||
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:
|
||||
conn.close()
|
||||
except Exception as exc:
|
||||
record_embedding(fallback_reason="REMOTE_INDEX_UNAVAILABLE")
|
||||
logger.debug("Remote vector search unavailable (%s); using local index", type(exc).__name__)
|
||||
return None
|
||||
|
||||
@@ -221,8 +221,9 @@ def test_config_snapshot_records_index_and_models() -> None:
|
||||
|
||||
snapshot = run.config_snapshot
|
||||
assert snapshot["index_meta"] is not None
|
||||
assert snapshot["embedding"]["version"]
|
||||
assert snapshot["embedding"]["dim"]
|
||||
assert snapshot["embedding"]["policy"] == "per_case"
|
||||
assert snapshot["local_embedding"]["version"]
|
||||
assert snapshot["local_embedding"]["dim"]
|
||||
assert snapshot["reranker"]["version"]
|
||||
assert snapshot["retrieval"]["rrf_k"] == 60
|
||||
|
||||
|
||||
@@ -150,6 +150,27 @@ def assert_local(rig, result, texts, reason):
|
||||
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
|
||||
def audio(tmp_path):
|
||||
source, reference = tmp_path / "audio.wav", tmp_path / "reference.wav"
|
||||
|
||||
@@ -66,6 +66,83 @@ async def seed():
|
||||
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"])
|
||||
def test_rebuild_failure_preserves_concurrent_configuration_and_all_indexes(runtime, monkeypatch, failure):
|
||||
from app.container import container
|
||||
|
||||
@@ -1567,3 +1567,7 @@ frontend/src/
|
||||
```
|
||||
|
||||
目录调整应按实际代码规模渐进进行。Router 只做参数接收和错误映射,状态机、第三方 SDK 与文件处理继续放在 Service/Adapter 层。
|
||||
|
||||
### Benchmark Embedding 运行归属(阶段 E 集成修复)
|
||||
|
||||
`config_snapshot.local_embedding` 仅表示本地基线;`config_snapshot.embedding` 为 `{ "policy": "per_case", "details": "cases[].embedding" }`。报告与 CaseCompleted 事件的逐样本 `embedding` 包含实际 source(api/local/not_used/unavailable)、model_id、dimensions,以及可选 version、fallback_reason、requested_route、route_version、attempted_space。requested_route 仅含提供商引用、模型、相对端点和维度,不包含 API Key 或凭据引用。FTS 不使用 Embedding,标记 not_used;远程失败或索引不完整回退时记录实际本地模型及原因。
|
||||
|
||||
@@ -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。记录使用任务局部上下文隔离,并发评测不会相互覆盖。
|
||||
|
||||
## 测试
|
||||
|
||||
|
||||
Reference in New Issue
Block a user