merge: #570 — /search and /similar retrieval endpoints (refs #570)
Some checks failed
CI / lint (push) Successful in 45s
CI / notebooks-smoke (push) Successful in 1m30s
Deploy / notebooks (push) Has been skipped
Deploy / zotero (push) Has been skipped
Deploy / docs (push) Has been skipped
Deploy / api (push) Has been skipped
Deploy / llm (push) Has been skipped
Deploy / mc (push) Has been skipped
Infra CI / notebooks (push) Successful in 55s
Infra CI / zotero (push) Successful in 23s
Infra CI / docs (push) Successful in 19s
Infra CI / api (push) Successful in 1m21s
Infra CI / llm (push) Successful in 50s
Infra CI / mc (push) Failing after 13s
Deploy / report (push) Successful in 15s
CI / test (push) Has started running

This commit is contained in:
kert
2026-09-11 17:18:37 -04:00
4 changed files with 843 additions and 0 deletions

View File

@@ -10,7 +10,9 @@ Run: ``uvicorn llm.api:app`` (or ``stack llm serve``).
from __future__ import annotations
import json
import logging
import re
import time
from contextlib import asynccontextmanager
from importlib import resources
from typing import AsyncIterator, Iterator
@@ -19,6 +21,8 @@ from fastapi import FastAPI, Header, HTTPException
from fastapi.responses import HTMLResponse, StreamingResponse
from pydantic import BaseModel
log = logging.getLogger(__name__)
@asynccontextmanager
async def _lifespan(app: FastAPI) -> AsyncIterator[None]:
@@ -45,6 +49,11 @@ app = FastAPI(
_ISO_DATE = re.compile(r"^\d{4}-\d{2}-\d{2}$")
_MODES = {"auto", "timeline", "recent"}
#: /search, /similar collection + kind query values — mirrors
#: llm.search.COLLECTIONS/KINDS (kept as plain literals here, like
#: _MODES above, rather than importing llm.search at module scope).
_COLLECTION_CHOICES = {"all", "comments", "rules", "corpus"}
_KIND_CHOICES = {"comment", "rule", "corpus"}
# OTel instrumentation — no-op if perf not installed or telemetry disabled.
try:
@@ -154,3 +163,107 @@ def hosts() -> dict:
with pool.acquire_generation() as host:
model = pick_model(cfg, pool, host)
return {"hosts": rows, "generation": {"host": host, "model": model}}
def _limit_ok(limit: int) -> bool:
return 1 <= limit <= 50
@app.get("/search")
def search_endpoint(
q: str = "",
collection: str = "all",
docket: str = "",
item_key: str = "",
year: str = "",
kind: str = "",
limit: int = 10,
offset: int = 0,
) -> dict:
"""Metadata-filtered similarity search over the indexed library.
``total`` is the number of results in this page (``len(results)``) —
the underlying search overfetches per collection rather than running
a separate COUNT query, so there is no cheaper true total to report.
"""
from llm import config as llm_config
from llm.pool import HostPool
from llm.search import search as run_search
query = q.strip()
if not query:
raise HTTPException(status_code=400, detail="q is required")
if collection not in _COLLECTION_CHOICES:
raise HTTPException(status_code=400, detail="unknown collection")
if kind and kind not in _KIND_CHOICES:
raise HTTPException(status_code=400, detail="unknown kind")
if not _limit_ok(limit):
raise HTTPException(status_code=400, detail="limit must be between 1 and 50")
if offset < 0:
raise HTTPException(status_code=400, detail="offset must be >= 0")
filters = {
k: v
for k, v in {
"docket": docket,
"item_key": item_key,
"year": year,
"kind": kind,
}.items()
if v
}
cfg = llm_config.load()
pool = HostPool.from_config(cfg)
t0 = time.monotonic()
results = run_search(
query,
cfg=cfg,
pool=pool,
collection=collection,
filters=filters,
limit=limit,
offset=offset,
)
log.info(
"search %r collection=%s filters=%s -> %d results in %.1fms",
query,
collection,
filters,
len(results),
(time.monotonic() - t0) * 1000,
)
return {
"query": query,
"filters": filters,
"total": len(results),
"results": results,
}
@app.get("/similar/{key}")
def similar_endpoint(key: str, collection: str = "all", limit: int = 10) -> dict:
"""Nearest neighbours of an already-indexed bib item's own chunk(s),
excluding the item itself. 404 when *key* has no indexed chunks."""
from llm import config as llm_config
from llm.pool import HostPool
from llm.search import similar as run_similar
if collection not in _COLLECTION_CHOICES:
raise HTTPException(status_code=400, detail="unknown collection")
if not _limit_ok(limit):
raise HTTPException(status_code=400, detail="limit must be between 1 and 50")
cfg = llm_config.load()
pool = HostPool.from_config(cfg)
t0 = time.monotonic()
results = run_similar(key, cfg=cfg, pool=pool, collection=collection, limit=limit)
if results is None:
raise HTTPException(status_code=404, detail=f"no indexed chunks for {key!r}")
log.info(
"similar %r collection=%s -> %d results in %.1fms",
key,
collection,
len(results),
(time.monotonic() - t0) * 1000,
)
return {"key": key, "total": len(results), "results": results}

206
src/llm/search.py Normal file
View File

@@ -0,0 +1,206 @@
"""Retrieval endpoints' core (refs #570): ``GET /search`` and
``GET /similar/{key}`` (``llm.api``).
``search`` embeds the question once (self-hosted — the Ollama pool, same
as ``llm.rag.retrieve``) and searches the selected pgvector collection(s)
directly, pushing docket/item_key/kind down to the store's own ``filter=``
argument and applying a ``year`` cutoff in Python (pgvector metadata has
no separate year column — ``date`` is an ISO string). ``similar`` instead
starts from an existing bib item's own chunk vector(s), read straight out
of ``langchain_pg_embedding`` by ``item_key``, and excludes that item from
its own results.
Both return the same source-dict shape ``llm.rag`` sends to the chat UI
(``llm.links.as_source``) plus a ``distance`` key — the raw pgvector
cosine distance behind the rounded similarity ``score`` — so a caller can
see the ranking signal directly instead of only the derived score.
"""
from __future__ import annotations
import logging
from typing import Any, Sequence
from llm.config import LlmConfig
from llm.index import _engine
from llm.links import as_source
from llm.pool import HostPool, PoolEmbeddings
from llm.rerank import _similarity
log = logging.getLogger(__name__)
#: pgvector collection names — also the valid ``collection`` query values
#: besides ``"all"``.
COLLECTIONS: tuple[str, ...] = ("comments", "rules", "corpus")
#: chunk metadata ``kind`` values — also the valid ``kind`` query values.
KINDS: tuple[str, ...] = ("comment", "rule", "corpus")
#: /similar: how many extra neighbours to fetch per collection beyond
#: *limit*, so excluding the source item's own (often top-ranked) other
#: chunks still leaves enough real neighbours to fill the page.
_SIMILAR_OVERFETCH = 10
#: /similar: at most this many of the item's own chunk vectors are
#: averaged into its "typical" vector — plenty of items have dozens of
#: chunks, and averaging them all would wash out the item's own signal.
_ITEM_VECTOR_CHUNKS = 3
def _collections_for(collection: str) -> tuple[str, ...]:
"""The pgvector collection(s) to search for *collection* — every
collection for ``""``/``"all"``, else that one collection alone.
Assumes *collection* was already validated by the API layer."""
if not collection or collection == "all":
return COLLECTIONS
return (collection,)
def _store_filter(filters: dict[str, str]) -> dict[str, str]:
"""The langchain PGVector ``filter=`` metadata dict for *filters* —
only docket/item_key/kind are pushed down to the store; ``year`` is
applied in Python (:func:`_year_ok`) since it's a prefix of the
``date`` field, not its own metadata key."""
return {
key: filters[key] for key in ("docket", "item_key", "kind") if filters.get(key)
}
def _year_ok(md: dict[str, Any], year: str) -> bool:
if not year:
return True
return str(md.get("date", ""))[:4] == year
def _hit_to_result(doc: Any, distance: float) -> dict:
md = {k: str(v) for k, v in (doc.metadata or {}).items()}
result = as_source(md, doc.page_content, _similarity(float(distance)))
result["distance"] = round(float(distance), 4)
return result
def search(
q: str,
*,
cfg: LlmConfig,
pool: HostPool,
collection: str = "all",
filters: dict[str, str] | None = None,
limit: int = 10,
offset: int = 0,
) -> list[dict]:
"""Metadata-filtered similarity search across the selected
collection(s), embedding *q* once.
*filters* may carry ``docket``, ``item_key``, ``year`` and ``kind``
(all optional; an absent/empty value is not filtered on). Each
selected collection is overfetched ``limit + offset`` hits (so a
later page can still be answered without a second round trip per
collection); results are merged, the ``year`` filter (if any) is
applied, the merged list is sorted by distance ascending (closest
first) and sliced to ``[offset : offset + limit]``.
Validation (empty *q*, ``limit``/``collection``/``kind`` ranges) is
the API layer's job — this function assumes its inputs are already
sane.
"""
from llm.index import vectorstore
filters = filters or {}
year = str(filters.get("year") or "")
store_filter = _store_filter(filters)
fetch = limit + offset
vec = PoolEmbeddings(pool, cfg.embed_model).embed_query(q)
hits: list[dict] = []
for name in _collections_for(collection):
store = vectorstore(name, cfg, pool)
kwargs: dict[str, Any] = {"k": fetch}
if store_filter:
kwargs["filter"] = store_filter
for doc, distance in store.similarity_search_with_score_by_vector(
vec, **kwargs
):
if not _year_ok(doc.metadata or {}, year):
continue
hits.append(_hit_to_result(doc, distance))
hits.sort(key=lambda h: h["distance"])
return hits[offset : offset + limit]
def _sa_text(sql: str) -> Any:
"""``sqlalchemy.text`` imported at call time — mirrors
``llm.evidence._sa_text``: sqlalchemy is part of the ``llm`` extra,
kept out of this module's top-level imports for the same reason."""
from sqlalchemy import text
return text(sql)
_ITEM_VECTOR_SQL = (
"SELECT embedding::text FROM langchain_pg_embedding "
"WHERE cmetadata->>'item_key' = :key "
"ORDER BY id LIMIT :n"
)
def _parse_vector(text: str) -> list[float]:
"""``"[0.1,0.2,...]"`` (pgvector's text output format) -> a plain
float list — parsed by hand rather than via the ``pgvector`` package
so this module needs no vector-type adapter registered on *engine*."""
return [float(x) for x in text.strip("[]").split(",") if x.strip()]
def _average_vector(vectors: Sequence[list[float]]) -> list[float]:
dim = len(vectors[0])
return [sum(v[i] for v in vectors) / len(vectors) for i in range(dim)]
def _item_vector(engine: Any, key: str) -> list[float] | None:
"""*key*'s own embedding — its first indexed chunk's vector, or the
average of up to :data:`_ITEM_VECTOR_CHUNKS` when it has more (a
stable "typical" vector, not weighted toward whichever chunk happens
to sort first). ``None`` when *key* has no indexed chunks in any
collection."""
with engine.begin() as conn:
rows = conn.execute(
_sa_text(_ITEM_VECTOR_SQL), {"key": key, "n": _ITEM_VECTOR_CHUNKS}
).fetchall()
if not rows:
return None
return _average_vector([_parse_vector(r[0]) for r in rows])
def similar(
key: str,
*,
cfg: LlmConfig,
pool: HostPool,
collection: str = "all",
limit: int = 10,
) -> list[dict] | None:
"""Nearest neighbours of bib item *key*'s own chunk(s) across the
selected collection(s), excluding *key* itself from the results.
``None`` when *key* has no indexed chunks anywhere — the API layer's
cue to 404. Each selected collection is overfetched by
:data:`_SIMILAR_OVERFETCH` beyond *limit* so dropping the source
item's own (often top-ranked) other chunks still leaves enough real
neighbours to fill the page; the merged, filtered list is sorted by
distance ascending and sliced to *limit*.
"""
vec = _item_vector(_engine(cfg), key)
if vec is None:
return None
from llm.index import vectorstore
fetch = limit + _SIMILAR_OVERFETCH
hits: list[dict] = []
for name in _collections_for(collection):
store = vectorstore(name, cfg, pool)
for doc, distance in store.similarity_search_with_score_by_vector(vec, k=fetch):
md = doc.metadata or {}
if str(md.get("item_key", "")) == key:
continue
hits.append(_hit_to_result(doc, distance))
hits.sort(key=lambda h: h["distance"])
return hits[:limit]

View File

@@ -208,3 +208,163 @@ class TestHosts:
r = client.get("/hosts")
assert r.json()["generation"] is None
assert "no Ollama host" in r.json()["error"]
class TestSearchEndpoint:
@patch("llm.search.search")
@patch("llm.config.load")
def test_empty_q_400(self, mock_load, mock_search):
mock_load.return_value = _cfg()
r = client.get("/search", params={"q": " "})
assert r.status_code == 400
mock_search.assert_not_called()
@patch("llm.search.search")
@patch("llm.config.load")
def test_unknown_collection_400(self, mock_load, mock_search):
mock_load.return_value = _cfg()
r = client.get("/search", params={"q": "ccm", "collection": "bogus"})
assert r.status_code == 400
mock_search.assert_not_called()
@patch("llm.search.search")
@patch("llm.config.load")
def test_unknown_kind_400(self, mock_load, mock_search):
mock_load.return_value = _cfg()
r = client.get("/search", params={"q": "ccm", "kind": "bogus"})
assert r.status_code == 400
mock_search.assert_not_called()
@patch("llm.search.search")
@patch("llm.config.load")
def test_limit_out_of_range_400(self, mock_load, mock_search):
mock_load.return_value = _cfg()
assert client.get("/search", params={"q": "ccm", "limit": 0}).status_code == 400
assert (
client.get("/search", params={"q": "ccm", "limit": 51}).status_code == 400
)
mock_search.assert_not_called()
@patch("llm.search.search")
@patch("llm.config.load")
def test_negative_offset_400(self, mock_load, mock_search):
mock_load.return_value = _cfg()
r = client.get("/search", params={"q": "ccm", "offset": -1})
assert r.status_code == 400
mock_search.assert_not_called()
@patch("llm.search.search")
@patch("llm.config.load")
def test_filtered_search_scopes_by_docket_and_returns_shape(
self, mock_load, mock_search
):
mock_load.return_value = _cfg()
src = {
"id": "CMS-2023-0121-1",
"label": "CMS-2023-0121-1",
"kind": "comment",
"url": "u",
"title": "t",
"date": "2024-01-01",
"docket": "CMS-2023-0121",
"comment_id": "CMS-2023-0121-1",
"snippet": "s",
"score": 0.9,
"item_key": "K1",
"p_id": "",
"seq": "1",
"section": "",
"distance": 0.1,
}
mock_search.return_value = [src]
r = client.get(
"/search",
params={
"q": "chronic care management",
"docket": "CMS-2023-0121",
"collection": "comments",
"limit": 10,
},
)
assert r.status_code == 200
body = r.json()
assert body["query"] == "chronic care management"
assert body["filters"] == {"docket": "CMS-2023-0121"}
assert body["total"] == 1
assert body["results"] == [src]
assert all(s["docket"] == "CMS-2023-0121" for s in body["results"])
kwargs = mock_search.call_args.kwargs
assert kwargs["collection"] == "comments"
assert kwargs["filters"] == {"docket": "CMS-2023-0121"}
assert kwargs["limit"] == 10
assert kwargs["offset"] == 0
@patch("llm.search.search")
@patch("llm.config.load")
def test_defaults(self, mock_load, mock_search):
mock_load.return_value = _cfg()
mock_search.return_value = []
r = client.get("/search", params={"q": "q"})
assert r.status_code == 200
kwargs = mock_search.call_args.kwargs
assert kwargs["collection"] == "all"
assert kwargs["filters"] == {}
assert kwargs["limit"] == 10
assert kwargs["offset"] == 0
class TestSimilarEndpoint:
@patch("llm.search.similar")
@patch("llm.config.load")
def test_unknown_collection_400(self, mock_load, mock_similar):
mock_load.return_value = _cfg()
r = client.get("/similar/K1", params={"collection": "bogus"})
assert r.status_code == 400
mock_similar.assert_not_called()
@patch("llm.search.similar")
@patch("llm.config.load")
def test_limit_out_of_range_400(self, mock_load, mock_similar):
mock_load.return_value = _cfg()
assert client.get("/similar/K1", params={"limit": 0}).status_code == 400
assert client.get("/similar/K1", params={"limit": 51}).status_code == 400
mock_similar.assert_not_called()
@patch("llm.search.similar", return_value=None)
@patch("llm.config.load")
def test_unknown_key_404(self, mock_load, mock_similar):
mock_load.return_value = _cfg()
r = client.get("/similar/NOPE")
assert r.status_code == 404
@patch("llm.search.similar")
@patch("llm.config.load")
def test_returns_neighbours_shape(self, mock_load, mock_similar):
mock_load.return_value = _cfg()
src = {
"id": "91 FR 43949 ¶1",
"label": "91 FR 43949 ¶1",
"kind": "rule",
"url": "u",
"title": "t",
"date": "2024-01-01",
"docket": "",
"comment_id": "",
"snippet": "s",
"score": 0.8,
"item_key": "R1",
"p_id": "1",
"seq": "",
"section": "",
"distance": 0.2,
}
mock_similar.return_value = [src]
r = client.get("/similar/K1", params={"limit": 5})
assert r.status_code == 200
body = r.json()
assert body["key"] == "K1"
assert body["total"] == 1
assert body["results"] == [src]
kwargs = mock_similar.call_args.kwargs
assert kwargs["collection"] == "all"
assert kwargs["limit"] == 5

364
tests/llm/test_search.py Normal file
View File

@@ -0,0 +1,364 @@
"""llm.search — /search and /similar retrieval core (refs #570)."""
from __future__ import annotations
from unittest.mock import MagicMock, patch
from langchain_core.documents import Document
import llm.search as search
from llm.config import LlmConfig
from llm.search import (
_average_vector,
_collections_for,
_item_vector,
_parse_vector,
_store_filter,
_year_ok,
)
from llm.search import search as run_search
from llm.search import similar as run_similar
CFG = LlmConfig(
ollama_hosts=("http://h1:11434",),
host_vram={"http://h1:11434": 24},
embed_model="embed",
instruct_model="chat",
embed_dim=3,
pg_host="x",
pg_port=5432,
pg_db="llm",
pg_user="llm",
build_ann_index=True,
)
def _doc(text, **md):
return Document(page_content=text, metadata=md)
def _stores(by_collection):
"""vectorstore(collection, cfg, pool) -> a MagicMock store whose
similarity_search_with_score_by_vector returns by_collection[name]
and records every call for assertions."""
def factory(collection, cfg, pool):
s = MagicMock()
s.similarity_search_with_score_by_vector.return_value = by_collection.get(
collection, []
)
return s
return factory
class TestCollectionsFor:
def test_all_returns_every_collection(self):
assert _collections_for("all") == search.COLLECTIONS
def test_empty_returns_every_collection(self):
assert _collections_for("") == search.COLLECTIONS
def test_named_collection_returns_just_that_one(self):
assert _collections_for("comments") == ("comments",)
class TestStoreFilter:
def test_only_docket_item_key_kind_pass_through(self):
out = _store_filter(
{
"docket": "CMS-2023-0121",
"item_key": "K1",
"kind": "comment",
"year": "2024",
}
)
assert out == {"docket": "CMS-2023-0121", "item_key": "K1", "kind": "comment"}
def test_empty_values_are_dropped(self):
assert _store_filter({"docket": "", "item_key": None}) == {}
def test_no_filters_is_empty_dict(self):
assert _store_filter({}) == {}
class TestYearOk:
def test_no_year_filter_is_a_noop(self):
assert _year_ok({"date": "2020-01-01"}, "") is True
def test_matching_year(self):
assert _year_ok({"date": "2024-06-01"}, "2024") is True
def test_non_matching_year(self):
assert _year_ok({"date": "2024-06-01"}, "2023") is False
def test_missing_date_fails_a_year_filter(self):
assert _year_ok({}, "2024") is False
class TestParseVector:
def test_parses_bracketed_csv(self):
assert _parse_vector("[0.1,0.2,0.3]") == [0.1, 0.2, 0.3]
class TestAverageVector:
def test_averages_componentwise(self):
assert _average_vector([[0.0, 2.0], [2.0, 4.0]]) == [1.0, 3.0]
def test_single_vector_is_unchanged(self):
assert _average_vector([[1.0, 2.0]]) == [1.0, 2.0]
class TestItemVector:
def test_no_rows_returns_none(self):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = []
assert _item_vector(engine, "NOPE") is None
def test_averages_up_to_three_rows(self):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = [
("[0.0,0.0]",),
("[2.0,4.0]",),
]
assert _item_vector(engine, "K1") == [1.0, 2.0]
def test_query_params_carry_key_and_chunk_cap(self):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = [("[1.0]",)]
_item_vector(engine, "K1")
assert conn.execute.call_args.args[1] == {"key": "K1", "n": 3}
class TestSearch:
@patch("llm.search.PoolEmbeddings")
@patch("llm.index.vectorstore")
def test_embeds_once_and_returns_source_shape_plus_distance(self, mock_vs, MockEmb):
MockEmb.return_value.embed_query.return_value = [0.1, 0.2, 0.3]
mock_vs.side_effect = _stores(
{
"comments": [
(
_doc(
"Telehealth comment.",
kind="comment",
comment_id="CMS-2023-0121-1",
docket="CMS-2023-0121",
item_key="K1",
date="2024-08-19",
),
0.2,
)
]
}
)
out = run_search("chronic care", cfg=CFG, pool=MagicMock(), limit=10)
assert MockEmb.return_value.embed_query.call_count == 1
assert len(out) == 1
r = out[0]
assert r["distance"] == 0.2
assert r["item_key"] == "K1"
assert r["docket"] == "CMS-2023-0121"
assert 0 < r["score"] <= 1
@patch("llm.search.PoolEmbeddings")
@patch("llm.index.vectorstore")
def test_collection_all_searches_every_collection(self, mock_vs, MockEmb):
MockEmb.return_value.embed_query.return_value = [0.0, 0.0, 0.0]
stores = {}
def factory(collection, cfg, pool):
s = MagicMock()
s.similarity_search_with_score_by_vector.return_value = []
stores[collection] = s
return s
mock_vs.side_effect = factory
run_search("q", cfg=CFG, pool=MagicMock())
assert set(stores) == {"comments", "rules", "corpus"}
@patch("llm.search.PoolEmbeddings")
@patch("llm.index.vectorstore")
def test_named_collection_searches_only_that_store(self, mock_vs, MockEmb):
MockEmb.return_value.embed_query.return_value = [0.0, 0.0, 0.0]
stores = {}
def factory(collection, cfg, pool):
s = MagicMock()
s.similarity_search_with_score_by_vector.return_value = []
stores[collection] = s
return s
mock_vs.side_effect = factory
run_search("q", cfg=CFG, pool=MagicMock(), collection="rules")
assert set(stores) == {"rules"}
@patch("llm.search.PoolEmbeddings")
@patch("llm.index.vectorstore")
def test_filters_pass_through_to_store_filter_kwarg(self, mock_vs, MockEmb):
MockEmb.return_value.embed_query.return_value = [0.0, 0.0, 0.0]
store = MagicMock()
store.similarity_search_with_score_by_vector.return_value = []
mock_vs.return_value = store
run_search(
"q",
cfg=CFG,
pool=MagicMock(),
collection="comments",
filters={"docket": "CMS-2023-0121", "item_key": "K1", "kind": "comment"},
limit=5,
offset=2,
)
kwargs = store.similarity_search_with_score_by_vector.call_args.kwargs
assert kwargs["filter"] == {
"docket": "CMS-2023-0121",
"item_key": "K1",
"kind": "comment",
}
assert kwargs["k"] == 7 # limit + offset
@patch("llm.search.PoolEmbeddings")
@patch("llm.index.vectorstore")
def test_year_filter_applied_in_python(self, mock_vs, MockEmb):
MockEmb.return_value.embed_query.return_value = [0.0, 0.0, 0.0]
mock_vs.side_effect = _stores(
{
"comments": [
(_doc("old", item_key="A", date="2019-01-01"), 0.1),
(_doc("new", item_key="B", date="2024-06-01"), 0.3),
]
}
)
out = run_search(
"q",
cfg=CFG,
pool=MagicMock(),
collection="comments",
filters={"year": "2024"},
)
assert [r["item_key"] for r in out] == ["B"]
@patch("llm.search.PoolEmbeddings")
@patch("llm.index.vectorstore")
def test_pagination_slices_the_merged_sorted_list(self, mock_vs, MockEmb):
MockEmb.return_value.embed_query.return_value = [0.0, 0.0, 0.0]
mock_vs.side_effect = _stores(
{
"comments": [
(_doc("a", item_key="A"), 0.1),
(_doc("b", item_key="B"), 0.2),
(_doc("c", item_key="C"), 0.3),
]
}
)
page1 = run_search(
"q", cfg=CFG, pool=MagicMock(), collection="comments", limit=2, offset=0
)
page2 = run_search(
"q", cfg=CFG, pool=MagicMock(), collection="comments", limit=2, offset=2
)
assert [r["item_key"] for r in page1] == ["A", "B"]
assert [r["item_key"] for r in page2] == ["C"]
@patch("llm.search.PoolEmbeddings")
@patch("llm.index.vectorstore")
def test_merges_and_sorts_by_distance_across_collections(self, mock_vs, MockEmb):
MockEmb.return_value.embed_query.return_value = [0.0, 0.0, 0.0]
mock_vs.side_effect = _stores(
{
"comments": [(_doc("c", item_key="C"), 0.30)],
"rules": [(_doc("r", item_key="R"), 0.05)],
"corpus": [(_doc("x", item_key="X"), 0.15)],
}
)
out = run_search("q", cfg=CFG, pool=MagicMock(), limit=10)
assert [r["item_key"] for r in out] == ["R", "X", "C"]
assert [r["distance"] for r in out] == [0.05, 0.15, 0.30]
class TestSimilar:
@patch("llm.search._engine")
def test_missing_item_returns_none(self, mock_engine):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = []
mock_engine.return_value = engine
assert run_similar("NOPE", cfg=CFG, pool=MagicMock()) is None
@patch("llm.index.vectorstore")
@patch("llm.search._engine")
def test_excludes_the_source_item_and_returns_neighbours(
self, mock_engine, mock_vs
):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = [("[0.1,0.2,0.3]",)]
mock_engine.return_value = engine
mock_vs.side_effect = _stores(
{
"comments": [
(_doc("self chunk", item_key="K1"), 0.0),
(_doc("neighbour", item_key="K2"), 0.1),
]
}
)
out = run_similar("K1", cfg=CFG, pool=MagicMock(), collection="comments")
assert [r["item_key"] for r in out] == ["K2"]
@patch("llm.index.vectorstore")
@patch("llm.search._engine")
def test_uses_the_items_own_vector_for_the_search(self, mock_engine, mock_vs):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = [("[0.5,0.25,0.0]",)]
mock_engine.return_value = engine
store = MagicMock()
store.similarity_search_with_score_by_vector.return_value = []
mock_vs.return_value = store
run_similar("K1", cfg=CFG, pool=MagicMock(), collection="rules")
args, kwargs = store.similarity_search_with_score_by_vector.call_args
assert args[0] == [0.5, 0.25, 0.0]
assert kwargs["k"] == 10 + search._SIMILAR_OVERFETCH
@patch("llm.index.vectorstore")
@patch("llm.search._engine")
def test_limit_caps_the_merged_result(self, mock_engine, mock_vs):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = [("[0.1,0.1]",)]
mock_engine.return_value = engine
mock_vs.side_effect = _stores(
{
"comments": [
(_doc("n1", item_key="N1"), 0.10),
(_doc("n2", item_key="N2"), 0.20),
(_doc("n3", item_key="N3"), 0.30),
]
}
)
out = run_similar(
"K1", cfg=CFG, pool=MagicMock(), collection="comments", limit=2
)
assert [r["item_key"] for r in out] == ["N1", "N2"]
@patch("llm.index.vectorstore")
@patch("llm.search._engine")
def test_all_collections_searched_by_default(self, mock_engine, mock_vs):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = [("[0.1,0.1]",)]
mock_engine.return_value = engine
stores = {}
def factory(collection, cfg, pool):
s = MagicMock()
s.similarity_search_with_score_by_vector.return_value = []
stores[collection] = s
return s
mock_vs.side_effect = factory
run_similar("K1", cfg=CFG, pool=MagicMock())
assert set(stores) == {"comments", "rules", "corpus"}