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
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:
@@ -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 report_check as report_check_svc
|
||||||
from scribe.services import shape_check as shape_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 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
|
from scribe.services.settings import get_admin_setting, set_setting
|
||||||
|
|
||||||
plugin_bp = Blueprint("plugin", __name__, url_prefix="/api/plugin")
|
plugin_bp = Blueprint("plugin", __name__, url_prefix="/api/plugin")
|
||||||
@@ -88,6 +89,7 @@ async def session_context():
|
|||||||
|
|
||||||
@plugin_bp.get("/retrieve")
|
@plugin_bp.get("/retrieve")
|
||||||
@login_required
|
@login_required
|
||||||
|
@memoized_query_embeddings
|
||||||
async def autoinject_retrieve():
|
async def autoinject_retrieve():
|
||||||
"""Title-first knowledge auto-inject for the plugin's UserPromptSubmit hook.
|
"""Title-first knowledge auto-inject for the plugin's UserPromptSubmit hook.
|
||||||
|
|
||||||
@@ -163,6 +165,7 @@ async def autoinject_retrieve():
|
|||||||
|
|
||||||
@plugin_bp.get("/tool-rules")
|
@plugin_bp.get("/tool-rules")
|
||||||
@login_required
|
@login_required
|
||||||
|
@memoized_query_embeddings
|
||||||
async def pre_tool_rules():
|
async def pre_tool_rules():
|
||||||
"""Standing rules for the plugin's PreToolUse hook on ACTIONS (#3476).
|
"""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")
|
@plugin_bp.get("/prior-art")
|
||||||
@login_required
|
@login_required
|
||||||
|
@memoized_query_embeddings
|
||||||
async def write_path_prior_art():
|
async def write_path_prior_art():
|
||||||
"""Prior-art hint for the plugin's PreToolUse hook on Write/Edit.
|
"""Prior-art hint for the plugin's PreToolUse hook on Write/Edit.
|
||||||
|
|
||||||
|
|||||||
@@ -10,11 +10,14 @@ volume so subsequent boots are instant).
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import functools
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
|
|
||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from contextvars import ContextVar
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
@@ -77,13 +80,68 @@ async def _get_model():
|
|||||||
return _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]:
|
async def get_embedding(text: str) -> list[float]:
|
||||||
"""Get an embedding vector for the given text.
|
"""Get an embedding vector for the given text.
|
||||||
|
|
||||||
Raises if the fastembed model fails to load. Callers should catch and
|
Raises if the fastembed model fails to load. Callers should catch and
|
||||||
degrade to keyword search.
|
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]]:
|
async def get_embeddings(texts: list[str]) -> list[list[float]]:
|
||||||
|
|||||||
@@ -72,3 +72,102 @@ async def test_get_embedding_propagates_model_load_failures():
|
|||||||
):
|
):
|
||||||
with pytest.raises(RuntimeError, match="model load failed"):
|
with pytest.raises(RuntimeError, match="model load failed"):
|
||||||
await get_embedding("x")
|
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"
|
||||||
|
|||||||
Reference in New Issue
Block a user