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
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:
113
src/llm/api.py
113
src/llm/api.py
@@ -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
206
src/llm/search.py
Normal 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]
|
||||
@@ -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
364
tests/llm/test_search.py
Normal 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"}
|
||||
Reference in New Issue
Block a user