From be62abf1426f71c3ab0b6ab4d78f7810da45e419 Mon Sep 17 00:00:00 2001 From: Bryan Van Deusen Date: Mon, 5 Oct 2026 07:55:06 -0400 Subject: [PATCH] perf(retrieval): one embedding per query per plugin request (milestone 456 step 1, #4903) 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 --- src/scribe/routes/plugin.py | 4 ++ src/scribe/services/embeddings.py | 60 ++++++++++++++++++- tests/test_embeddings.py | 99 +++++++++++++++++++++++++++++++ 3 files changed, 162 insertions(+), 1 deletion(-) diff --git a/src/scribe/routes/plugin.py b/src/scribe/routes/plugin.py index 4c1ad3da..e1ce5986 100644 --- a/src/scribe/routes/plugin.py +++ b/src/scribe/routes/plugin.py @@ -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. diff --git a/src/scribe/services/embeddings.py b/src/scribe/services/embeddings.py index e62a84bc..cb49d97c 100644 --- a/src/scribe/services/embeddings.py +++ b/src/scribe/services/embeddings.py @@ -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]]: diff --git a/tests/test_embeddings.py b/tests/test_embeddings.py index ffaf4b16..d9c4301f 100644 --- a/tests/test_embeddings.py +++ b/tests/test_embeddings.py @@ -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"