Always use the Store-owned embedder

This commit is contained in:
Yiorgis Gozadinos 2026-05-29 11:27:08 +03:00
parent 0a9a444996
commit e7c7df2915
No known key found for this signature in database
16 changed files with 157 additions and 69 deletions

View file

@ -9,6 +9,8 @@
- A successful DELETE job auto-prunes dead jobs with the same `(source_id, uri)`. New `JobRepo.prune_dead(source_id, uri)`. - A successful DELETE job auto-prunes dead jobs with the same `(source_id, uri)`. New `JobRepo.prune_dead(source_id, uri)`.
- Reranker built once per `HaikuRAG` client instead of per `search()` call. - Reranker built once per `HaikuRAG` client instead of per `search()` call.
- Embedder built once per `HaikuRAG` client instead of rebuilt per ingest/search operation. `embed_chunks` now takes it explicitly: `embed_chunks(chunks, embedder, config)` (was `embed_chunks(chunks, config)`). New `HaikuRAG.embedder` property exposes it.
- `HaikuRAG` close runs a final vacuum after draining in-flight background vacuums when `storage.auto_vacuum` is enabled.
## [0.50.0] - 2026-05-27 ## [0.50.0] - 2026-05-27

View file

@ -26,7 +26,7 @@ The client exposes four primitives that can be composed into custom workflows:
|-----------|-------|--------|---------| |-----------|-------|--------|---------|
| `convert()` | file, URL, or text | `DoclingDocument` | Convert source to structured document | | `convert()` | file, URL, or text | `DoclingDocument` | Convert source to structured document |
| `chunk()` | `DoclingDocument` | `list[Chunk]` | Split document into chunks | | `chunk()` | `DoclingDocument` | `list[Chunk]` | Split document into chunks |
| `embed_chunks()` | `list[Chunk]` | `list[Chunk]` | Generate embeddings for chunks (includes contextualization) | | `embed_chunks()` | `list[Chunk]`, embedder | `list[Chunk]` | Generate embeddings for chunks (includes contextualization) |
| `contextualize()` | `list[Chunk]` | `list[str]` | Get embedding-ready text (for custom embedders only) | | `contextualize()` | `list[Chunk]` | `list[str]` | Get embedding-ready text (for custom embedders only) |
## Basic Pipeline ## Basic Pipeline
@ -45,7 +45,7 @@ async with HaikuRAG("database.lancedb", create=True) as client:
chunks = await client.chunk(docling_doc) chunks = await client.chunk(docling_doc)
# 3. Generate embeddings # 3. Generate embeddings
embedded_chunks = await embed_chunks(chunks) embedded_chunks = await embed_chunks(chunks, client.embedder)
# 4. Store the document with chunks # 4. Store the document with chunks
doc = await client.import_document( doc = await client.import_document(
@ -117,13 +117,13 @@ Chunks are returned with:
## Embed ## Embed
`embed_chunks()` generates embeddings for chunks. It automatically contextualizes chunks (prepends section headings) before embedding for better semantic search, without modifying the stored content: `embed_chunks()` generates embeddings for chunks using the client's embedder. It automatically contextualizes chunks (prepends section headings) before embedding for better semantic search, without modifying the stored content:
```python ```python
from haiku.rag.embeddings import embed_chunks from haiku.rag.embeddings import embed_chunks
# Generate embeddings (returns new Chunk objects) # Generate embeddings (returns new Chunk objects)
embedded_chunks = await embed_chunks(chunks) embedded_chunks = await embed_chunks(chunks, client.embedder)
# Original chunks unchanged # Original chunks unchanged
assert chunks[0].embedding is None assert chunks[0].embedding is None
@ -175,7 +175,7 @@ async with HaikuRAG("database.lancedb", create=True) as client:
# Continue with standard pipeline # Continue with standard pipeline
chunks = await client.chunk(processed_doc) chunks = await client.chunk(processed_doc)
embedded_chunks = await embed_chunks(chunks) embedded_chunks = await embed_chunks(chunks, client.embedder)
await client.import_document( await client.import_document(
chunks=embedded_chunks, chunks=embedded_chunks,
@ -203,7 +203,7 @@ async with HaikuRAG("database.lancedb", create=True) as client:
for i, chunk in enumerate(filtered): for i, chunk in enumerate(filtered):
chunk.order = i chunk.order = i
embedded_chunks = await embed_chunks(filtered) embedded_chunks = await embed_chunks(filtered, client.embedder)
await client.import_document( await client.import_document(
docling_document=docling_doc, docling_document=docling_doc,

View file

@ -31,6 +31,7 @@ if TYPE_CHECKING:
from docling_core.types.doc.document import DoclingDocument from docling_core.types.doc.document import DoclingDocument
from PIL import Image as PILImage from PIL import Image as PILImage
from haiku.rag.embeddings import EmbedderWrapper
from haiku.rag.ingester.sources.base import Source from haiku.rag.ingester.sources.base import Source
from haiku.rag.reranking.base import RerankerBase from haiku.rag.reranking.base import RerankerBase
from haiku.rag.sandbox import AnalysisResult from haiku.rag.sandbox import AnalysisResult
@ -89,6 +90,11 @@ class HaikuRAG:
"""Whether the client is in read-only mode.""" """Whether the client is in read-only mode."""
return self.store.is_read_only return self.store.is_read_only
@property
def embedder(self) -> "EmbedderWrapper":
"""The embedder owned by the Store, reused across all operations."""
return self.store.embedder
@cached_property @cached_property
def reranker(self) -> "RerankerBase | None": def reranker(self) -> "RerankerBase | None":
"""The configured reranker, built once and reused across searches. """The configured reranker, built once and reused across searches.
@ -128,14 +134,23 @@ class HaikuRAG:
return False return False
async def _await_vacuum_tasks(self) -> None: async def _await_vacuum_tasks(self) -> None:
"""Wait for all in-flight background vacuum tasks to complete. """Drain background vacuum work before tearing down the connection.
Each create_document / update_document can schedule its own vacuum task; Each create_document / update_document can schedule its own vacuum task;
all must be awaited before tearing down the connection, not just the all must be awaited, not just the most recently scheduled one. Vacuum
most recently scheduled one. skips when another is already running, so the cleanup for the final
writes may have been a no-op. Run one more pass once the in-flight tasks
are done to collapse versions created after the last vacuum took the lock.
""" """
if self._vacuum_tasks: if not self._vacuum_tasks:
await asyncio.gather(*self._vacuum_tasks, return_exceptions=True) return
await asyncio.gather(*self._vacuum_tasks, return_exceptions=True)
# __aexit__ runs during exception unwinding; a raising vacuum here would
# mask the original exception, so the drain stays best-effort.
try:
await self.store.vacuum()
except Exception:
logger.debug("Final vacuum on close failed", exc_info=True)
def _schedule_vacuum(self) -> None: def _schedule_vacuum(self) -> None:
"""Schedule a background vacuum and track the task for later awaiting.""" """Schedule a background vacuum and track the task for later awaiting."""
@ -180,6 +195,7 @@ class HaikuRAG:
return await chunk( return await chunk(
self._config, self._config,
docling_document, docling_document,
embedder=self.embedder,
existing_picture_data=existing_picture_data, existing_picture_data=existing_picture_data,
document_id=document_id, document_id=document_id,
) )

View file

@ -59,7 +59,7 @@ async def _store_document_with_chunks(
Handles versioning/rollback on failure. Handles versioning/rollback on failure.
""" """
chunks = await ensure_chunks_embedded(client._config, chunks) chunks = await ensure_chunks_embedded(client._config, chunks, client.embedder)
versions = await client.store.current_table_versions() versions = await client.store.current_table_versions()
@ -109,7 +109,7 @@ async def _update_document_with_chunks(
await client.document_item_repository.get_all_picture_data(document.id) await client.document_item_repository.get_all_picture_data(document.id)
) )
chunks = await ensure_chunks_embedded(client._config, chunks) chunks = await ensure_chunks_embedded(client._config, chunks, client.embedder)
versions = await client.store.current_table_versions() versions = await client.store.current_table_versions()
@ -160,7 +160,7 @@ async def create_document(
converter = get_converter(client._config) converter = get_converter(client._config)
docling_document = await converter.convert_text(content, format=format) docling_document = await converter.convert_text(content, format=format)
chunks = await client.chunk(docling_document) chunks = await client.chunk(docling_document)
embedded_chunks = await embed_chunks(chunks, client._config) embedded_chunks = await embed_chunks(chunks, client.embedder, client._config)
stored_content = docling_document.export_to_markdown() stored_content = docling_document.export_to_markdown()
@ -286,7 +286,9 @@ async def _ingest_fetch_result(
chunks = await client.chunk(docling_document) chunks = await client.chunk(docling_document)
chunk_span.set_attribute("chunks_created", len(chunks)) chunk_span.set_attribute("chunks_created", len(chunks))
with logfire.span("document.embed", uri=result.uri): with logfire.span("document.embed", uri=result.uri):
embedded_chunks = await embed_chunks(chunks, client._config) embedded_chunks = await embed_chunks(
chunks, client.embedder, client._config
)
finally: finally:
if cleanup_path is not None: if cleanup_path is not None:
cleanup_path.unlink(missing_ok=True) cleanup_path.unlink(missing_ok=True)
@ -635,7 +637,9 @@ async def update_document(
existing_doc.set_docling(docling_document) existing_doc.set_docling(docling_document)
new_chunks = await client.chunk(docling_document) new_chunks = await client.chunk(docling_document)
embedded_chunks = await embed_chunks(new_chunks, client._config) embedded_chunks = await embed_chunks(
new_chunks, client.embedder, client._config
)
return await _update_document_with_chunks( return await _update_document_with_chunks(
client, existing_doc, embedded_chunks, docling_document client, existing_doc, embedded_chunks, docling_document
) )
@ -647,7 +651,7 @@ async def update_document(
existing_doc.set_docling(converted_docling) existing_doc.set_docling(converted_docling)
new_chunks = await client.chunk(converted_docling) new_chunks = await client.chunk(converted_docling)
embedded_chunks = await embed_chunks(new_chunks, client._config) embedded_chunks = await embed_chunks(new_chunks, client.embedder, client._config)
return await _update_document_with_chunks( return await _update_document_with_chunks(
client, existing_doc, embedded_chunks, converted_docling client, existing_doc, embedded_chunks, converted_docling
) )

View file

@ -15,6 +15,8 @@ from haiku.rag.store.models.document_item import _picture_description_text
if TYPE_CHECKING: if TYPE_CHECKING:
from docling_core.types.doc.document import DoclingDocument from docling_core.types.doc.document import DoclingDocument
from haiku.rag.embeddings import EmbedderWrapper
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -169,6 +171,7 @@ async def chunk(
config: AppConfig, config: AppConfig,
docling_document: "DoclingDocument", docling_document: "DoclingDocument",
*, *,
embedder: "EmbedderWrapper",
existing_picture_data: dict[str, bytes] | None = None, existing_picture_data: dict[str, bytes] | None = None,
document_id: str | None = None, document_id: str | None = None,
) -> list[Chunk]: ) -> list[Chunk]:
@ -184,12 +187,11 @@ async def chunk(
path where the docling is loaded from the stored blob. path where the docling is loaded from the stored blob.
""" """
from haiku.rag.chunkers import get_chunker from haiku.rag.chunkers import get_chunker
from haiku.rag.embeddings import get_embedder
chunker = get_chunker(config) chunker = get_chunker(config)
text_chunks = await chunker.chunk(docling_document) text_chunks = await chunker.chunk(docling_document)
if not get_embedder(config).supports_images: if not embedder.supports_images:
for i, c in enumerate(text_chunks): for i, c in enumerate(text_chunks):
c.order = i c.order = i
return text_chunks return text_chunks
@ -277,7 +279,9 @@ def build_picture_chunks(
return chunks return chunks
async def ensure_chunks_embedded(config: AppConfig, chunks: list[Chunk]) -> list[Chunk]: async def ensure_chunks_embedded(
config: AppConfig, chunks: list[Chunk], embedder: "EmbedderWrapper"
) -> list[Chunk]:
"""Ensure all chunks have embeddings, embedding any that don't. """Ensure all chunks have embeddings, embedding any that don't.
Chunks that already have embeddings are passed through unchanged; missing Chunks that already have embeddings are passed through unchanged; missing
@ -290,7 +294,7 @@ async def ensure_chunks_embedded(config: AppConfig, chunks: list[Chunk]) -> list
if not chunks_to_embed: if not chunks_to_embed:
return chunks return chunks
embedded = await embed_chunks(chunks_to_embed, config) embedded = await embed_chunks(chunks_to_embed, embedder, config)
# Build result maintaining original order # Build result maintaining original order
embedded_map = {(c.content, c.order): c for c in embedded} embedded_map = {(c.content, c.order): c for c in embedded}

View file

@ -503,11 +503,11 @@ async def _rebuild_rechunk(
client: "HaikuRAG", documents: list[Document] client: "HaikuRAG", documents: list[Document]
) -> AsyncGenerator[str, None]: ) -> AsyncGenerator[str, None]:
"""Re-chunk and re-embed each document from its stored docling blob.""" """Re-chunk and re-embed each document from its stored docling blob."""
from haiku.rag.embeddings import embed_chunks, get_embedder from haiku.rag.embeddings import embed_chunks
pending_chunks: list[Chunk] = [] pending_chunks: list[Chunk] = []
pending_docs: list[Document] = [] pending_docs: list[Document] = []
embedder = get_embedder(client._config) embedder = client.embedder
async for doc in _hydrate(client, documents): async for doc in _hydrate(client, documents):
assert doc.id is not None assert doc.id is not None
@ -530,7 +530,7 @@ async def _rebuild_rechunk(
existing_picture_data=existing_picture_data, existing_picture_data=existing_picture_data,
document_id=doc.id, document_id=doc.id,
) )
embedded_chunks = await embed_chunks(chunks, client._config) embedded_chunks = await embed_chunks(chunks, embedder, client._config)
# Prepare chunks with document_id and order # Prepare chunks with document_id and order
for order, chunk in enumerate(embedded_chunks): for order, chunk in enumerate(embedded_chunks):
@ -638,7 +638,7 @@ async def _rebuild_descriptions(
cost remains. Idempotent: pictures whose ``meta.description.text`` is cost remains. Idempotent: pictures whose ``meta.description.text`` is
already populated are not re-described. already populated are not re-described.
""" """
from haiku.rag.embeddings import embed_chunks, get_embedder from haiku.rag.embeddings import embed_chunks
if client._config.processing.pictures != "description": if client._config.processing.pictures != "description":
raise ValueError( raise ValueError(
@ -648,7 +648,7 @@ async def _rebuild_descriptions(
pending_chunks: list[Chunk] = [] pending_chunks: list[Chunk] = []
pending_docs: list[Document] = [] pending_docs: list[Document] = []
embedder = get_embedder(client._config) embedder = client.embedder
described_total = 0 described_total = 0
async for doc in _hydrate(client, documents): async for doc in _hydrate(client, documents):
@ -676,7 +676,7 @@ async def _rebuild_descriptions(
existing_picture_data=existing_picture_data, existing_picture_data=existing_picture_data,
document_id=doc.id, document_id=doc.id,
) )
embedded_chunks = await embed_chunks(chunks, client._config) embedded_chunks = await embed_chunks(chunks, embedder, client._config)
for order, chunk in enumerate(embedded_chunks): for order, chunk in enumerate(embedded_chunks):
chunk.document_id = doc.id chunk.document_id = doc.id
@ -710,6 +710,7 @@ async def _rebuild_full(
pending_chunks: list[Chunk] = [] pending_chunks: list[Chunk] = []
pending_docs: list[Document] = [] pending_docs: list[Document] = []
converter = get_converter(client._config) converter = get_converter(client._config)
embedder = client.embedder
for light_doc in documents: for light_doc in documents:
assert light_doc.id is not None assert light_doc.id is not None
@ -751,7 +752,7 @@ async def _rebuild_full(
docling_document = await converter.convert_text(doc.content, format="md") docling_document = await converter.convert_text(doc.content, format="md")
chunks = await client.chunk(docling_document) chunks = await client.chunk(docling_document)
embedded_chunks = await embed_chunks(chunks, client._config) embedded_chunks = await embed_chunks(chunks, embedder, client._config)
doc.set_docling(docling_document) doc.set_docling(docling_document)

View file

@ -55,9 +55,7 @@ async def search(
chunks = [chunk for chunk, _ in raw_results] chunks = [chunk for chunk, _ in raw_results]
chunk_results = await reranker.rerank(query, chunks, top_n=limit) chunk_results = await reranker.rerank(query, chunks, top_n=limit)
else: else:
from haiku.rag.embeddings import get_embedder embedder = client.embedder
embedder = get_embedder(client._config)
if not embedder.supports_images: if not embedder.supports_images:
raise ValueError( raise ValueError(
"Image queries require a multimodal embedder. Configure " "Image queries require a multimodal embedder. Configure "

View file

@ -83,7 +83,7 @@ def contextualize(chunks: list["Chunk"]) -> list[str]:
async def embed_chunks( async def embed_chunks(
chunks: list["Chunk"], config: AppConfig = Config chunks: list["Chunk"], embedder: "EmbedderWrapper", config: AppConfig = Config
) -> list["Chunk"]: ) -> list["Chunk"]:
"""Generate embeddings for chunks, dispatching text vs picture variants. """Generate embeddings for chunks, dispatching text vs picture variants.
@ -97,8 +97,6 @@ async def embed_chunks(
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
embedder = get_embedder(config)
text_chunks: list[Chunk] = [] text_chunks: list[Chunk] = []
picture_chunks: list[Chunk] = [] picture_chunks: list[Chunk] = []
for chunk in chunks: for chunk in chunks:

View file

@ -108,7 +108,9 @@ def create_mcp_server(
return [] return []
# Image-as-query tool, only registered when the configured embedder # Image-as-query tool, only registered when the configured embedder
# supports image embeddings. # supports image embeddings. Probed at server-build time when no Store is
# open, so there is no cached embedder to read; this is the one place
# outside Store that builds one.
from haiku.rag.embeddings import get_embedder from haiku.rag.embeddings import get_embedder
if get_embedder(config).supports_images: if get_embedder(config).supports_images:

View file

@ -108,7 +108,7 @@ async def test_chunking_pipeline(qa_corpus: Dataset, temp_db_path):
# Use client primitives: convert → chunk → embed # Use client primitives: convert → chunk → embed
docling_document = await client.convert(document_text) docling_document = await client.convert(document_text)
chunks = await client.chunk(docling_document) chunks = await client.chunk(docling_document)
embedded_chunks = await embed_chunks(chunks) embedded_chunks = await embed_chunks(chunks, client.embedder)
# Verify chunks were created with embeddings # Verify chunks were created with embeddings
assert len(chunks) > 0 assert len(chunks) > 0

View file

@ -1580,7 +1580,7 @@ async def test_sql_injection_is_blocked_with_escaping(temp_db_path):
def _patch_embed_chunks(monkeypatch): def _patch_embed_chunks(monkeypatch):
async def fake_embed_chunks(chunks, config): async def fake_embed_chunks(chunks, embedder, config):
for chunk in chunks: for chunk in chunks:
chunk.embedding = [0.0] * 2560 chunk.embedding = [0.0] * 2560
return chunks return chunks

View file

@ -14,7 +14,9 @@ async def create_document_with_docling(
): ):
"""Helper to create a document from a DoclingDocument using import_document.""" """Helper to create a document from a DoclingDocument using import_document."""
chunks = await client.chunk(docling_doc) chunks = await client.chunk(docling_doc)
embedded_chunks = await ensure_chunks_embedded(client._config, chunks) embedded_chunks = await ensure_chunks_embedded(
client._config, chunks, client.embedder
)
return await client.import_document( return await client.import_document(
docling_document=docling_doc, docling_document=docling_doc,
chunks=embedded_chunks, chunks=embedded_chunks,

View file

@ -3,8 +3,18 @@ from pathlib import Path
import numpy as np import numpy as np
import pytest import pytest
from haiku.rag.config import AppConfig, EmbeddingModelConfig, EmbeddingsConfig from haiku.rag.config import (
from haiku.rag.embeddings import contextualize, embed_chunks, get_embedder AppConfig,
Config,
EmbeddingModelConfig,
EmbeddingsConfig,
)
from haiku.rag.embeddings import (
EmbedderWrapper,
contextualize,
embed_chunks,
get_embedder,
)
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
@ -105,6 +115,45 @@ def test_contextualize_empty_list():
assert texts == [] assert texts == []
class _StubEmbedder(EmbedderWrapper):
def __init__(self):
super().__init__(embedder=None, vector_dim=8)
self.doc_batches = 0
async def embed_documents(self, texts):
self.doc_batches += 1
return [[0.1] * 8 for _ in texts]
async def test_embed_chunks_uses_provided_embedder(monkeypatch):
"""embed_chunks embeds via the embedder it is given and never builds one."""
import haiku.rag.embeddings as embeddings_mod
def fail(*args, **kwargs):
raise AssertionError("embed_chunks must not build its own embedder")
monkeypatch.setattr(embeddings_mod, "get_embedder", fail)
embedder = _StubEmbedder()
chunks = [
Chunk(id="a", content="alpha", order=0),
Chunk(id="b", content="beta", order=1),
]
embedded = await embed_chunks(chunks, embedder, AppConfig())
assert [c.embedding for c in embedded] == [[0.1] * 8, [0.1] * 8]
assert embedder.doc_batches == 1
async def test_client_embedder_is_store_embedder(temp_db_path):
"""The client exposes the Store's cached embedder rather than its own."""
from haiku.rag.client import HaikuRAG
async with HaikuRAG(temp_db_path, create=True) as client:
assert client.embedder is client.store.embedder
@pytest.mark.vcr() @pytest.mark.vcr()
async def test_embed_chunks_basic(allow_model_requests): async def test_embed_chunks_basic(allow_model_requests):
"""Test that embed_chunks generates embeddings for chunks.""" """Test that embed_chunks generates embeddings for chunks."""
@ -125,7 +174,7 @@ async def test_embed_chunks_basic(allow_model_requests):
), ),
] ]
embedded_chunks = await embed_chunks(chunks) embedded_chunks = await embed_chunks(chunks, get_embedder(Config))
assert len(embedded_chunks) == 2 assert len(embedded_chunks) == 2
# Check that all original fields are preserved # Check that all original fields are preserved
@ -144,7 +193,7 @@ async def test_embed_chunks_basic(allow_model_requests):
async def test_embed_chunks_returns_new_objects(allow_model_requests): async def test_embed_chunks_returns_new_objects(allow_model_requests):
"""Test that embed_chunks returns new Chunk objects (immutable pattern).""" """Test that embed_chunks returns new Chunk objects (immutable pattern)."""
original = Chunk(id="orig", content="Test content.") original = Chunk(id="orig", content="Test content.")
embedded = await embed_chunks([original]) embedded = await embed_chunks([original], get_embedder(Config))
# Original should be unchanged # Original should be unchanged
assert original.embedding is None assert original.embedding is None
@ -156,7 +205,7 @@ async def test_embed_chunks_returns_new_objects(allow_model_requests):
async def test_embed_chunks_empty_list(): async def test_embed_chunks_empty_list():
"""Test that embed_chunks handles empty list.""" """Test that embed_chunks handles empty list."""
result = await embed_chunks([]) result = await embed_chunks([], _StubEmbedder())
assert result == [] assert result == []
@ -176,7 +225,7 @@ async def test_embed_chunks_picture_with_text_only_embedder_raises():
) )
with pytest.raises(ValueError, match="multimodal embedder"): with pytest.raises(ValueError, match="multimodal embedder"):
await embed_chunks([chunk], config) await embed_chunks([chunk], get_embedder(config), config)
async def test_embed_chunks_respects_configured_batch_size(monkeypatch): async def test_embed_chunks_respects_configured_batch_size(monkeypatch):
@ -206,7 +255,7 @@ async def test_embed_chunks_respects_configured_batch_size(monkeypatch):
for i in range(num_chunks) for i in range(num_chunks)
] ]
result = await embed_chunks(chunks, config) result = await embed_chunks(chunks, get_embedder(config), config)
assert len(result) == num_chunks assert len(result) == num_chunks
# 20 chunks / 7 per batch -> 7, 7, 6 # 20 chunks / 7 per batch -> 7, 7, 6
@ -230,7 +279,7 @@ async def test_embed_chunks_preserves_all_fields(allow_model_requests):
document_meta={"author": "Test"}, document_meta={"author": "Test"},
) )
embedded = await embed_chunks([chunk]) embedded = await embed_chunks([chunk], get_embedder(Config))
assert embedded[0].id == "test-id" assert embedded[0].id == "test-id"
assert embedded[0].document_id == "doc-id" assert embedded[0].document_id == "doc-id"

View file

@ -440,10 +440,6 @@ async def test_chunk_interleaves_picture_in_structural_order(monkeypatch):
), ),
] ]
monkeypatch.setattr(
"haiku.rag.embeddings.get_embedder",
lambda *a, **kw: StubMultimodalEmbedder(),
)
monkeypatch.setattr( monkeypatch.setattr(
"haiku.rag.chunkers.get_chunker", lambda *a, **kw: StubChunker() "haiku.rag.chunkers.get_chunker", lambda *a, **kw: StubChunker()
) )
@ -459,7 +455,7 @@ async def test_chunk_interleaves_picture_in_structural_order(monkeypatch):
doc.add_picture(image=ImageRef.from_pil(img, dpi=72)) doc.add_picture(image=ImageRef.from_pil(img, dpi=72))
doc.add_text(label=DocItemLabel.PARAGRAPH, text="C") doc.add_text(label=DocItemLabel.PARAGRAPH, text="C")
chunks = await chunk(AppConfig(), doc) chunks = await chunk(AppConfig(), doc, embedder=StubMultimodalEmbedder())
contents = [c.content for c in chunks] contents = [c.content for c in chunks]
assert contents == ["before", "", "after"], ( assert contents == ["before", "", "after"], (
@ -474,7 +470,7 @@ async def test_chunk_interleaves_picture_in_structural_order(monkeypatch):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_embed_chunks_dispatches_text_vs_picture(monkeypatch): async def test_embed_chunks_dispatches_text_vs_picture():
"""embed_chunks routes text chunks through embed_documents (batched) and """embed_chunks routes text chunks through embed_documents (batched) and
picture chunks through embed_image (one at a time), reassembling picture chunks through embed_image (one at a time), reassembling
in original order.""" in original order."""
@ -498,10 +494,6 @@ async def test_embed_chunks_dispatches_text_vs_picture(monkeypatch):
image_calls.append(image) image_calls.append(image)
return [0.9, 0.8, 0.7, 0.6] return [0.9, 0.8, 0.7, 0.6]
monkeypatch.setattr(
"haiku.rag.embeddings.get_embedder", lambda *a, **kw: StubEmbedder()
)
text_chunk = Chunk(content="hello", order=0) text_chunk = Chunk(content="hello", order=0)
pic_chunk = Chunk( pic_chunk = Chunk(
content="figure 1", content="figure 1",
@ -510,7 +502,9 @@ async def test_embed_chunks_dispatches_text_vs_picture(monkeypatch):
) )
pic_chunk._picture_data = b"PNGBYTES" pic_chunk._picture_data = b"PNGBYTES"
embedded = await embed_chunks([text_chunk, pic_chunk, text_chunk.model_copy()]) embedded = await embed_chunks(
[text_chunk, pic_chunk, text_chunk.model_copy()], StubEmbedder()
)
assert len(embedded) == 3 assert len(embedded) == 3
assert embedded[0].embedding == [0.1, 0.2, 0.3, 0.4] assert embedded[0].embedding == [0.1, 0.2, 0.3, 0.4]
@ -521,9 +515,7 @@ async def test_embed_chunks_dispatches_text_vs_picture(monkeypatch):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_embed_chunks_raises_on_picture_chunks_with_text_only_embedder( async def test_embed_chunks_raises_on_picture_chunks_with_text_only_embedder():
monkeypatch,
):
from haiku.rag.embeddings import EmbedderWrapper, embed_chunks from haiku.rag.embeddings import EmbedderWrapper, embed_chunks
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
@ -534,14 +526,10 @@ async def test_embed_chunks_raises_on_picture_chunks_with_text_only_embedder(
async def embed_documents(self, texts): async def embed_documents(self, texts):
return [[0.0] * 4 for _ in texts] return [[0.0] * 4 for _ in texts]
monkeypatch.setattr(
"haiku.rag.embeddings.get_embedder", lambda *a, **kw: TextOnlyEmbedder()
)
pic_chunk = Chunk(content="x", metadata={"labels": ["picture"]}, order=0) pic_chunk = Chunk(content="x", metadata={"labels": ["picture"]}, order=0)
pic_chunk._picture_data = b"PNG" pic_chunk._picture_data = b"PNG"
with pytest.raises(ValueError, match="multimodal embedder"): with pytest.raises(ValueError, match="multimodal embedder"):
await embed_chunks([pic_chunk]) await embed_chunks([pic_chunk], TextOnlyEmbedder())
@pytest.mark.asyncio @pytest.mark.asyncio
@ -569,7 +557,7 @@ async def test_ingest_emits_picture_chunks_with_multimodal_embedder(
return [0.9] * 4 return [0.9] * 4
monkeypatch.setattr( monkeypatch.setattr(
"haiku.rag.embeddings.get_embedder", "haiku.rag.store.engine.get_embedder",
lambda *a, **kw: StubMultimodalEmbedder(), lambda *a, **kw: StubMultimodalEmbedder(),
) )
@ -585,7 +573,7 @@ async def test_ingest_emits_picture_chunks_with_multimodal_embedder(
async with HaikuRAG(temp_db_path, config=config, create=True) as rag: async with HaikuRAG(temp_db_path, config=config, create=True) as rag:
chunks = await rag.chunk(docling_doc) chunks = await rag.chunk(docling_doc)
embedded = await embed_chunks(chunks, rag._config) embedded = await embed_chunks(chunks, rag.embedder, rag._config)
document = Document(content="x", uri="test://doc") document = Document(content="x", uri="test://doc")
document.set_docling(docling_doc) document.set_docling(docling_doc)

View file

@ -354,7 +354,7 @@ async def test_search_with_bytes_query_uses_multimodal_embedder(
return [0.5, 0.5, 0.5, 0.5] return [0.5, 0.5, 0.5, 0.5]
monkeypatch.setattr( monkeypatch.setattr(
"haiku.rag.embeddings.get_embedder", "haiku.rag.store.engine.get_embedder",
lambda *a, **kw: StubMultimodal(), lambda *a, **kw: StubMultimodal(),
) )
@ -446,7 +446,7 @@ async def test_search_with_pil_image_works_like_bytes(temp_db_path, monkeypatch)
return [0.1] * 4 return [0.1] * 4
monkeypatch.setattr( monkeypatch.setattr(
"haiku.rag.embeddings.get_embedder", "haiku.rag.store.engine.get_embedder",
lambda *a, **kw: StubMultimodal(), lambda *a, **kw: StubMultimodal(),
) )

View file

@ -320,3 +320,27 @@ async def test_auto_vacuum_enabled_triggers_vacuum(temp_db_path, monkeypatch):
assert final_versions <= 2, ( assert final_versions <= 2, (
f"With auto-vacuum enabled, should have minimal versions, got {final_versions}" f"With auto-vacuum enabled, should have minimal versions, got {final_versions}"
) )
async def test_close_suppresses_failing_drain_vacuum(temp_db_path, monkeypatch):
"""A failing final vacuum on close must not raise out of __aexit__,
where it would mask an in-flight exception from the context body."""
client = HaikuRAG(db_path=temp_db_path, create=True)
await client.__aenter__()
calls: list[int] = []
async def boom(*args, **kwargs):
calls.append(1)
raise RuntimeError("vacuum boom")
# A finished task in the set forces the drain branch to run.
task = asyncio.create_task(asyncio.sleep(0))
await task
client._vacuum_tasks.add(task)
monkeypatch.setattr(client.store, "vacuum", boom)
# Must not raise despite the drain vacuum erroring.
await client.__aexit__(None, None, None)
assert calls, "drain vacuum should have been attempted"