Adds documents-ingest-embed: chunks source='paperless' envelope content (paragraph-preferring, ~600 tok/chunk, ~150 tok overlap, hard char-fallback for oversized paragraphs per plan §2 decision 3), embeds each chunk via Ollama (bge-m3, dim validated against document_chunk's VECTOR(1024) on every response) and inserts into document_chunk. Lives in documents-ingest per the plan's own recommendation (§6 step 6) rather than a new package — reuses the job family's existing idempotency/stats-balance/dry-run conventions (paperless_adapter.py, gmail-header-backfill). A 5-angle multi-agent code review of the initial implementation surfaced three real bugs, fixed here: hard_split() could infinite-loop if --chunk-overlap >= --chunk-size (now guarded in both hard_split() and main()); insert_chunk() wasn't error-isolated like embed_chunk(), so a DB write failure would crash the whole run instead of being counted and skipped; and ON CONFLICT DO NOTHING's outcome was discarded, so a silently skipped row (the known gap where document_chunk's UNIQUE constraint doesn't include `model`) would have been miscounted as a successful insert - now tracked separately as chunks_conflict_skipped and treated as a run failure. Smoke-tested and run to completion live on SOLARIA against the real Ollama instance and kb-postgres@PIHA: dry-run matched the known phase-2-step-5 figures exactly (186 fetched, 26 empty_content, 2684 chunks planned), a --limit 10 apply + idempotent re-run + DB/distance sanity checks all passed, and the full 186-document run inserted 2683/2684 chunks (1 isolated error - Ollama's runtime context window rejected one pathological dot-leader table-of-contents chunk that tokenized far more densely than estimated; documented as a known limitation, not fixed here given it's a single-chunk edge case). Timing: ~0.79s/chunk average on CPU (SOLARIA's Ollama runs GPU-less per the recent GPU-reservation-disabled fix), ~35 min wall-clock for the full pilot - the real input for scoping the later mail-corpus embedding phase (plan §7's GPU-based estimate doesn't hold here). pytest: 101 passed. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
456 lines
17 KiB
Python
456 lines
17 KiB
Python
"""Unit tests for the chunk + embed job — no DB, no real HTTP, no real Ollama."""
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
|
|
import pytest
|
|
|
|
from documents_ingest.chunk_embed import (
|
|
EmbeddingDimensionError,
|
|
OVERLAP_CHARS,
|
|
TARGET_CHARS,
|
|
chunk_text,
|
|
embed_chunk,
|
|
extract_content,
|
|
fetch_documents,
|
|
fetch_existing_chunk_keys,
|
|
hard_split,
|
|
run,
|
|
split_paragraphs,
|
|
_vector_literal,
|
|
)
|
|
|
|
|
|
def _paragraph(char: str, length: int) -> str:
|
|
return char * length
|
|
|
|
|
|
class TestSplitParagraphs:
|
|
def test_splits_on_blank_lines(self):
|
|
text = "first para\n\nsecond para\n\nthird para"
|
|
assert split_paragraphs(text) == ["first para", "second para", "third para"]
|
|
|
|
def test_no_blank_lines_returns_single_paragraph(self):
|
|
text = "line one\nline two\nline three"
|
|
assert split_paragraphs(text) == [text]
|
|
|
|
def test_drops_empty_fragments(self):
|
|
text = "a\n\n\n\n\n\nb"
|
|
assert split_paragraphs(text) == ["a", "b"]
|
|
|
|
|
|
class TestHardSplit:
|
|
def test_overlap_equal_to_size_raises(self):
|
|
# step = size - overlap would be 0 -> `start` never advances -> infinite loop.
|
|
with pytest.raises(ValueError):
|
|
hard_split("x" * 1000, size=100, overlap=100)
|
|
|
|
def test_overlap_greater_than_size_raises(self):
|
|
with pytest.raises(ValueError):
|
|
hard_split("x" * 1000, size=100, overlap=150)
|
|
|
|
def test_short_text_returns_single_chunk(self):
|
|
assert hard_split("short", size=100, overlap=10) == ["short"]
|
|
|
|
def test_splits_with_overlap(self):
|
|
text = "A" * 5000
|
|
chunks = hard_split(text, size=2400, overlap=600)
|
|
assert len(chunks) == 3
|
|
# consecutive windows overlap by exactly `overlap` characters
|
|
assert chunks[0][-600:] == chunks[1][:600]
|
|
assert chunks[1][-600:] == chunks[2][:600]
|
|
# covers the full text, in order, no gaps
|
|
assert chunks[0] + chunks[1][600:] + chunks[2][600:] == text
|
|
|
|
def test_last_chunk_reaches_end_of_text(self):
|
|
text = "B" * 5000
|
|
chunks = hard_split(text, size=2400, overlap=600)
|
|
assert chunks[-1] == text[-len(chunks[-1]):]
|
|
assert text.endswith(chunks[-1])
|
|
|
|
|
|
class TestChunkText:
|
|
def test_empty_content_returns_no_chunks(self):
|
|
assert chunk_text("") == []
|
|
assert chunk_text(None) == []
|
|
assert chunk_text(" \n\n ") == []
|
|
|
|
def test_document_shorter_than_one_chunk_returns_single_chunk(self):
|
|
text = "A short document."
|
|
assert chunk_text(text, size=TARGET_CHARS, overlap=OVERLAP_CHARS) == [text]
|
|
|
|
def test_multi_paragraph_document_splits_on_boundaries(self):
|
|
para1 = _paragraph("a", 1000)
|
|
para2 = _paragraph("b", 1000)
|
|
para3 = _paragraph("c", 1000)
|
|
text = f"{para1}\n\n{para2}\n\n{para3}"
|
|
|
|
chunks = chunk_text(text, size=2400, overlap=600)
|
|
|
|
assert len(chunks) == 2
|
|
assert para1 in chunks[0]
|
|
assert para2 in chunks[0]
|
|
assert para3 in chunks[-1]
|
|
|
|
def test_overlap_between_consecutive_chunks(self):
|
|
para1 = _paragraph("a", 1000)
|
|
para2 = _paragraph("b", 1000)
|
|
para3 = _paragraph("c", 1000)
|
|
text = f"{para1}\n\n{para2}\n\n{para3}"
|
|
|
|
chunks = chunk_text(text, size=2400, overlap=600)
|
|
|
|
# the tail of chunk[0] (the overlap window) reappears at the start of chunk[1]
|
|
assert chunks[0][-600:] == chunks[1][: len(chunks[0][-600:])]
|
|
|
|
def test_single_oversized_paragraph_falls_back_to_hard_split(self):
|
|
text = "X" * 6000 # no blank lines at all
|
|
chunks = chunk_text(text, size=2400, overlap=600)
|
|
assert len(chunks) > 1
|
|
assert all(len(c) <= 2400 for c in chunks)
|
|
|
|
def test_paragraph_larger_than_target_is_hard_split_within_mixed_document(self):
|
|
small = _paragraph("s", 100)
|
|
huge = _paragraph("h", 6000)
|
|
text = f"{small}\n\n{huge}"
|
|
|
|
chunks = chunk_text(text, size=2400, overlap=600)
|
|
|
|
assert chunks[0] == small
|
|
assert len(chunks) > 2
|
|
assert all(len(c) <= 2400 for c in chunks[1:])
|
|
|
|
|
|
class TestExtractContent:
|
|
def test_finds_content_entity(self):
|
|
entities = [{"type": "filename", "value": "a.pdf"}, {"type": "content", "text": "hello"}]
|
|
assert extract_content(entities) == "hello"
|
|
|
|
def test_missing_content_entity_returns_empty_string(self):
|
|
entities = [{"type": "filename", "value": "a.pdf"}]
|
|
assert extract_content(entities) == ""
|
|
|
|
def test_none_text_returns_empty_string(self):
|
|
entities = [{"type": "content", "text": None}]
|
|
assert extract_content(entities) == ""
|
|
|
|
def test_empty_entities_list(self):
|
|
assert extract_content([]) == ""
|
|
assert extract_content(None) == ""
|
|
|
|
|
|
class TestVectorLiteral:
|
|
def test_formats_as_bracketed_csv(self):
|
|
assert _vector_literal([0.1, 0.2, -0.3]) == "[0.1,0.2,-0.3]"
|
|
|
|
|
|
class _FakeEmbedResponse:
|
|
def __init__(self, payload, status=200):
|
|
self._payload = payload
|
|
self._status = status
|
|
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, *exc):
|
|
return False
|
|
|
|
async def json(self):
|
|
return self._payload
|
|
|
|
def raise_for_status(self):
|
|
if self._status >= 400:
|
|
raise RuntimeError(f"HTTP {self._status}")
|
|
|
|
|
|
class _FakeOllamaSession:
|
|
"""Serves a fixed embedding vector for every /api/embeddings POST, or errors by filename."""
|
|
|
|
def __init__(self, dim=1024, fail_for=None):
|
|
self._dim = dim
|
|
self._fail_for = fail_for or set()
|
|
self.requests: list[dict] = []
|
|
|
|
def post(self, url, json):
|
|
self.requests.append({"url": url, "json": json})
|
|
if json["prompt"] in self._fail_for:
|
|
return _FakeEmbedResponse({}, status=500)
|
|
return _FakeEmbedResponse({"embedding": [0.01] * self._dim})
|
|
|
|
async def close(self):
|
|
pass
|
|
|
|
|
|
class TestEmbedChunk:
|
|
async def test_returns_embedding_and_elapsed(self):
|
|
session = _FakeOllamaSession(dim=1024)
|
|
embedding, elapsed = await embed_chunk(session, "http://fake-ollama", "bge-m3", "hello world")
|
|
assert len(embedding) == 1024
|
|
assert elapsed >= 0
|
|
assert session.requests == [
|
|
{"url": "http://fake-ollama/api/embeddings", "json": {"model": "bge-m3", "prompt": "hello world"}}
|
|
]
|
|
|
|
async def test_missing_embedding_key_raises(self):
|
|
class _EmptyResponse(_FakeEmbedResponse):
|
|
pass
|
|
|
|
class _Session:
|
|
def post(self, url, json):
|
|
return _EmptyResponse({})
|
|
|
|
with pytest.raises(ValueError):
|
|
await embed_chunk(_Session(), "http://fake-ollama", "bge-m3", "text")
|
|
|
|
|
|
class _FakeConn:
|
|
def __init__(self, docs=None, existing_keys=None, existing_model="bge-m3", execute_results=None):
|
|
self._docs = docs or []
|
|
self._existing_keys = list(existing_keys or [])
|
|
self._existing_model = existing_model
|
|
# Optional queue of command tags returned by successive execute() calls, in order
|
|
# (e.g. ["INSERT 0 0"] to simulate an ON CONFLICT DO NOTHING no-op). Defaults to a
|
|
# real insert every time.
|
|
self._execute_results = list(execute_results) if execute_results is not None else None
|
|
self.execute_calls: list[tuple] = []
|
|
|
|
async def fetch(self, query, *params):
|
|
if "FROM document_chunk" in query:
|
|
# Mirrors the real `WHERE model = $1` filter: existing_keys were "written" under
|
|
# existing_model, so a query for a different model must not see them.
|
|
if params and params[0] != self._existing_model:
|
|
return []
|
|
return [{"envelope_id": eid, "chunk_index": idx} for eid, idx in self._existing_keys]
|
|
return self._docs
|
|
|
|
async def execute(self, query, *params):
|
|
self.execute_calls.append(params)
|
|
if self._execute_results is not None:
|
|
return self._execute_results.pop(0)
|
|
return "INSERT 0 1"
|
|
|
|
async def close(self):
|
|
pass
|
|
|
|
|
|
def _row(envelope_id, content):
|
|
return {"id": envelope_id, "entities": json.dumps([{"type": "content", "text": content}])}
|
|
|
|
|
|
class TestFetchHelpers:
|
|
async def test_fetch_documents_no_limit_offset(self):
|
|
conn = _FakeConn(docs=[_row("paperless:1", "hi")])
|
|
docs = await fetch_documents(conn, limit=None, offset=None)
|
|
assert docs == [_row("paperless:1", "hi")]
|
|
|
|
async def test_fetch_existing_chunk_keys(self):
|
|
conn = _FakeConn(existing_keys=[("paperless:1", 0), ("paperless:1", 1)])
|
|
keys = await fetch_existing_chunk_keys(conn, model="bge-m3")
|
|
assert keys == {("paperless:1", 0), ("paperless:1", 1)}
|
|
|
|
|
|
class TestRun:
|
|
def _patch(self, monkeypatch, conn, ollama_session):
|
|
async def _fake_connect(dsn):
|
|
return conn
|
|
monkeypatch.setattr("documents_ingest.chunk_embed.asyncpg.connect", _fake_connect)
|
|
|
|
def _fake_session_factory(*args, **kwargs):
|
|
return ollama_session
|
|
monkeypatch.setattr("documents_ingest.chunk_embed.aiohttp.ClientSession", _fake_session_factory)
|
|
|
|
async def test_dry_run_counts_without_calling_ollama_or_db(self, monkeypatch):
|
|
conn = _FakeConn(docs=[_row("paperless:1", "short doc")])
|
|
ollama = _FakeOllamaSession()
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=False)
|
|
|
|
assert stats["documents_fetched"] == 1
|
|
assert stats["documents_chunked"] == 1
|
|
assert stats["chunks_total"] == 1
|
|
assert stats["chunks_inserted"] == 1
|
|
assert ollama.requests == []
|
|
assert conn.execute_calls == []
|
|
|
|
async def test_empty_content_counted_separately_not_as_error(self, monkeypatch):
|
|
conn = _FakeConn(docs=[_row("paperless:1", ""), _row("paperless:2", "some text")])
|
|
ollama = _FakeOllamaSession()
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=False)
|
|
|
|
assert stats["empty_content"] == 1
|
|
assert stats["documents_chunked"] == 1
|
|
assert stats["documents_fetched"] == 2
|
|
assert stats["chunks_errors"] == 0
|
|
|
|
async def test_apply_embeds_and_inserts(self, monkeypatch):
|
|
conn = _FakeConn(docs=[_row("paperless:1", "some real content")])
|
|
ollama = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=True)
|
|
|
|
assert stats["chunks_inserted"] == 1
|
|
assert stats["embed_calls"] == 1
|
|
assert len(conn.execute_calls) == 1
|
|
assert len(ollama.requests) == 1
|
|
|
|
async def test_idempotent_skips_already_embedded_chunks(self, monkeypatch):
|
|
conn = _FakeConn(
|
|
docs=[_row("paperless:1", "some real content")],
|
|
existing_keys=[("paperless:1", 0)],
|
|
)
|
|
ollama = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=True)
|
|
|
|
assert stats["chunks_already_embedded"] == 1
|
|
assert stats["chunks_inserted"] == 0
|
|
assert ollama.requests == []
|
|
assert conn.execute_calls == []
|
|
|
|
async def test_rerun_after_apply_inserts_nothing_new(self, monkeypatch):
|
|
doc = _row("paperless:1", "some real content")
|
|
|
|
conn1 = _FakeConn(docs=[doc])
|
|
ollama1 = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn1, ollama1)
|
|
first = await run(dsn="postgresql://fake", apply=True)
|
|
assert first["chunks_inserted"] == 1
|
|
|
|
conn2 = _FakeConn(docs=[doc], existing_keys=[("paperless:1", 0)])
|
|
ollama2 = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn2, ollama2)
|
|
second = await run(dsn="postgresql://fake", apply=True)
|
|
|
|
assert second["chunks_inserted"] == 0
|
|
assert second["chunks_already_embedded"] == 1
|
|
assert ollama2.requests == []
|
|
|
|
async def test_dimension_mismatch_aborts(self, monkeypatch):
|
|
conn = _FakeConn(docs=[_row("paperless:1", "some real content")])
|
|
ollama = _FakeOllamaSession(dim=768) # wrong dim
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
with pytest.raises(EmbeddingDimensionError):
|
|
await run(dsn="postgresql://fake", apply=True)
|
|
|
|
async def test_embed_error_is_isolated_and_counted(self, monkeypatch):
|
|
# Two chunks: force one to fail via a document long enough to produce 2 chunks
|
|
# (chunk_size/overlap kept small to make the test fast and explicit).
|
|
big = "a" * 100 + "\n\n" + "b" * 100
|
|
conn = _FakeConn(docs=[_row("paperless:1", big)])
|
|
ollama = _FakeOllamaSession(dim=1024, fail_for={"a" * 100})
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=True, chunk_size=100, chunk_overlap=20)
|
|
|
|
assert stats["chunks_errors"] == 1
|
|
assert stats["chunks_inserted"] == 1
|
|
assert stats["chunks_total"] == (
|
|
stats["chunks_already_embedded"] + stats["chunks_inserted"]
|
|
+ stats["chunks_conflict_skipped"] + stats["chunks_errors"]
|
|
)
|
|
|
|
async def test_insert_error_is_isolated_and_counted(self, monkeypatch):
|
|
# Two chunks; the first chunk's INSERT raises (simulated DB blip) but the run
|
|
# continues and the second chunk still gets embedded and inserted normally.
|
|
big = "a" * 100 + "\n\n" + "b" * 100
|
|
conn = _FakeConn(docs=[_row("paperless:1", big)])
|
|
|
|
real_execute = conn.execute
|
|
|
|
async def _flaky_execute(query, *params):
|
|
if params[1] == 0: # chunk_index 0
|
|
conn.execute_calls.append(params)
|
|
raise RuntimeError("simulated db error")
|
|
return await real_execute(query, *params)
|
|
|
|
conn.execute = _flaky_execute
|
|
|
|
ollama = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=True, chunk_size=100, chunk_overlap=20)
|
|
|
|
assert stats["chunks_errors"] == 1
|
|
assert stats["chunks_inserted"] == 1
|
|
assert stats["chunks_total"] == (
|
|
stats["chunks_already_embedded"] + stats["chunks_inserted"]
|
|
+ stats["chunks_conflict_skipped"] + stats["chunks_errors"]
|
|
)
|
|
|
|
async def test_conflict_skipped_counted_separately_from_inserted(self, monkeypatch):
|
|
# ON CONFLICT DO NOTHING no-op: command tag reports 0 rows affected.
|
|
conn = _FakeConn(
|
|
docs=[_row("paperless:1", "some real content")], execute_results=["INSERT 0 0"]
|
|
)
|
|
ollama = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=True)
|
|
|
|
assert stats["chunks_conflict_skipped"] == 1
|
|
assert stats["chunks_inserted"] == 0
|
|
assert stats["chunks_total"] == (
|
|
stats["chunks_already_embedded"] + stats["chunks_inserted"]
|
|
+ stats["chunks_conflict_skipped"] + stats["chunks_errors"]
|
|
)
|
|
|
|
async def test_existing_keys_scoped_to_model(self, monkeypatch):
|
|
# existing_keys were "written" under a different model — the WHERE model = $1
|
|
# filter must not treat them as already-embedded for the current --model.
|
|
conn = _FakeConn(
|
|
docs=[_row("paperless:1", "some real content")],
|
|
existing_keys=[("paperless:1", 0)],
|
|
existing_model="some-other-model",
|
|
)
|
|
ollama = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=True, model="bge-m3")
|
|
|
|
assert stats["chunks_already_embedded"] == 0
|
|
assert stats["chunks_inserted"] == 1
|
|
|
|
async def test_stats_balance_documents_and_chunks(self, monkeypatch):
|
|
conn = _FakeConn(docs=[
|
|
_row("paperless:1", "content one"),
|
|
_row("paperless:2", ""),
|
|
_row("paperless:3", "content three"),
|
|
])
|
|
ollama = _FakeOllamaSession(dim=1024)
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
stats = await run(dsn="postgresql://fake", apply=True)
|
|
|
|
assert stats["documents_fetched"] == stats["empty_content"] + stats["documents_chunked"]
|
|
assert stats["chunks_total"] == (
|
|
stats["chunks_already_embedded"] + stats["chunks_inserted"]
|
|
+ stats["chunks_conflict_skipped"] + stats["chunks_errors"]
|
|
)
|
|
|
|
async def test_limit_and_offset_passed_through_to_query(self, monkeypatch):
|
|
captured = {}
|
|
|
|
class _Conn(_FakeConn):
|
|
async def fetch(self, query, *params):
|
|
if "FROM document_chunk" in query:
|
|
return []
|
|
captured["query"] = query
|
|
captured["params"] = params
|
|
return []
|
|
|
|
conn = _Conn()
|
|
ollama = _FakeOllamaSession()
|
|
self._patch(monkeypatch, conn, ollama)
|
|
|
|
await run(dsn="postgresql://fake", apply=False, limit=10, offset=5)
|
|
|
|
assert "LIMIT $1" in captured["query"]
|
|
assert "OFFSET $2" in captured["query"]
|
|
assert captured["params"] == (10, 5)
|