homelab-codex-ws/jobs/documents-ingest/tests/test_chunk_embed.py

456 lines
17 KiB
Python
Raw Normal View History

feat(documents-ingest): chunk + embed job (module 5 phase 2 step 6) 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>
2026-07-15 20:41:00 +02:00
"""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)