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>
174 lines
6.0 KiB
Python
174 lines
6.0 KiB
Python
"""Tests for services.embeddings — fastembed backend.
|
|
|
|
We don't actually load the fastembed model in tests (heavy download).
|
|
Instead, mock _get_model to return a fake that produces deterministic vectors.
|
|
"""
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
from scribe.services.embeddings import (
|
|
_cosine_similarity, get_embedding,
|
|
)
|
|
|
|
|
|
# ─── cosine_similarity (pure logic) ──────────────────────────────────────────
|
|
|
|
|
|
def test_cosine_similarity_orthogonal_is_zero():
|
|
assert _cosine_similarity([1.0, 0.0], [0.0, 1.0]) == 0.0
|
|
|
|
|
|
def test_cosine_similarity_identical_is_one():
|
|
assert _cosine_similarity([1.0, 0.0], [1.0, 0.0]) == pytest.approx(1.0)
|
|
|
|
|
|
def test_cosine_similarity_opposite_is_negative_one():
|
|
assert _cosine_similarity([1.0, 0.0], [-1.0, 0.0]) == pytest.approx(-1.0)
|
|
|
|
|
|
def test_cosine_similarity_zero_length_safe():
|
|
"""Zero-magnitude vector must not divide-by-zero."""
|
|
assert _cosine_similarity([0.0, 0.0], [1.0, 0.0]) == 0.0
|
|
assert _cosine_similarity([1.0, 0.0], [0.0, 0.0]) == 0.0
|
|
|
|
|
|
def test_cosine_similarity_mismatched_dim_returns_zero():
|
|
"""Cross-migration safety: a 768-dim vs 384-dim comparison must not crash."""
|
|
assert _cosine_similarity([1.0] * 5, [1.0] * 3) == 0.0
|
|
|
|
|
|
def test_cosine_similarity_empty_inputs():
|
|
assert _cosine_similarity([], []) == 0.0
|
|
assert _cosine_similarity([], [1.0]) == 0.0
|
|
|
|
|
|
# ─── get_embedding (fastembed path) ──────────────────────────────────────────
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_embedding_returns_list_of_floats():
|
|
"""get_embedding wraps the embedder; the result is a Python list of floats."""
|
|
fake_vec = MagicMock()
|
|
fake_vec.tolist.return_value = [0.1, 0.2, 0.3, 0.4]
|
|
fake_embedder = MagicMock()
|
|
fake_embedder.embed = MagicMock(return_value=iter([fake_vec]))
|
|
with patch(
|
|
"scribe.services.embeddings._get_model",
|
|
AsyncMock(return_value=fake_embedder),
|
|
):
|
|
out = await get_embedding("hello world")
|
|
assert out == [0.1, 0.2, 0.3, 0.4]
|
|
fake_embedder.embed.assert_called_once_with(["hello world"])
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_embedding_propagates_model_load_failures():
|
|
"""If fastembed can't initialize, the error propagates — callers catch
|
|
and degrade to keyword search."""
|
|
with patch(
|
|
"scribe.services.embeddings._get_model",
|
|
AsyncMock(side_effect=RuntimeError("model load failed")),
|
|
):
|
|
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"
|