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
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).
265 lines
9.3 KiB
Python
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
|