Files
stack/tests/llm/test_search.py

365 lines
13 KiB
Python

"""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"}