diff --git a/backend/app/benchmarks/rag.py b/backend/app/benchmarks/rag.py index 9d298d0..e364f46 100644 --- a/backend/app/benchmarks/rag.py +++ b/backend/app/benchmarks/rag.py @@ -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, diff --git a/backend/app/benchmarks/service.py b/backend/app/benchmarks/service.py index 9b04a27..2db1de3 100644 --- a/backend/app/benchmarks/service.py +++ b/backend/app/benchmarks/service.py @@ -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, diff --git a/backend/app/contracts.py b/backend/app/contracts.py index dcba971..22b47e8 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -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 diff --git a/backend/app/providers/routing.py b/backend/app/providers/routing.py index dd44572..ea7dff3 100644 --- a/backend/app/providers/routing.py +++ b/backend/app/providers/routing.py @@ -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: diff --git a/backend/app/retrieval/engine.py b/backend/app/retrieval/engine.py index 18f4320..1125989 100644 --- a/backend/app/retrieval/engine.py +++ b/backend/app/retrieval/engine.py @@ -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} diff --git a/backend/app/retrieval/provenance.py b/backend/app/retrieval/provenance.py new file mode 100644 index 0000000..b11af6d --- /dev/null +++ b/backend/app/retrieval/provenance.py @@ -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) diff --git a/backend/app/retrieval/routed_vectors.py b/backend/app/retrieval/routed_vectors.py index bd74bff..7a5be34 100644 --- a/backend/app/retrieval/routed_vectors.py +++ b/backend/app/retrieval/routed_vectors.py @@ -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 diff --git a/backend/tests/test_benchmark.py b/backend/tests/test_benchmark.py index 0b9e0d2..0087a8e 100644 --- a/backend/tests/test_benchmark.py +++ b/backend/tests/test_benchmark.py @@ -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 diff --git a/backend/tests/test_model_routing.py b/backend/tests/test_model_routing.py index dbcf8e1..10bc496 100644 --- a/backend/tests/test_model_routing.py +++ b/backend/tests/test_model_routing.py @@ -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" diff --git a/backend/tests/test_routed_retrieval.py b/backend/tests/test_routed_retrieval.py index e4f820d..ec30999 100644 --- a/backend/tests/test_routed_retrieval.py +++ b/backend/tests/test_routed_retrieval.py @@ -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 diff --git a/docs/contracts/第二阶段接口契约-开发版.md b/docs/contracts/第二阶段接口契约-开发版.md index f64b4b8..dd0da1f 100644 --- a/docs/contracts/第二阶段接口契约-开发版.md +++ b/docs/contracts/第二阶段接口契约-开发版.md @@ -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;远程失败或索引不完整回退时记录实际本地模型及原因。 diff --git a/docs/development/Benchmark开发说明.md b/docs/development/Benchmark开发说明.md index bd9f393..628b5bf 100644 --- a/docs/development/Benchmark开发说明.md +++ b/docs/development/Benchmark开发说明.md @@ -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。记录使用任务局部上下文隔离,并发评测不会相互覆盖。 ## 测试