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
+4
View File
@@ -18,6 +18,7 @@ from scribe.services import repo_bindings as repo_bindings_svc
from scribe.services import report_check as report_check_svc
from scribe.services import shape_check as shape_check_svc
from scribe.services import task_claims as task_claims_svc
from scribe.services.embeddings import memoized_query_embeddings
from scribe.services.settings import get_admin_setting, set_setting
plugin_bp = Blueprint("plugin", __name__, url_prefix="/api/plugin")
@@ -88,6 +89,7 @@ async def session_context():
@plugin_bp.get("/retrieve")
@login_required
@memoized_query_embeddings
async def autoinject_retrieve():
"""Title-first knowledge auto-inject for the plugin's UserPromptSubmit hook.
@@ -163,6 +165,7 @@ async def autoinject_retrieve():
@plugin_bp.get("/tool-rules")
@login_required
@memoized_query_embeddings
async def pre_tool_rules():
"""Standing rules for the plugin's PreToolUse hook on ACTIONS (#3476).
@@ -240,6 +243,7 @@ async def pre_tool_rules():
@plugin_bp.get("/prior-art")
@login_required
@memoized_query_embeddings
async def write_path_prior_art():
"""Prior-art hint for the plugin's PreToolUse hook on Write/Edit.
+59 -1
View File
@@ -10,11 +10,14 @@ volume so subsequent boots are instant).
"""
import asyncio
import functools
import logging
import math
import os
from collections.abc import Sequence
from contextlib import contextmanager
from contextvars import ContextVar
from typing import TYPE_CHECKING
@@ -77,13 +80,68 @@ async def _get_model():
return _model
# ONE EMBEDDING PER QUERY PER REQUEST (milestone 456 step 1).
#
# A single plugin request runs several ranked arms over the SAME string: the
# operator's message is searched by auto_inject, its reuse and lesson slots,
# prompt_rule, the preference slot and rule_via_lesson — each of which called
# `get_embedding` on its own, so one prompt cost up to six model calls for one
# vector. The arms are not wrong to search separately; they were wrong to pay
# for the vector separately.
#
# A memo, not a cache: it lives for one request and is opened explicitly by
# `query_embedding_memo()`. Outside a scope the default is None and nothing is
# remembered, so every other caller behaves exactly as before. Keyed on the
# exact text — a query that differs by a character is a different query.
#
# Only `get_embedding` (a QUERY) consults it. `get_embeddings` embeds a
# document's chunks on the write path, where a repeat inside one request does
# not happen and a memo would only hold memory.
_query_memo: ContextVar[dict[str, list[float]] | None] = ContextVar(
"_query_memo", default=None,
)
@contextmanager
def query_embedding_memo():
"""Within this block, each distinct query text is embedded once."""
token = _query_memo.set({})
try:
yield
finally:
_query_memo.reset(token)
def memoized_query_embeddings(fn):
"""Run an async handler inside `query_embedding_memo()`.
For the plugin routes, each of which fans one query out to several arms.
Innermost decorator, so the scope opens after auth has passed.
"""
@functools.wraps(fn)
async def wrapper(*args, **kwargs):
with query_embedding_memo():
return await fn(*args, **kwargs)
return wrapper
async def get_embedding(text: str) -> list[float]:
"""Get an embedding vector for the given text.
Raises if the fastembed model fails to load. Callers should catch and
degrade to keyword search.
Inside `query_embedding_memo()`, a text already embedded in this request
returns the same vector without a model call. A failure is not
remembered, so the next arm tries again and degrades on its own.
"""
return (await get_embeddings([text]))[0]
memo = _query_memo.get()
if memo is not None and text in memo:
return memo[text]
vector = (await get_embeddings([text]))[0]
if memo is not None:
memo[text] = vector
return vector
async def get_embeddings(texts: list[str]) -> list[list[float]]:
+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"