perf(retrieval): one embedding per query per plugin request (milestone 456 step 1, #4903)
CI & Build / Plugin hooks (push) Successful in 17s
CI & Build / Python lint (push) Successful in 2s
CI & Build / TypeScript typecheck (push) Successful in 55s
CI & Build / integration (push) Successful in 1m25s
CI & Build / Python tests (push) Successful in 2m6s
CI & Build / Build & push image (push) Successful in 1m25s

A single plugin request fans one query out to several ranked arms. The
operator message is searched by auto_inject, its reuse and lesson slots,
prompt_rule, the preference slot and rule_via_lesson, and each arm called
get_embedding on its own: up to six model calls for one vector.

- embeddings.query_embedding_memo(): a request-scoped ContextVar memo that
  get_embedding consults. Outside a scope the default is None, so every
  other caller is unchanged. A failed embedding is not remembered.
- memoized_query_embeddings decorates /retrieve, /tool-rules and
  /prior-art.
- get_embeddings (document chunks) is untouched.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
2026-10-05 07:55:06 -04:00
co-authored by Claude Opus 5.5
parent d2ac7bf220
commit be62abf142
3 changed files with 162 additions and 1 deletions
+99
View File
@@ -72,3 +72,102 @@ async def test_get_embedding_propagates_model_load_failures():
):
with pytest.raises(RuntimeError, match="model load failed"):
await get_embedding("x")
# ─── query_embedding_memo (milestone 456 step 1) ─────────────────────────────
def _counting_embeddings():
"""A stand-in for get_embeddings that counts model calls and returns a
vector derived from the text, so a wrong memo key returns a wrong vector."""
calls: list[list[str]] = []
async def fake(texts):
calls.append(list(texts))
return [[float(len(t)), 1.0] for t in texts]
return fake, calls
@pytest.mark.asyncio
async def test_memo_embeds_a_repeated_query_once_per_scope():
from scribe.services import embeddings as emb
fake, calls = _counting_embeddings()
with patch.object(emb, "get_embeddings", fake):
with emb.query_embedding_memo():
first = await emb.get_embedding("push to dev")
again = await emb.get_embedding("push to dev")
other = await emb.get_embedding("merge to main")
assert first == again == [11.0, 1.0]
assert other == [13.0, 1.0]
assert calls == [["push to dev"], ["merge to main"]]
@pytest.mark.asyncio
async def test_without_a_scope_every_call_reaches_the_model():
"""The default is None: outside a request scope nothing is remembered, so
every caller other than the plugin routes behaves exactly as before."""
from scribe.services import embeddings as emb
fake, calls = _counting_embeddings()
with patch.object(emb, "get_embeddings", fake):
await emb.get_embedding("q")
await emb.get_embedding("q")
assert calls == [["q"], ["q"]]
@pytest.mark.asyncio
async def test_a_scope_does_not_leak_into_the_next_one():
from scribe.services import embeddings as emb
fake, calls = _counting_embeddings()
with patch.object(emb, "get_embeddings", fake):
with emb.query_embedding_memo():
await emb.get_embedding("q")
with emb.query_embedding_memo():
await emb.get_embedding("q")
await emb.get_embedding("q")
assert calls == [["q"], ["q"], ["q"]]
assert emb._query_memo.get() is None
@pytest.mark.asyncio
async def test_a_failed_embedding_is_not_remembered():
"""The next arm in the request must try again and degrade on its own,
not inherit a failure as if it were a vector."""
from scribe.services import embeddings as emb
attempts = {"n": 0}
async def flaky(texts):
attempts["n"] += 1
if attempts["n"] == 1:
raise RuntimeError("model busy")
return [[0.5, 0.5] for _ in texts]
with patch.object(emb, "get_embeddings", flaky):
with emb.query_embedding_memo():
with pytest.raises(RuntimeError):
await emb.get_embedding("q")
assert await emb.get_embedding("q") == [0.5, 0.5]
assert attempts["n"] == 2
@pytest.mark.asyncio
async def test_the_route_decorator_opens_one_scope_per_call():
from scribe.services import embeddings as emb
fake, calls = _counting_embeddings()
@emb.memoized_query_embeddings
async def handler(q):
await emb.get_embedding(q)
await emb.get_embedding(q)
return emb._query_memo.get() is not None
with patch.object(emb, "get_embeddings", fake):
assert await handler("a") is True
assert await handler("a") is True
assert calls == [["a"], ["a"]]
assert handler.__name__ == "handler"