Files
stack/tests/llm/test_index.py
kert 26f01bd0ba
Some checks failed
CI / lint (push) Successful in 58s
CI / notebooks-smoke (push) Successful in 1m48s
Deploy / notebooks (push) Has been skipped
Deploy / zotero (push) Has been skipped
CI / test (push) Failing after 2m58s
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 / zotero (push) Successful in 25s
Infra CI / notebooks (push) Successful in 58s
Infra CI / docs (push) Successful in 1m32s
Infra CI / api (push) Successful in 1m11s
Infra CI / mc (push) Failing after 41s
Infra CI / llm (push) Successful in 1m9s
Deploy / report (push) Successful in 19s
fix(llm): filtered vector search returns results — hnsw.iterative_scan = relaxed_order on every connection and at the database level
pgvector's HNSW scan collects ef_search (40) nearest candidates and
only then applies the WHERE clause, so a selective metadata filter
(one theme inside one doctype; one docket's comments) could drop every
candidate and answer nothing. pgvector 0.8's iterative scan keeps
walking the graph until the LIMIT is met. The engine now SETs
hnsw.iterative_scan = relaxed_order on each pooled connection (a
placeholder SET needs no privilege) and PGVector is handed that engine
instead of a URL so its sessions inherit it; migrate() also tries
ALTER DATABASE (best effort — the llm role may not; applied once as
the superuser on the live database).
2026-09-24 20:07:38 -04:00

265 lines
9.3 KiB
Python

"""llm.index — incremental, resumable embedding indexer."""
from unittest.mock import MagicMock, patch
from sqlalchemy.exc import ProgrammingError
from llm.chunk import Doc, content_hash
from llm.config import LlmConfig
from llm.index import index_docs
CFG = LlmConfig(
ollama_hosts=("http://h1:11434",),
embed_model="m",
instruct_model="g",
embed_dim=768,
build_ann_index=False,
pg_host="x",
pg_port=5432,
pg_db="llm",
pg_user="llm",
)
DOC = Doc(key="K1", text="Some body text.", metadata={"docket": "D"})
def _run(docs, state_rows, force=False, cfg=CFG):
"""Run index_docs with everything external mocked; return mocks."""
store = MagicMock()
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = state_rows
with (
patch("llm.index._engine", return_value=engine),
patch("llm.index.vectorstore", return_value=store),
patch("llm.index.embed_texts", return_value=[[0.0] * 3]),
patch("llm.index.ensure_hnsw") as mock_ensure_hnsw,
patch("llm.index.HostPool") as MockPool,
# Real `_code_index` opens the live DuckDB replica — never in a
# unit test.
patch("llm.index._code_index", return_value={}),
):
MockPool.return_value.check.return_value = ["http://h1:11434"]
stats = index_docs(
docs,
collection="comments",
cfg=cfg,
pool=MockPool.return_value,
force=force,
)
return stats, store, conn, mock_ensure_hnsw
class TestIndexDocs:
def test_new_doc_embedded_and_recorded(self):
stats, store, conn, _ensure_hnsw = _run([DOC], state_rows=[])
assert stats == {"indexed": 1, "skipped": 0, "chunks": 1}
store.add_embeddings.assert_called_once()
kwargs = store.add_embeddings.call_args.kwargs
assert kwargs["ids"][0].startswith("K1:")
def test_unchanged_doc_skipped(self):
h = content_hash(DOC.text)
stats, store, _, _ensure_hnsw = _run([DOC], state_rows=[("K1", h, "")])
assert stats["skipped"] == 1
store.add_embeddings.assert_not_called()
def test_changed_doc_deletes_old_chunks_first(self):
stats, store, conn, _ensure_hnsw = _run(
[DOC], state_rows=[("K1", "stalehash", "")]
)
assert stats["indexed"] == 1
deletes = [
c
for c in conn.execute.call_args_list
if "DELETE FROM langchain_pg_embedding" in str(c.args[0])
]
assert len(deletes) == 1
sql = str(deletes[0].args[0])
assert "item_key" in sql # DELETE ... cmetadata->>'item_key'
# Scoped to the target collection, not a global delete by item_key.
assert "collection_id" in sql
assert "langchain_pg_collection" in sql
params = deletes[0].args[1]
assert params["k"] == "K1"
assert params["c"] == "comments"
def test_force_reembeds_unchanged(self):
h = content_hash(DOC.text)
stats, store, _, _ensure_hnsw = _run(
[DOC], state_rows=[("K1", h, "")], force=True
)
assert stats["indexed"] == 1
def test_empty_doc_counts_skipped(self):
empty = Doc(key="K2", text=" ", metadata={})
stats, store, _, _ensure_hnsw = _run([empty], state_rows=[])
assert stats == {"indexed": 0, "skipped": 1, "chunks": 0}
class TestPageEnrichmentHook:
def test_enriches_only_docs_that_embed(self, monkeypatch):
"""index_docs runs llm.pages.enrich_pdf_pages only on docs that reach
the embed path — a hash match must not open the doc's PDFs."""
from llm import index as index_mod
seen = []
monkeypatch.setattr(
index_mod,
"enrich_pdf_pages",
lambda doc, chunks: seen.append(doc.key) or chunks,
)
_run([DOC], [])
assert seen == ["K1"]
seen.clear()
_run([DOC], [("K1", content_hash(DOC.text), "")])
assert seen == []
class TestDeleteOldChunksFallback:
def test_programming_error_swallowed_and_indexing_continues(self):
"""A fresh DB where langchain_pg_embedding doesn't exist yet raises
ProgrammingError on DELETE; _delete_old_chunks must swallow it and
indexing must still proceed to add_embeddings + record state."""
store = MagicMock()
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
def fake_execute(clause, *args, **kwargs):
sql = str(clause)
if "DELETE FROM langchain_pg_embedding" in sql:
raise ProgrammingError("stmt", {}, Exception("relation missing"))
result = MagicMock()
result.fetchall.return_value = [] # no prior state
return result
conn.execute.side_effect = fake_execute
with (
patch("llm.index._engine", return_value=engine),
patch("llm.index.vectorstore", return_value=store),
patch("llm.index.embed_texts", return_value=[[0.0] * 3]),
patch("llm.index.ensure_hnsw"),
patch("llm.index.HostPool") as MockPool,
patch("llm.index._code_index", return_value={}),
):
MockPool.return_value.check.return_value = ["http://h1:11434"]
stats = index_docs(
[DOC],
collection="comments",
cfg=CFG,
pool=MockPool.return_value,
force=False,
)
assert stats == {"indexed": 1, "skipped": 0, "chunks": 1}
store.add_embeddings.assert_called_once()
class TestAnnIndexGating:
def test_build_ann_index_false_skips_ensure_hnsw(self):
_stats, _store, _conn, mock_ensure_hnsw = _run([DOC], state_rows=[], cfg=CFG)
mock_ensure_hnsw.assert_not_called()
def test_build_ann_index_true_calls_ensure_hnsw(self):
cfg = LlmConfig(
ollama_hosts=("http://h1:11434",),
embed_model="m",
instruct_model="g",
embed_dim=768,
build_ann_index=True,
pg_host="x",
pg_port=5432,
pg_db="llm",
pg_user="llm",
)
_stats, _store, _conn, mock_ensure_hnsw = _run([DOC], state_rows=[], cfg=cfg)
assert mock_ensure_hnsw.called
class TestPruneRulesFromCorpus:
def test_deletes_corpus_copies_and_stamps_rule_kind(self, tmp_path, monkeypatch):
from unittest.mock import MagicMock
from bib.item import Rule, Source
from bib.store import Store
from llm.index import prune_rules_from_corpus, rule_item_keys
store = Store(":memory:", storage_dir=tmp_path / "st")
r1 = store.create(
Rule(
title="CY 2027 PFS Proposed Rule",
url="https://www.federalregister.gov/d/1",
)
)
r2 = store.create(
Rule(
title="CY 2024 PFS Correction",
url="https://www.federalregister.gov/d/2",
)
)
store.create(Source(title="paper", url="https://doi.org/10.1/x"))
assert rule_item_keys(store) == sorted([r1, r2])
monkeypatch.setattr(
"llm.source.rule_kind_of",
lambda store, key, title="": {r1: "proposed", r2: "correction"}[key],
)
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.rowcount = 3
stats = prune_rules_from_corpus(engine, store)
assert stats == {
"rule_items": 2,
"corpus_chunks_deleted": 3,
"state_rows_deleted": 3,
"stamped": 6,
}
sql = [str(c.args[0]) for c in conn.execute.call_args_list]
assert (
"DELETE FROM langchain_pg_embedding" in sql[0]
and "name = 'corpus'" in sql[0]
)
assert "DELETE FROM index_state" in sql[1]
assert all("rule_kind" in q and "name = 'rules'" in q for q in sql[2:])
kinds = {
c.args[1]["key"]: c.args[1]["kind"] for c in conn.execute.call_args_list[2:]
}
assert kinds == {r1: "proposed", r2: "correction"}
store.close()
def test_nothing_to_prune(self, tmp_path):
from unittest.mock import MagicMock
from bib.store import Store
from llm.index import prune_rules_from_corpus
store = Store(":memory:", storage_dir=tmp_path / "st")
engine = MagicMock()
assert prune_rules_from_corpus(engine, store)["rule_items"] == 0
engine.begin.assert_not_called()
store.close()
class TestIterativeScanOnConnect:
def test_sets_the_guc_on_each_new_connection(self):
from unittest.mock import MagicMock
from llm.index import ITERATIVE_SCAN_SQL, set_iterative_scan
conn = MagicMock()
set_iterative_scan(conn, None)
conn.cursor.return_value.execute.assert_called_once_with(ITERATIVE_SCAN_SQL)
conn.cursor.return_value.close.assert_called_once()
def test_failure_never_propagates(self):
from unittest.mock import MagicMock
from llm.index import set_iterative_scan
conn = MagicMock()
conn.cursor.return_value.execute.side_effect = RuntimeError(
"unrecognized parameter"
)
set_iterative_scan(conn, None) # older pgvector: no raise