365 lines
13 KiB
Python
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"}
|