Adds the 0.75.0 upgrade, which brings a pre-existing database up to the index
set `_init_tables` now creates. It rewrites no table data, so unlike the earlier
data migrations its cost is the index builds alone, each of which reads the
column it indexes.
`ensure_indexes` ensures an index of the declared *type* covers each declared
column, rather than checking that the column is indexed at all. The distinction
is what makes it safe to run against a database of unknown provenance:
- A wrong-typed index no longer satisfies the check. A BTree on `label` covers
the column while losing the low-cardinality equality lookup the Bitmap is for.
- Nothing is dropped or converted away from. Two index types over one column can
be deliberate, serving different query shapes, so an index this function did
not declare survives even on a column it does. The one thing it overwrites is
an index at LanceDB's default name, `{column}_idx`, which is the name it
creates itself.
- A column already carrying the declared type is skipped, so a database with the
full set migrates instantly rather than re-sorting every indexed column.
- Undeclared columns are untouched, so a vector index on `chunks` survives.
It returns the columns it acted on, because a change is not always visible from
outside: adding a Bitmap beside an existing BTree leaves the column indexed
before and after.
The version bump to 0.75.0 is required, not incidental: `_set_initial_version`
stamps a new database with the installed package version, so a migration
numbered above it would be pending the moment the database was created.
`test_client_update_document_replaces_rows_with_bounded_versions` turns
auto_vacuum off. Indexing `documents` means a background vacuum now has an index
to maintain on that table, so `optimize()` writes a version where it previously
had nothing to do, and it landed inside the window the test measures. The
document update itself is still one version, so the bound stays exact.
2708 lines
102 KiB
Python
2708 lines
102 KiB
Python
import asyncio
|
|
import json
|
|
import tempfile
|
|
import threading
|
|
from pathlib import Path
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
import httpx
|
|
import pytest
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
from haiku.rag.client import HaikuRAG
|
|
from haiku.rag.client.documents import (
|
|
DocumentImport,
|
|
_prepare_document_from_docling,
|
|
_write_fetch_body,
|
|
check_source_accessible,
|
|
)
|
|
from haiku.rag.config import Config
|
|
from haiku.rag.embeddings import EmbedderWrapper
|
|
from haiku.rag.store.compression import decompress_json
|
|
from haiku.rag.store.models.chunk import Chunk
|
|
from haiku.rag.store.models.document import Document
|
|
|
|
|
|
@pytest.fixture(scope="module")
|
|
def vcr_cassette_dir():
|
|
return str(Path(__file__).parent / "cassettes" / "test_client")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_prepare_document_from_docling_runs_off_event_loop_thread(monkeypatch):
|
|
import haiku.rag.client.documents as documents
|
|
|
|
event_loop_thread = threading.current_thread()
|
|
called_from: list[threading.Thread] = []
|
|
|
|
docling_doc = DoclingDocument(name="thread-check")
|
|
docling_doc.add_text(label=DocItemLabel.TEXT, text="Threaded content")
|
|
document = Document(content="")
|
|
original = documents._prepare_document_from_docling_sync
|
|
|
|
def spy(doc, docling):
|
|
called_from.append(threading.current_thread())
|
|
return original(doc, docling)
|
|
|
|
monkeypatch.setattr(documents, "_prepare_document_from_docling_sync", spy)
|
|
|
|
content = await _prepare_document_from_docling(document, docling_doc)
|
|
|
|
assert content == "Threaded content"
|
|
assert document.content == "Threaded content"
|
|
assert document.docling_document is not None
|
|
assert called_from, "Document.set_docling was never called"
|
|
assert called_from[0] is not event_loop_thread, (
|
|
"Document.set_docling ran on the event-loop thread; document prep must "
|
|
"be dispatched via asyncio.to_thread"
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_write_fetch_body_runs_off_event_loop_thread(monkeypatch):
|
|
import haiku.rag.client.documents as documents
|
|
|
|
event_loop_thread = threading.current_thread()
|
|
called_from: list[threading.Thread] = []
|
|
original = documents._write_fetch_body_sync
|
|
|
|
def spy(body, suffix):
|
|
called_from.append(threading.current_thread())
|
|
return original(body, suffix)
|
|
|
|
monkeypatch.setattr(documents, "_write_fetch_body_sync", spy)
|
|
|
|
path = await _write_fetch_body(b"payload", ".bin")
|
|
try:
|
|
assert path.read_bytes() == b"payload"
|
|
finally:
|
|
path.unlink(missing_ok=True)
|
|
|
|
assert called_from, "_write_fetch_body_sync was never called"
|
|
assert called_from[0] is not event_loop_thread, (
|
|
"_write_fetch_body_sync ran on the event-loop thread; fetched body "
|
|
"writes must be dispatched via asyncio.to_thread"
|
|
)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_document_crud(qa_corpus: list[dict[str, str]], temp_db_path):
|
|
"""Test HaikuRAG CRUD operations for documents."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Get test data
|
|
first_doc = qa_corpus[0]
|
|
document_text = first_doc["document_extracted"]
|
|
test_uri = "file:///path/to/test.txt"
|
|
test_metadata = {"source": "test", "topic": "testing"}
|
|
|
|
# Test create_document
|
|
created_doc = await client.create_document(
|
|
content=document_text, uri=test_uri, metadata=test_metadata
|
|
)
|
|
|
|
assert created_doc.id is not None
|
|
# Content is stored as markdown export, check key text is preserved
|
|
assert "Jakarta" in created_doc.content
|
|
assert created_doc.uri == test_uri
|
|
assert created_doc.metadata == test_metadata
|
|
|
|
# Test get_document_by_id
|
|
retrieved_doc = await client.get_document_by_id(created_doc.id)
|
|
assert retrieved_doc is not None
|
|
assert retrieved_doc.id == created_doc.id
|
|
assert "Jakarta" in retrieved_doc.content
|
|
assert retrieved_doc.uri == test_uri
|
|
|
|
# Test get_document_by_uri
|
|
retrieved_by_uri = await client.get_document_by_uri(test_uri)
|
|
assert retrieved_by_uri is not None
|
|
assert retrieved_by_uri.id == created_doc.id
|
|
assert "Jakarta" in retrieved_by_uri.content
|
|
|
|
# Test get_document_by_uri with non-existent URI
|
|
non_existent = await client.get_document_by_uri("file:///non/existent.txt")
|
|
assert non_existent is None
|
|
|
|
# Test update_document
|
|
updated_doc = await client.update_document(
|
|
document_id=retrieved_doc.id,
|
|
content="Updated content",
|
|
)
|
|
assert updated_doc.content == "Updated content"
|
|
|
|
# Test list_documents
|
|
all_docs = await client.list_documents()
|
|
assert len(all_docs) == 1
|
|
assert all_docs[0].id == created_doc.id
|
|
|
|
# Test list_documents with pagination
|
|
limited_docs = await client.list_documents(limit=10, offset=0)
|
|
assert len(limited_docs) == 1
|
|
|
|
# Test delete_document
|
|
deleted = await client.delete_document(created_doc.id)
|
|
assert deleted is True
|
|
|
|
# Verify document is gone
|
|
retrieved_doc = await client.get_document_by_id(created_doc.id)
|
|
assert retrieved_doc is None
|
|
|
|
# Test delete non-existent document
|
|
deleted_again = await client.delete_document(created_doc.id)
|
|
assert deleted_again is False
|
|
|
|
|
|
async def test_client_resolve_document(temp_db_path):
|
|
"""Test resolve_document finds documents by ID, title, or URI."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Insert document directly via repository (no embeddings needed)
|
|
doc = Document(
|
|
content="Test content",
|
|
uri="test://resolve-test",
|
|
title="Resolve Test Doc",
|
|
)
|
|
doc = await client.document_repository.create(doc)
|
|
|
|
# Resolve by ID
|
|
by_id = await client.resolve_document(doc.id)
|
|
assert by_id is not None
|
|
assert by_id.id == doc.id
|
|
|
|
# Resolve by title
|
|
by_title = await client.resolve_document("Resolve Test Doc")
|
|
assert by_title is not None
|
|
assert by_title.id == doc.id
|
|
|
|
# Resolve by URI
|
|
by_uri = await client.resolve_document("test://resolve-test")
|
|
assert by_uri is not None
|
|
assert by_uri.id == doc.id
|
|
|
|
# Not found returns None
|
|
not_found = await client.resolve_document("nonexistent")
|
|
assert not_found is None
|
|
|
|
# SQL injection is escaped
|
|
injection = "x' OR title LIKE '%"
|
|
injected = await client.resolve_document(injection)
|
|
assert injected is None
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_update_document(qa_corpus: list[dict[str, str]], temp_db_path):
|
|
"""Test updating document with individual parameters."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Get test data
|
|
first_doc = qa_corpus[0]
|
|
document_text = first_doc["document_extracted"]
|
|
test_uri = "file:///path/to/test.txt"
|
|
test_metadata = {"source": "test", "topic": "testing"}
|
|
|
|
# Create a document
|
|
created_doc = await client.create_document(
|
|
content=document_text,
|
|
uri=test_uri,
|
|
title="Original Title",
|
|
metadata=test_metadata,
|
|
)
|
|
assert created_doc.id is not None
|
|
original_id = created_doc.id
|
|
|
|
# Test updating only content
|
|
updated_doc = await client.update_document(
|
|
document_id=original_id, content="Updated content only"
|
|
)
|
|
assert updated_doc.id == original_id
|
|
assert updated_doc.content == "Updated content only"
|
|
assert updated_doc.title == "Original Title"
|
|
assert updated_doc.uri == test_uri
|
|
|
|
# Test updating only metadata
|
|
new_metadata = {"source": "updated", "version": "2.0"}
|
|
updated_doc = await client.update_document(
|
|
document_id=original_id, metadata=new_metadata
|
|
)
|
|
assert updated_doc.metadata == new_metadata
|
|
assert (
|
|
updated_doc.content == "Updated content only"
|
|
) # Should keep previous update
|
|
|
|
# Test updating only title
|
|
updated_doc = await client.update_document(
|
|
document_id=original_id, title="New Title"
|
|
)
|
|
assert updated_doc.title == "New Title"
|
|
assert updated_doc.content == "Updated content only"
|
|
assert updated_doc.metadata == new_metadata
|
|
|
|
# Test updating multiple fields at once
|
|
custom_chunks = [
|
|
Chunk(content="Custom chunk 1", order=0),
|
|
Chunk(content="Custom chunk 2", order=1),
|
|
]
|
|
updated_doc = await client.update_document(
|
|
document_id=original_id,
|
|
content="Content with custom chunks",
|
|
title="Final Title",
|
|
metadata={"final": "true"},
|
|
chunks=custom_chunks,
|
|
)
|
|
assert updated_doc.id == original_id
|
|
assert updated_doc.content == "Content with custom chunks"
|
|
assert updated_doc.title == "Final Title"
|
|
assert updated_doc.metadata == {"final": "true"}
|
|
|
|
# Verify the custom chunks were created
|
|
doc_chunks = await client.chunk_repository.get_by_document_id(original_id)
|
|
assert len(doc_chunks) == 2
|
|
assert doc_chunks[0].content == "Custom chunk 1"
|
|
assert doc_chunks[1].content == "Custom chunk 2"
|
|
|
|
# Test updating only the uri
|
|
new_uri = "file:///path/to/new.txt"
|
|
updated_doc = await client.update_document(document_id=original_id, uri=new_uri)
|
|
assert updated_doc.uri == new_uri
|
|
refetched = await client.get_document_by_id(original_id)
|
|
assert refetched is not None
|
|
assert refetched.uri == new_uri
|
|
assert (await client.get_document_by_uri(new_uri)) is not None
|
|
assert (await client.get_document_by_uri(test_uri)) is None
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_source(temp_db_path):
|
|
"""Test creating a document from a file source."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
test_content = "This is test content from a file."
|
|
temp_path = Path(temp_dir) / "test.txt"
|
|
temp_path.write_text(test_content)
|
|
|
|
# Test create_document_from_source with Path
|
|
doc = await client.create_document_from_source(source=temp_path)
|
|
assert isinstance(doc, Document)
|
|
|
|
assert doc.id is not None
|
|
assert doc.content == test_content
|
|
assert doc.uri == temp_path.as_uri()
|
|
assert "content_type" in doc.metadata
|
|
assert "md5" in doc.metadata
|
|
assert doc.metadata["content_type"] == "text/plain"
|
|
|
|
# Test create_document_from_source with string path
|
|
doc2 = await client.create_document_from_source(source=str(temp_path))
|
|
assert isinstance(doc2, Document)
|
|
|
|
assert doc2.id is not None
|
|
assert doc2.content == test_content
|
|
assert doc2.uri == temp_path.as_uri()
|
|
assert "content_type" in doc2.metadata
|
|
assert "md5" in doc2.metadata
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_update_title_noop_behavior(temp_db_path):
|
|
"""When content is unchanged, updating title should update document without re-chunking."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test_update_title.txt"
|
|
temp_path.write_text("Original content")
|
|
|
|
doc1 = await client.create_document_from_source(temp_path, title="Title A")
|
|
assert isinstance(doc1, Document)
|
|
assert doc1.id is not None
|
|
assert doc1.title == "Title A"
|
|
|
|
# Re-add with same content but new title
|
|
doc2 = await client.create_document_from_source(temp_path, title="Title B")
|
|
assert isinstance(doc2, Document)
|
|
assert doc2.id == doc1.id
|
|
# Fetch and verify title updated
|
|
got = await client.get_document_by_id(doc1.id)
|
|
assert got is not None
|
|
assert got.title == "Title B"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_source_with_uri_override(temp_db_path):
|
|
"""A `uri` override is honored as the canonical document identifier."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "2412.06611v2.pdf-ish.txt"
|
|
temp_path.write_text("Synthetic content for URI-override test.")
|
|
|
|
doc = await client.create_document_from_source(
|
|
source=temp_path, uri="2412.06611v2"
|
|
)
|
|
assert isinstance(doc, Document)
|
|
assert doc.uri == "2412.06611v2"
|
|
assert doc.uri != temp_path.as_uri()
|
|
|
|
# The override URI is the lookup key for subsequent reads.
|
|
looked_up = await client.get_document_by_uri("2412.06611v2")
|
|
assert looked_up is not None
|
|
assert looked_up.id == doc.id
|
|
|
|
# The original file URI is NOT a key.
|
|
assert await client.get_document_by_uri(temp_path.as_uri()) is None
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_source_uri_override_dedupes(temp_db_path):
|
|
"""Re-creating from the same source with the same override is a no-op."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "doc.txt"
|
|
temp_path.write_text("Stable content for dedup test.")
|
|
|
|
doc1 = await client.create_document_from_source(
|
|
source=temp_path, uri="paper-id-1"
|
|
)
|
|
doc2 = await client.create_document_from_source(
|
|
source=temp_path, uri="paper-id-1"
|
|
)
|
|
assert isinstance(doc1, Document) and isinstance(doc2, Document)
|
|
assert doc1.id == doc2.id
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_source_uri_override_rejected_for_dir(
|
|
temp_db_path,
|
|
):
|
|
"""Directory sources reject the `uri` override (would collide on multiple files)."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
(Path(temp_dir) / "a.txt").write_text("a")
|
|
(Path(temp_dir) / "b.txt").write_text("b")
|
|
|
|
with pytest.raises(ValueError, match="directory sources"):
|
|
await client.create_document_from_source(
|
|
source=Path(temp_dir), uri="some-uri"
|
|
)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_source_unsupported(temp_db_path):
|
|
"""Test creating a document from an unsupported file type."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a temporary file with unsupported extension
|
|
with tempfile.NamedTemporaryFile(
|
|
mode="w", suffix=".unsupported", delete=False
|
|
) as f:
|
|
f.write("content")
|
|
temp_path = Path(f.name)
|
|
|
|
# Should raise ValueError for unsupported extension
|
|
with pytest.raises(ValueError, match="Unsupported file extension"):
|
|
await client.create_document_from_source(temp_path)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_source_nonexistent(temp_db_path):
|
|
"""Test creating a document from a non-existent file."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
non_existent_path = Path("/non/existent/file.txt")
|
|
|
|
# Should raise ValueError when file doesn't exist
|
|
with pytest.raises(ValueError, match="File does not exist"):
|
|
await client.create_document_from_source(non_existent_path)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_directory(temp_db_path):
|
|
"""Test creating documents from a directory recursively."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
test_dir = Path(temp_dir) / "test_docs"
|
|
test_dir.mkdir()
|
|
|
|
(test_dir / "doc1.txt").write_text("Content of doc1")
|
|
(test_dir / "doc2.md").write_text("# Content of doc2")
|
|
|
|
subdir = test_dir / "subdir"
|
|
subdir.mkdir()
|
|
(subdir / "doc3.py").write_text("print('hello')")
|
|
|
|
(test_dir / "unsupported.xyz").write_text("unsupported file")
|
|
|
|
result = await client.create_document_from_source(test_dir)
|
|
|
|
assert isinstance(result, list)
|
|
assert len(result) == 3
|
|
|
|
for doc in result:
|
|
assert doc.id is not None
|
|
assert doc.uri is not None
|
|
assert "md5" in doc.metadata
|
|
assert "content_type" in doc.metadata
|
|
|
|
uris = [doc.uri for doc in result if doc.uri]
|
|
assert any("doc1.txt" in uri for uri in uris)
|
|
assert any("doc2.md" in uri for uri in uris)
|
|
assert any("doc3.py" in uri for uri in uris)
|
|
assert not any("unsupported.xyz" in uri for uri in uris)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_url(temp_db_path):
|
|
"""Test creating a document from a URL."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Mock the HTTP response
|
|
mock_response = AsyncMock()
|
|
mock_response.content = b"<html><body><h1>Test Page</h1><p>This is test content from a webpage.</p></body></html>"
|
|
mock_response.headers = {"content-type": "text/html"}
|
|
mock_response.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_response):
|
|
doc = await client.create_document_from_source(
|
|
source="https://example.com/test.html", metadata={"source_type": "web"}
|
|
)
|
|
assert isinstance(doc, Document)
|
|
|
|
assert doc.id is not None
|
|
assert "Test Page" in doc.content
|
|
assert "test content" in doc.content
|
|
assert doc.uri == "https://example.com/test.html"
|
|
assert doc.metadata["source_type"] == "web"
|
|
assert "content_type" in doc.metadata
|
|
assert "md5" in doc.metadata
|
|
assert doc.metadata["content_type"] == "text/html"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_url_with_different_content_types(
|
|
temp_db_path,
|
|
):
|
|
"""Test creating documents from URLs with different content types."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Test JSON content
|
|
mock_json_response = AsyncMock()
|
|
mock_json_response.content = (
|
|
b'{"title": "Test JSON", "content": "This is JSON content"}'
|
|
)
|
|
mock_json_response.headers = {"content-type": "application/json"}
|
|
mock_json_response.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_json_response):
|
|
doc = await client.create_document_from_source(
|
|
"https://api.example.com/data.json"
|
|
)
|
|
assert isinstance(doc, Document)
|
|
|
|
assert doc.id is not None
|
|
assert "Test JSON" in doc.content
|
|
assert doc.uri == "https://api.example.com/data.json"
|
|
assert "content_type" in doc.metadata
|
|
assert "md5" in doc.metadata
|
|
assert doc.metadata["content_type"] == "application/json"
|
|
|
|
# Test plain text content
|
|
mock_text_response = AsyncMock()
|
|
mock_text_response.content = b"This is plain text content from a URL."
|
|
mock_text_response.headers = {"content-type": "text/plain"}
|
|
mock_text_response.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_text_response):
|
|
doc = await client.create_document_from_source(
|
|
"https://example.com/readme.txt"
|
|
)
|
|
assert isinstance(doc, Document)
|
|
|
|
assert doc.id is not None
|
|
assert doc.content == "This is plain text content from a URL."
|
|
assert doc.uri == "https://example.com/readme.txt"
|
|
assert "content_type" in doc.metadata
|
|
assert "md5" in doc.metadata
|
|
assert doc.metadata["content_type"] == "text/plain"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_url_unsupported_content(temp_db_path):
|
|
"""Test creating a document from URL with unsupported content type."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Mock response with unsupported content type
|
|
mock_response = AsyncMock()
|
|
mock_response.content = b"binary content"
|
|
mock_response.headers = {"content-type": "application/octet-stream"}
|
|
mock_response.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_response):
|
|
with pytest.raises(ValueError, match="Unsupported content type"):
|
|
await client.create_document_from_source(
|
|
"https://example.com/binary.bin"
|
|
)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_url_http_error(temp_db_path):
|
|
"""Test handling HTTP errors when creating document from URL."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with patch("httpx.AsyncClient.get") as mock_get:
|
|
mock_get.side_effect = httpx.HTTPStatusError(
|
|
"404 Not Found",
|
|
request=httpx.Request("GET", "https://example.com/notfound.html"),
|
|
response=httpx.Response(404),
|
|
)
|
|
|
|
with pytest.raises(httpx.HTTPStatusError):
|
|
await client.create_document_from_source(
|
|
"https://example.com/notfound.html"
|
|
)
|
|
|
|
|
|
def test_get_extension_from_content_type_or_url():
|
|
"""Test the helper function for determining file extensions."""
|
|
from haiku.rag.client.processing import get_extension_from_content_type_or_url
|
|
|
|
# Content type mappings
|
|
assert get_extension_from_content_type_or_url("", "text/html") == ".html"
|
|
assert get_extension_from_content_type_or_url("", "application/pdf") == ".pdf"
|
|
assert get_extension_from_content_type_or_url("", "text/plain") == ".txt"
|
|
|
|
# URL extension detection
|
|
assert (
|
|
get_extension_from_content_type_or_url("https://example.com/doc.pdf", "")
|
|
== ".pdf"
|
|
)
|
|
assert (
|
|
get_extension_from_content_type_or_url("https://example.com/data.json", "")
|
|
== ".json"
|
|
)
|
|
|
|
# Default fallback
|
|
assert get_extension_from_content_type_or_url("https://example.com/", "") == ".html"
|
|
|
|
# Content type priority over URL extension
|
|
assert (
|
|
get_extension_from_content_type_or_url(
|
|
"https://example.com/file.txt", "application/pdf"
|
|
)
|
|
== ".pdf"
|
|
)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_metadata_content_type_and_md5(temp_db_path):
|
|
"""Test that content_type and md5 metadata are correctly set."""
|
|
import hashlib
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a temporary file with known content
|
|
test_content = "Test content for MD5 calculation."
|
|
expected_md5 = hashlib.md5(test_content.encode()).hexdigest()
|
|
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test.txt"
|
|
temp_path.write_text(test_content)
|
|
|
|
doc = await client.create_document_from_source(temp_path)
|
|
assert isinstance(doc, Document)
|
|
|
|
assert doc.metadata["content_type"] == "text/plain"
|
|
assert doc.metadata["md5"] == expected_md5
|
|
|
|
mock_response = AsyncMock()
|
|
mock_response.content = test_content.encode()
|
|
mock_response.headers = {"content-type": "text/plain"}
|
|
mock_response.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_response):
|
|
url_doc = await client.create_document_from_source(
|
|
"https://example.com/test.txt"
|
|
)
|
|
assert isinstance(url_doc, Document)
|
|
|
|
assert url_doc.metadata["content_type"] == "text/plain"
|
|
assert url_doc.metadata["md5"] == expected_md5
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_update_no_op_behavior(temp_db_path):
|
|
"""Test create/update/no-op behavior based on MD5 changes."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a temporary file
|
|
test_content = "Original content for testing."
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test.txt"
|
|
temp_path.write_text(test_content)
|
|
|
|
# First call - should create new document
|
|
doc1 = await client.create_document_from_source(temp_path)
|
|
assert isinstance(doc1, Document)
|
|
assert doc1.id is not None
|
|
assert doc1.content == test_content
|
|
original_id = doc1.id
|
|
original_updated_at = doc1.updated_at
|
|
|
|
# Second call with same content - should return existing document (no-op)
|
|
doc2 = await client.create_document_from_source(temp_path)
|
|
assert isinstance(doc2, Document)
|
|
assert doc2.id == original_id # Same document
|
|
assert doc2.content == test_content
|
|
assert doc2.updated_at == original_updated_at # No-op leaves it untouched
|
|
|
|
# Modify file content
|
|
updated_content = "Updated content for testing."
|
|
temp_path.write_text(updated_content)
|
|
|
|
# Third call with changed content - should update existing document
|
|
doc3 = await client.create_document_from_source(temp_path)
|
|
assert isinstance(doc3, Document)
|
|
assert doc3.id == original_id # Same document ID
|
|
assert doc3.content == updated_content # Updated content
|
|
|
|
# Verify the document was actually updated in database
|
|
retrieved_doc = await client.get_document_by_id(original_id)
|
|
assert retrieved_doc is not None
|
|
assert retrieved_doc.content == updated_content
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_url_create_update_no_op_behavior(temp_db_path):
|
|
"""Test create/update/no-op behavior for URLs based on MD5 changes."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
url = "https://example.com/test.txt"
|
|
original_content = b"Original URL content"
|
|
updated_content = b"Updated URL content"
|
|
|
|
# Mock first response
|
|
mock_response1 = AsyncMock()
|
|
mock_response1.content = original_content
|
|
mock_response1.headers = {"content-type": "text/plain"}
|
|
mock_response1.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_response1):
|
|
# First call - should create new document
|
|
doc1 = await client.create_document_from_source(url)
|
|
assert isinstance(doc1, Document)
|
|
assert doc1.id is not None
|
|
original_id = doc1.id
|
|
|
|
# Second call with same content - should return existing document (no-op)
|
|
doc2 = await client.create_document_from_source(url)
|
|
assert isinstance(doc2, Document)
|
|
assert doc2.id == original_id # Same document
|
|
|
|
mock_response2 = AsyncMock()
|
|
mock_response2.content = updated_content
|
|
mock_response2.headers = {"content-type": "text/plain"}
|
|
mock_response2.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_response2):
|
|
# Third call with changed content - should update existing document
|
|
doc3 = await client.create_document_from_source(url)
|
|
assert isinstance(doc3, Document)
|
|
assert doc3.id == original_id # Same document ID
|
|
assert doc3.content == updated_content.decode() # Updated content
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_search(temp_db_path):
|
|
"""Test HaikuRAG search functionality."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Add multiple documents to search from
|
|
doc1_text = "Python is a high-level programming language known for its simplicity and readability."
|
|
doc2_text = "Machine learning algorithms help computers learn patterns from data without explicit programming."
|
|
doc3_text = "Data science combines statistics, programming, and domain expertise to extract insights."
|
|
|
|
# Create documents
|
|
doc1 = await client.create_document(
|
|
content=doc1_text, uri="doc1.txt", metadata={"topic": "python"}
|
|
)
|
|
doc2 = await client.create_document(
|
|
content=doc2_text, uri="doc2.txt", metadata={"topic": "ml"}
|
|
)
|
|
await client.create_document(
|
|
content=doc3_text, uri="doc3.txt", metadata={"topic": "data_science"}
|
|
)
|
|
|
|
# Test search with keyword that should match doc1
|
|
results = await client.search("Python programming", limit=3)
|
|
|
|
assert len(results) > 0
|
|
# Verify results are SearchResult objects with expected fields
|
|
first_result = results[0]
|
|
assert first_result.content
|
|
assert first_result.score >= 0
|
|
assert first_result.document_id == doc1.id
|
|
|
|
# Test search with different query
|
|
ml_results = await client.search("machine learning algorithms", limit=2)
|
|
assert len(ml_results) > 0
|
|
|
|
# Verify first result is from the machine learning document (doc2)
|
|
first_ml_result = ml_results[0]
|
|
assert first_ml_result.document_id == doc2.id
|
|
|
|
# Test search with limit parameter
|
|
limited_results = await client.search("programming", limit=1)
|
|
assert len(limited_results) <= 1
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_async_context_manager(temp_db_path):
|
|
"""Test HaikuRAG as async context manager."""
|
|
|
|
# Test that context manager works and auto-closes
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a document to ensure the client works
|
|
doc = await client.create_document(
|
|
content="Test content for context manager",
|
|
uri="test://context",
|
|
metadata={"test": "context_manager"},
|
|
)
|
|
|
|
assert doc.id is not None
|
|
assert doc.content == "Test content for context manager"
|
|
|
|
# Test search works within context
|
|
results = await client.search("Test content", limit=1)
|
|
assert len(results) > 0
|
|
|
|
# Context manager should have automatically closed the connection
|
|
# We can't easily test that the connection is closed without accessing internals,
|
|
# but the test passing means the context manager methods work correctly
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_import_document_with_custom_chunks(temp_db_path):
|
|
"""Test importing a document with pre-created chunks."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a DoclingDocument
|
|
docling_doc = DoclingDocument(name="test")
|
|
docling_doc.add_text(label=DocItemLabel.TEXT, text="Full document content")
|
|
|
|
# Create some custom chunks with and without embeddings
|
|
chunks = [
|
|
Chunk(
|
|
content="This is the first chunk",
|
|
metadata={"custom": "metadata1"},
|
|
order=0,
|
|
),
|
|
Chunk(
|
|
content="This is the second chunk",
|
|
metadata={"custom": "metadata2"},
|
|
embedding=[0.1] * Config.embeddings.model.vector_dim,
|
|
order=1,
|
|
), # With embedding
|
|
Chunk(
|
|
content="This is the third chunk",
|
|
metadata={"custom": "metadata3"},
|
|
order=2,
|
|
),
|
|
]
|
|
|
|
# Import document with custom chunks
|
|
document = await client.import_document(
|
|
docling_document=docling_doc, chunks=chunks
|
|
)
|
|
|
|
assert document.id is not None
|
|
assert "Full document content" in document.content
|
|
|
|
# Verify the chunks were created correctly
|
|
doc_chunks = await client.chunk_repository.get_by_document_id(document.id)
|
|
assert len(doc_chunks) == 3
|
|
|
|
# Check chunks have correct content, document_id, and order from list position
|
|
for i, chunk in enumerate(doc_chunks):
|
|
assert chunk.document_id == document.id
|
|
assert chunk.content == chunks[i].content
|
|
assert chunk.order == i # Order should be set from list position
|
|
assert (
|
|
chunk.metadata["custom"] == f"metadata{i + 1}"
|
|
) # Original metadata preserved
|
|
|
|
|
|
def _docling_doc(name: str, text: str):
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
doc = DoclingDocument(name=name)
|
|
doc.add_text(label=DocItemLabel.TEXT, text=text)
|
|
return doc
|
|
|
|
|
|
def _import(name: str, text: str, **overrides) -> "DocumentImport":
|
|
"""Build a DocumentImport with one pre-embedded chunk (no embedder call)."""
|
|
dim = Config.embeddings.model.vector_dim
|
|
chunk = Chunk(content=text, embedding=[0.1] * dim, order=0)
|
|
return DocumentImport(
|
|
docling_document=_docling_doc(name, text),
|
|
chunks=[chunk],
|
|
**overrides,
|
|
)
|
|
|
|
|
|
async def test_client_import_documents_single_version_per_table(temp_db_path):
|
|
"""import_documents writes documents/chunks/document_items once for the
|
|
whole batch (issue #287)."""
|
|
config = Config.model_copy(deep=True)
|
|
config.storage.auto_vacuum = False
|
|
|
|
async with HaikuRAG(temp_db_path, config=config, create=True) as client:
|
|
imports = [
|
|
_import("a", "Alpha document body", uri="mem://a", title="Alpha"),
|
|
_import("b", "Beta document body", uri="mem://b", title="Beta"),
|
|
_import("c", "Gamma document body", uri="mem://c", title="Gamma"),
|
|
]
|
|
|
|
before = await client.store.current_table_versions()
|
|
docs = await client.import_documents(imports)
|
|
after = await client.store.current_table_versions()
|
|
|
|
assert [d.title for d in docs] == ["Alpha", "Beta", "Gamma"]
|
|
assert all(d.id is not None for d in docs)
|
|
assert len({d.id for d in docs}) == 3
|
|
|
|
for table in ("documents", "chunks", "document_items"):
|
|
assert after[table] - before[table] == 1, table
|
|
|
|
for doc, expected in zip(docs, ("Alpha", "Beta", "Gamma")):
|
|
assert doc.id is not None
|
|
stored = await client.get_document_by_id(doc.id)
|
|
assert stored is not None and stored.title == expected
|
|
chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(chunks) == 1 and chunks[0].document_id == doc.id
|
|
items = await client.document_item_repository.get_all_items(doc.id)
|
|
assert len(items) >= 1
|
|
assert all(i.document_id == doc.id for i in items)
|
|
|
|
|
|
async def test_client_import_documents_rolls_back_on_failure(temp_db_path):
|
|
"""A failure mid-batch restores all tables: nothing is persisted."""
|
|
dim = Config.embeddings.model.vector_dim
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
good = _import("good", "Good document body", uri="mem://good")
|
|
bad = DocumentImport(
|
|
docling_document=_docling_doc("bad", "Bad document body"),
|
|
chunks=[Chunk(content="bad", embedding=[0.1] * (dim + 1), order=0)],
|
|
uri="mem://bad",
|
|
)
|
|
|
|
with pytest.raises(Exception):
|
|
await client.import_documents([good, bad])
|
|
|
|
assert await client.count_documents() == 0
|
|
|
|
|
|
async def test_client_import_documents_empty(temp_db_path):
|
|
"""import_documents([]) returns [] and bumps no versions."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
before = await client.store.current_table_versions()
|
|
result = await client.import_documents([])
|
|
after = await client.store.current_table_versions()
|
|
|
|
assert result == []
|
|
assert after == before
|
|
|
|
|
|
class _CountingEmbedder(EmbedderWrapper):
|
|
def __init__(self, vector_dim: int):
|
|
super().__init__(None, vector_dim)
|
|
self.batches: list[int] = []
|
|
|
|
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
|
self.batches.append(len(texts))
|
|
return [[0.1] * self.vector_dim for _ in texts]
|
|
|
|
|
|
async def test_client_import_documents_batches_embeddings(temp_db_path):
|
|
"""Chunks missing embeddings are embedded in one pass across the whole
|
|
batch, not one embedder call per document. Duplicate chunk texts across
|
|
documents keep their per-document embeddings."""
|
|
dim = Config.embeddings.model.vector_dim
|
|
embedder = _CountingEmbedder(dim)
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
client.store.embedder = embedder
|
|
imports = [
|
|
DocumentImport(
|
|
docling_document=_docling_doc(name, text),
|
|
chunks=[Chunk(content=text, order=0)],
|
|
uri=f"mem://{name}",
|
|
title=name,
|
|
)
|
|
for name, text in (
|
|
("a", "Alpha document body"),
|
|
("b", "Beta document body"),
|
|
("c", "Alpha document body"),
|
|
)
|
|
]
|
|
|
|
docs = await client.import_documents(imports)
|
|
|
|
assert embedder.batches == [3]
|
|
rows = await (
|
|
client.store.chunks_table.query()
|
|
.select(["document_id", "vector"])
|
|
.to_list()
|
|
)
|
|
assert {row["document_id"] for row in rows} == {doc.id for doc in docs}
|
|
assert all(len(row["vector"]) == dim for row in rows)
|
|
|
|
|
|
async def test_client_import_documents_mixed_embeddings(temp_db_path):
|
|
"""Pre-embedded chunks keep their vectors; only the unembedded ones go
|
|
through the embedder, in one batch."""
|
|
dim = Config.embeddings.model.vector_dim
|
|
embedder = _CountingEmbedder(dim)
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
client.store.embedder = embedder
|
|
pre_embedded = DocumentImport(
|
|
docling_document=_docling_doc("b", "Beta document body"),
|
|
chunks=[
|
|
Chunk(content="Beta document body", embedding=[0.5] * dim, order=0)
|
|
],
|
|
uri="mem://b",
|
|
title="b",
|
|
)
|
|
unembedded = [
|
|
DocumentImport(
|
|
docling_document=_docling_doc(name, text),
|
|
chunks=[Chunk(content=text, order=0)],
|
|
uri=f"mem://{name}",
|
|
title=name,
|
|
)
|
|
for name, text in (("a", "Alpha document body"), ("c", "Gamma body"))
|
|
]
|
|
|
|
docs = await client.import_documents(
|
|
[unembedded[0], pre_embedded, unembedded[1]]
|
|
)
|
|
|
|
assert embedder.batches == [2]
|
|
by_uri = {doc.uri: doc.id for doc in docs}
|
|
rows = await (
|
|
client.store.chunks_table.query()
|
|
.select(["document_id", "vector"])
|
|
.to_list()
|
|
)
|
|
vectors = {row["document_id"]: list(row["vector"]) for row in rows}
|
|
assert vectors[by_uri["mem://b"]] == pytest.approx([0.5] * dim)
|
|
assert vectors[by_uri["mem://a"]] == pytest.approx([0.1] * dim)
|
|
assert vectors[by_uri["mem://c"]] == pytest.approx([0.1] * dim)
|
|
|
|
|
|
async def test_client_update_document_replaces_rows_with_bounded_versions(
|
|
temp_db_path,
|
|
):
|
|
"""Updating one document should replace stale rows with bounded versions.
|
|
|
|
auto_vacuum is off: a background vacuum optimizes every table, which writes
|
|
versions of its own and would land inside the window being measured.
|
|
"""
|
|
dim = Config.embeddings.model.vector_dim
|
|
config = Config.model_copy(deep=True)
|
|
config.storage.auto_vacuum = False
|
|
|
|
async with HaikuRAG(temp_db_path, config=config, create=True) as client:
|
|
created = await client.import_document(
|
|
_docling_doc("original", "Original body"),
|
|
[Chunk(content="Original body", embedding=[0.1] * dim, order=0)],
|
|
uri="mem://replace",
|
|
title="Replace",
|
|
)
|
|
assert created.id is not None
|
|
|
|
updated_docling = _docling_doc("updated", "Updated body")
|
|
updated_chunks = [
|
|
Chunk(content="Updated body A", embedding=[0.2] * dim, order=0),
|
|
Chunk(content="Updated body B", embedding=[0.3] * dim, order=1),
|
|
]
|
|
|
|
before = await client.store.current_table_versions()
|
|
updated = await client.update_document(
|
|
created.id,
|
|
docling_document=updated_docling,
|
|
chunks=updated_chunks,
|
|
)
|
|
after = await client.store.current_table_versions()
|
|
|
|
assert updated.id == created.id
|
|
assert after["documents"] - before["documents"] == 1
|
|
# Indexed LanceDB tables record one additional physical version for
|
|
# merge replacement in 0.30.x.
|
|
assert after["chunks"] - before["chunks"] <= 2
|
|
assert after["document_items"] - before["document_items"] <= 2
|
|
|
|
stored_chunks = await client.chunk_repository.get_by_document_id(created.id)
|
|
assert [chunk.content for chunk in stored_chunks] == [
|
|
"Updated body A",
|
|
"Updated body B",
|
|
]
|
|
stored_items = await client.document_item_repository.get_all_items(created.id)
|
|
assert len(stored_items) == 1
|
|
assert stored_items[0].text == "Updated body"
|
|
|
|
|
|
async def test_metadata_only_update_does_not_advance_documents_table(temp_db_path):
|
|
"""Metadata/title-only updates must not rewrite the heavy documents row.
|
|
|
|
This is the blob-bloat fix: source_revision rolling on every ingester sweep
|
|
used to rewrite the multi-MB docling row each time. Mutable attributes now
|
|
live in document_meta, so the documents table version must stay frozen while
|
|
only metadata/title change — and reads must still hydrate the full Document.
|
|
"""
|
|
dim = Config.embeddings.model.vector_dim
|
|
config = Config.model_copy(deep=True)
|
|
config.storage.auto_vacuum = False
|
|
|
|
async with HaikuRAG(temp_db_path, config=config, create=True) as client:
|
|
created = await client.import_document(
|
|
_docling_doc("doc", "Body text"),
|
|
[Chunk(content="Body text", embedding=[0.1] * dim, order=0)],
|
|
uri="mem://meta-bloat",
|
|
title="Original",
|
|
metadata={"source_revision": "rev-0"},
|
|
)
|
|
assert created.id is not None
|
|
|
|
docs_v0 = await client.store.documents_table.version()
|
|
meta_v0 = await client.store.document_meta_table.version()
|
|
|
|
for i in range(1, 6):
|
|
await client.update_document(
|
|
created.id,
|
|
metadata={"source_revision": f"rev-{i}"},
|
|
title=f"Title {i}",
|
|
)
|
|
|
|
# The heavy documents table must not advance on metadata-only updates.
|
|
assert await client.store.documents_table.version() == docs_v0
|
|
# The light document_meta table absorbs the updates.
|
|
assert await client.store.document_meta_table.version() > meta_v0
|
|
|
|
# Reads still hydrate content and the mutable attributes together.
|
|
fetched = await client.get_document_by_id(created.id)
|
|
assert fetched is not None
|
|
assert fetched.metadata["source_revision"] == "rev-5"
|
|
assert fetched.title == "Title 5"
|
|
assert fetched.content == "Body text"
|
|
|
|
# And the untouched docling blob is still there.
|
|
docling = await client.document_repository.get_docling_data(created.id)
|
|
assert docling is not None
|
|
assert docling.get_docling_document() is not None
|
|
|
|
|
|
async def test_delete_marks_vacuum_dirty(temp_db_path):
|
|
"""A delete adds tombstone/table versions, so it must enter the auto-vacuum
|
|
lifecycle — otherwise a delete-only run closes without a final vacuum."""
|
|
dim = Config.embeddings.model.vector_dim
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(
|
|
_docling_doc("d", "body"),
|
|
[Chunk(content="body", embedding=[0.1] * dim, order=0)],
|
|
uri="mem://del",
|
|
)
|
|
assert doc.id is not None
|
|
client._vacuum_dirty = False # isolate the delete
|
|
|
|
assert await client.delete_document(doc.id) is True
|
|
assert client._vacuum_dirty is True
|
|
|
|
|
|
async def test_delete_rolls_back_on_partial_failure(temp_db_path, monkeypatch):
|
|
"""A multi-table delete is atomic: if a later table delete fails, the write
|
|
lock + version restore bring every table back, leaving no orphaned rows."""
|
|
dim = Config.embeddings.model.vector_dim
|
|
config = Config.model_copy(deep=True)
|
|
config.storage.auto_vacuum = False
|
|
|
|
async with HaikuRAG(temp_db_path, config=config, create=True) as client:
|
|
doc = await client.import_document(
|
|
_docling_doc("d", "body"),
|
|
[Chunk(content="body", embedding=[0.1] * dim, order=0)],
|
|
uri="mem://del",
|
|
title="T",
|
|
metadata={"k": "v"},
|
|
)
|
|
assert doc.id is not None
|
|
|
|
async def boom(*_a, **_k):
|
|
raise RuntimeError("meta delete failed")
|
|
|
|
# Fail the final step (document_meta) after chunks/items/documents deleted.
|
|
monkeypatch.setattr(client.store.document_meta_table, "delete", boom)
|
|
with pytest.raises(RuntimeError, match="meta delete failed"):
|
|
await client.delete_document(doc.id)
|
|
monkeypatch.undo()
|
|
|
|
# Rollback restored every table — the document is fully intact.
|
|
restored = await client.get_document_by_id(doc.id)
|
|
assert restored is not None
|
|
assert restored.title == "T"
|
|
assert restored.metadata["k"] == "v"
|
|
assert restored.content == "body"
|
|
chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(chunks) == 1
|
|
|
|
|
|
async def test_cascade_delete_is_atomic(temp_db_path, monkeypatch):
|
|
"""Deleting a parent cascades to children under one lock + snapshot. If any
|
|
delete in the subtree fails, the whole subtree is restored — a child isn't
|
|
left deleted while its parent survives."""
|
|
dim = Config.embeddings.model.vector_dim
|
|
config = Config.model_copy(deep=True)
|
|
config.storage.auto_vacuum = False
|
|
|
|
async with HaikuRAG(temp_db_path, config=config, create=True) as client:
|
|
parent = await client.import_document(
|
|
_docling_doc("p", "parent"),
|
|
[Chunk(content="parent", embedding=[0.1] * dim, order=0)],
|
|
uri="mem://parent",
|
|
)
|
|
child = await client.import_document(
|
|
_docling_doc("c", "child"),
|
|
[Chunk(content="child", embedding=[0.2] * dim, order=0)],
|
|
uri="mem://child",
|
|
metadata={"parent_uri": "mem://parent"},
|
|
)
|
|
assert parent.id is not None and child.id is not None
|
|
|
|
orig_delete = client.document_repository.delete
|
|
|
|
async def delete_failing_on_child(doc_id):
|
|
if doc_id == child.id:
|
|
raise RuntimeError("child delete failed")
|
|
return await orig_delete(doc_id)
|
|
|
|
monkeypatch.setattr(
|
|
client.document_repository, "delete", delete_failing_on_child
|
|
)
|
|
with pytest.raises(RuntimeError, match="child delete failed"):
|
|
await client.delete_document(parent.id)
|
|
monkeypatch.undo()
|
|
|
|
# Atomic: the parent delete was rolled back too — both survive.
|
|
assert await client.get_document_by_id(parent.id) is not None
|
|
assert await client.get_document_by_id(child.id) is not None
|
|
assert await client.count_documents() == 2
|
|
|
|
|
|
async def test_delete_missing_id_returns_false_without_vacuum(temp_db_path):
|
|
"""Deleting an id that doesn't exist returns False and owes no vacuum (the
|
|
existence check is inside the lock, so a no-op delete stays a no-op)."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
client._vacuum_dirty = False
|
|
assert await client.delete_document("does-not-exist") is False
|
|
assert client._vacuum_dirty is False
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_ask(allow_model_requests, temp_db_path):
|
|
"""Test asking questions through the native RAG capability."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a test document for the agent to search
|
|
await client.create_document(
|
|
content="Python is a high-level programming language.", uri="test.txt"
|
|
)
|
|
|
|
answer, citations = await client.ask("What is Python?")
|
|
|
|
# Should return a valid response
|
|
assert answer is not None
|
|
assert isinstance(answer, str)
|
|
assert isinstance(citations, list)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_expand_context(temp_db_path):
|
|
"""Test that expand_context method exists and works with basic input."""
|
|
from haiku.rag.store.models import SearchResult
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.create_document(content="Simple test content")
|
|
assert doc.id is not None
|
|
chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
search_results = [SearchResult.from_chunk(chunks[0], 0.9)]
|
|
expanded_results = await client.expand_context(search_results)
|
|
|
|
assert len(expanded_results) == 1
|
|
assert expanded_results[0].score == 0.9
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_stores_docling_json(temp_db_path):
|
|
"""Test that create_document stores DoclingDocument JSON."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.create_document(
|
|
content="Test content for docling storage",
|
|
uri="test://docling",
|
|
metadata={"test": "docling_storage"},
|
|
)
|
|
|
|
assert doc.id is not None
|
|
assert doc.docling_document is not None
|
|
assert doc.docling_version is not None
|
|
|
|
# Verify JSON is valid and can be parsed
|
|
import json
|
|
|
|
from haiku.rag.store.compression import decompress_json
|
|
|
|
parsed = json.loads(decompress_json(doc.docling_document))
|
|
assert "version" in parsed
|
|
assert parsed["version"] == doc.docling_version
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_import_document_stores_docling_data(temp_db_path):
|
|
"""Test that import_document stores DoclingDocument data correctly."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a docling document with some content
|
|
docling_doc = DoclingDocument(name="test")
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.TEXT, text="Content from docling document"
|
|
)
|
|
|
|
custom_chunks = [Chunk(content="Chunk content", order=0)]
|
|
|
|
# Import with DoclingDocument
|
|
doc = await client.import_document(
|
|
docling_document=docling_doc,
|
|
chunks=custom_chunks,
|
|
)
|
|
|
|
assert doc.id is not None
|
|
assert "Content from docling document" in doc.content
|
|
assert doc.docling_document is not None
|
|
assert doc.docling_version == docling_doc.version
|
|
# Structure is stored without pages
|
|
structure = json.loads(decompress_json(doc.docling_document))
|
|
assert "pages" not in structure
|
|
assert structure["name"] == "test"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_from_file_stores_docling_json(temp_db_path):
|
|
"""Test that create_document_from_source stores DoclingDocument JSON for files."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test.txt"
|
|
temp_path.write_text("Test file content")
|
|
|
|
doc = await client.create_document_from_source(temp_path)
|
|
assert isinstance(doc, Document)
|
|
|
|
assert doc.id is not None
|
|
assert doc.docling_document is not None
|
|
assert doc.docling_version is not None
|
|
|
|
# Verify the stored document also has the JSON
|
|
retrieved = await client.document_repository.get_by_id(
|
|
doc.id, include_blobs=True
|
|
)
|
|
assert retrieved is not None
|
|
assert retrieved.docling_document == doc.docling_document
|
|
assert retrieved.docling_version == doc.docling_version
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_update_document_stores_docling_json(temp_db_path):
|
|
"""Test that update_document stores DoclingDocument JSON when content changes."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create initial document
|
|
doc = await client.create_document(content="Initial content")
|
|
assert doc.id is not None
|
|
original_json = doc.docling_document
|
|
|
|
# Update content via update_document
|
|
updated_doc = await client.update_document(
|
|
document_id=doc.id, content="New content via fields update"
|
|
)
|
|
|
|
assert updated_doc.docling_document is not None
|
|
assert updated_doc.docling_version is not None
|
|
# JSON should be different because content changed
|
|
assert updated_doc.docling_document != original_json
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_update_document_with_custom_chunks_no_docling_json(
|
|
temp_db_path,
|
|
):
|
|
"""Test that update_document with custom chunks does not update docling JSON."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create initial document
|
|
doc = await client.create_document(content="Initial content")
|
|
assert doc.id is not None
|
|
original_json = doc.docling_document
|
|
|
|
# Update with custom chunks
|
|
custom_chunks = [Chunk(content="Custom chunk", order=0)]
|
|
updated_doc = await client.update_document(
|
|
document_id=doc.id, content="New content", chunks=custom_chunks
|
|
)
|
|
|
|
# Docling JSON should remain unchanged (no conversion when custom chunks provided)
|
|
assert updated_doc.docling_document == original_json
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_update_document_content_docling_mutually_exclusive(
|
|
temp_db_path,
|
|
):
|
|
"""Test that content and docling_document cannot both be provided."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.create_document(content="Initial content")
|
|
assert doc.id is not None
|
|
|
|
# Create a docling document
|
|
docling_doc = DoclingDocument(name="test")
|
|
docling_doc.add_text(label=DocItemLabel.TEXT, text="Some text")
|
|
|
|
with pytest.raises(ValueError, match="mutually exclusive"):
|
|
await client.update_document(
|
|
document_id=doc.id,
|
|
content="New content",
|
|
docling_document=docling_doc,
|
|
)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_update_document_with_docling_rechunks(temp_db_path):
|
|
"""Test that providing docling_document without chunks triggers rechunk."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create initial document
|
|
doc = await client.create_document(content="Initial content")
|
|
assert doc.id is not None
|
|
original_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
# Create a new docling document with different content
|
|
docling_doc = DoclingDocument(name="updated")
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.TEXT,
|
|
text="Completely different text from docling document",
|
|
)
|
|
|
|
# Update with docling document only - should rechunk from it
|
|
updated_doc = await client.update_document(
|
|
document_id=doc.id,
|
|
docling_document=docling_doc,
|
|
)
|
|
|
|
# Content should be extracted from docling document
|
|
assert "Completely different text" in updated_doc.content
|
|
assert updated_doc.docling_document is not None
|
|
assert updated_doc.docling_version == docling_doc.version
|
|
# Structure is stored without pages
|
|
structure = json.loads(decompress_json(updated_doc.docling_document))
|
|
assert "pages" not in structure
|
|
assert structure["name"] == "updated"
|
|
|
|
# Chunks should be regenerated
|
|
new_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(new_chunks) > 0
|
|
# Content should differ from original
|
|
assert new_chunks[0].content != original_chunks[0].content
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_update_document_docling_with_chunks(temp_db_path):
|
|
"""Test that providing both docling_document and chunks stores both."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create initial document
|
|
doc = await client.create_document(content="Initial content")
|
|
assert doc.id is not None
|
|
|
|
# Create a docling document
|
|
docling_doc = DoclingDocument(name="custom")
|
|
docling_doc.add_text(label=DocItemLabel.TEXT, text="Text from docling")
|
|
|
|
# Provide both docling and custom chunks
|
|
custom_chunks = [
|
|
Chunk(content="Custom chunk 1", order=0),
|
|
Chunk(content="Custom chunk 2", order=1),
|
|
]
|
|
|
|
updated_doc = await client.update_document(
|
|
document_id=doc.id,
|
|
chunks=custom_chunks,
|
|
docling_document=docling_doc,
|
|
)
|
|
|
|
# Content should be extracted from docling (since content wasn't provided)
|
|
assert "Text from docling" in updated_doc.content
|
|
assert updated_doc.docling_document is not None
|
|
structure = json.loads(decompress_json(updated_doc.docling_document))
|
|
assert "pages" not in structure
|
|
|
|
# Custom chunks should be used (not rechunked from docling)
|
|
chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(chunks) == 2
|
|
assert chunks[0].content == "Custom chunk 1"
|
|
assert chunks[1].content == "Custom chunk 2"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_file_update_stores_docling_json(temp_db_path):
|
|
"""Test that updating a file re-stores DoclingDocument JSON."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test.txt"
|
|
temp_path.write_text("Original content")
|
|
|
|
# Create initial document
|
|
doc1 = await client.create_document_from_source(temp_path)
|
|
assert isinstance(doc1, Document)
|
|
original_json = doc1.docling_document
|
|
original_version = doc1.docling_version
|
|
|
|
# Modify file
|
|
temp_path.write_text("Modified content")
|
|
|
|
# Update document from source
|
|
doc2 = await client.create_document_from_source(temp_path)
|
|
assert isinstance(doc2, Document)
|
|
assert doc2.id == doc1.id # Same document
|
|
|
|
# Docling JSON should be updated
|
|
assert doc2.docling_document is not None
|
|
assert doc2.docling_document != original_json
|
|
assert doc2.docling_version == original_version # Version stays same
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_visualize_chunk_no_document(temp_db_path):
|
|
"""Test visualize_chunk returns empty list when chunk has no document_id."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
chunk = Chunk(content="Orphan chunk", order=0)
|
|
images = await client.visualize_chunk(chunk)
|
|
assert images == []
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_visualize_chunk_no_bounding_boxes(temp_db_path):
|
|
"""Test visualize_chunk returns empty list when chunk has no bounding boxes."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create document from text (will have DoclingDocument but no page images)
|
|
doc = await client.create_document(
|
|
content="Simple text content without structure",
|
|
uri="test://simple",
|
|
)
|
|
|
|
assert doc.id is not None
|
|
assert doc.docling_document is not None
|
|
|
|
chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(chunks) >= 1
|
|
|
|
# Text documents converted via markdown won't have page images
|
|
# so visualize_chunk should return empty list
|
|
images = await client.visualize_chunk(chunks[0])
|
|
assert images == []
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_visualize_chunk_returns_list(temp_db_path):
|
|
"""Test visualize_chunk returns a list (empty or with images)."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a structured document
|
|
markdown_content = """# Chapter 1
|
|
|
|
This is paragraph one about topic A.
|
|
|
|
This is paragraph two about topic A continued.
|
|
|
|
# Chapter 2
|
|
|
|
This is paragraph four about topic C.
|
|
"""
|
|
doc = await client.create_document(
|
|
content=markdown_content,
|
|
uri="test://structured",
|
|
)
|
|
|
|
assert doc.id is not None
|
|
chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
# Find a chunk with doc_item_refs
|
|
chunks_with_refs = [c for c in chunks if c.get_chunk_metadata().doc_item_refs]
|
|
|
|
if chunks_with_refs:
|
|
# visualize_chunk should return a list (possibly empty if no page images)
|
|
images = await client.visualize_chunk(chunks_with_refs[0])
|
|
assert isinstance(images, list)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_visualize_chunk_with_pdf(temp_db_path, doclaynet_first_page_pdf):
|
|
"""Test visualize_chunk returns images with bounding boxes for PDF documents."""
|
|
from PIL.Image import Image as PILImage
|
|
|
|
from haiku.rag.config import AppConfig
|
|
|
|
pdf_path = doclaynet_first_page_pdf
|
|
config = AppConfig()
|
|
config.processing.conversion_options.do_ocr = False
|
|
|
|
async with HaikuRAG(temp_db_path, config=config, create=True) as client:
|
|
doc = await client.create_document_from_source(pdf_path)
|
|
assert isinstance(doc, Document)
|
|
assert doc.id is not None
|
|
assert doc.docling_document is not None
|
|
|
|
chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(chunks) > 0
|
|
|
|
# Find a chunk with doc_item_refs (bounding box info)
|
|
chunks_with_refs = [c for c in chunks if c.get_chunk_metadata().doc_item_refs]
|
|
assert len(chunks_with_refs) > 0, "PDF should have chunks with doc_item_refs"
|
|
|
|
# Visualize a chunk - should return images with bounding boxes drawn
|
|
images = await client.visualize_chunk(chunks_with_refs[0])
|
|
|
|
assert isinstance(images, list)
|
|
assert len(images) > 0, "PDF with page images should return visualizations"
|
|
|
|
# Verify returned objects are PIL Images
|
|
for img in images:
|
|
assert isinstance(img, PILImage)
|
|
|
|
|
|
async def test_client_visualize_chunk_multi_page(temp_db_path):
|
|
"""Test visualize_chunk returns one highlighted image per page for multi-page chunks."""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import (
|
|
DoclingDocument,
|
|
ImageRef,
|
|
ProvenanceItem,
|
|
)
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
from PIL import Image as PilImageModule
|
|
from PIL.Image import Image as PILImage
|
|
|
|
docling_doc = DoclingDocument(name="multi-page-test")
|
|
page_size = Size(width=612.0, height=792.0)
|
|
img1 = PilImageModule.new("RGB", (612, 792), color="white")
|
|
img2 = PilImageModule.new("RGB", (612, 792), color="white")
|
|
docling_doc.add_page(
|
|
page_no=1, size=page_size, image=ImageRef.from_pil(img1, dpi=72)
|
|
)
|
|
docling_doc.add_page(
|
|
page_no=2, size=page_size, image=ImageRef.from_pil(img2, dpi=72)
|
|
)
|
|
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text="Content on page one.",
|
|
prov=ProvenanceItem(
|
|
page_no=1,
|
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
|
charspan=(0, 20),
|
|
),
|
|
)
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text="Content on page two.",
|
|
prov=ProvenanceItem(
|
|
page_no=2,
|
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
|
charspan=(0, 20),
|
|
),
|
|
)
|
|
|
|
chunks = [
|
|
Chunk(
|
|
content="Content on page one.\nContent on page two.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0", "#/texts/1"],
|
|
"page_numbers": [1, 2],
|
|
"labels": ["paragraph", "paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://multi-page")
|
|
|
|
stored_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(stored_chunks) == 1
|
|
|
|
chunk = stored_chunks[0]
|
|
images = await client.visualize_chunk(chunk)
|
|
assert len(images) == 2
|
|
|
|
for img in images:
|
|
assert isinstance(img, PILImage)
|
|
assert img.size == (612, 792)
|
|
|
|
# Bounding boxes should have been drawn — images should differ from blank white
|
|
blank = PilImageModule.new("RGB", (612, 792), color="white")
|
|
for img in images:
|
|
assert img.tobytes() != blank.tobytes()
|
|
|
|
|
|
async def test_client_visualize_chunk_merged_chunks_union_pages(temp_db_path):
|
|
"""Visualizing all chunks of a merged result covers the union of their
|
|
expansions, which a single constituent chunk alone does not reach."""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import (
|
|
DoclingDocument,
|
|
ImageRef,
|
|
ProvenanceItem,
|
|
)
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
from PIL import Image as PilImageModule
|
|
|
|
docling_doc = DoclingDocument(name="merged-viz-test")
|
|
page_size = Size(width=612.0, height=792.0)
|
|
for page_no in (1, 2):
|
|
docling_doc.add_page(
|
|
page_no=page_no,
|
|
size=page_size,
|
|
image=ImageRef.from_pil(
|
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
|
),
|
|
)
|
|
|
|
# Two sections, one per page, each large enough to be returned whole and
|
|
# under budget (so no cross-boundary expansion and no clip). A single
|
|
# chunk visualizes its own section's page; both chunks cover both pages.
|
|
layout = [
|
|
(DocItemLabel.SECTION_HEADER, "Section One", 1),
|
|
(DocItemLabel.PARAGRAPH, "Page one body. " + "x" * 3000, 1),
|
|
(DocItemLabel.SECTION_HEADER, "Section Two", 2),
|
|
(DocItemLabel.PARAGRAPH, "Page two body. " + "y" * 3000, 2),
|
|
]
|
|
for i, (label, text, page_no) in enumerate(layout):
|
|
docling_doc.add_text(
|
|
label=label,
|
|
text=text,
|
|
prov=ProvenanceItem(
|
|
page_no=page_no,
|
|
bbox=BoundingBox(l=50, t=700 - (i % 2) * 100, r=550, b=650),
|
|
charspan=(0, 20),
|
|
),
|
|
)
|
|
|
|
chunks = [
|
|
Chunk(
|
|
content="Page one body. " + "x" * 3000,
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/1"],
|
|
"page_numbers": [1],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
),
|
|
Chunk(
|
|
content="Page two body. " + "y" * 3000,
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/3"],
|
|
"page_numbers": [2],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=1,
|
|
embedding=[0.1] * 2560,
|
|
),
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://merged")
|
|
|
|
stored_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
stored_chunks.sort(key=lambda c: c.order)
|
|
assert len(stored_chunks) == 2
|
|
c1, c2 = stored_chunks
|
|
|
|
solo_images = await client.visualize_chunk(c1)
|
|
assert len(solo_images) == 1
|
|
|
|
merged_images = await client.visualize_chunk([c1, c2])
|
|
assert len(merged_images) == 2
|
|
|
|
|
|
async def test_client_visualize_chunk_two_tone_highlights(temp_db_path):
|
|
"""Matched content draws stronger than context swept in by expansion."""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import (
|
|
DoclingDocument,
|
|
ImageRef,
|
|
ProvenanceItem,
|
|
)
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
from PIL import Image as PilImageModule
|
|
|
|
docling_doc = DoclingDocument(name="two-tone-test")
|
|
page_size = Size(width=612.0, height=792.0)
|
|
docling_doc.add_page(
|
|
page_no=1,
|
|
size=page_size,
|
|
image=ImageRef.from_pil(
|
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
|
),
|
|
)
|
|
|
|
# Three small paragraphs; the chunk matches only the middle one, so
|
|
# expansion sweeps in its neighbors.
|
|
for i in range(3):
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text=f"Paragraph {i}.",
|
|
prov=ProvenanceItem(
|
|
page_no=1,
|
|
bbox=BoundingBox(l=50, t=700 - i * 100, r=550, b=650 - i * 100),
|
|
charspan=(0, 12),
|
|
),
|
|
)
|
|
|
|
chunks = [
|
|
Chunk(
|
|
content="Paragraph 1.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/1"],
|
|
"page_numbers": [1],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://two-tone")
|
|
stored_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(stored_chunks) == 1
|
|
|
|
images = await client.visualize_chunk(stored_chunks[0])
|
|
assert len(images) == 1
|
|
image = images[0]
|
|
|
|
# Page dpi 72 == document coords, bottom-left origin flipped to
|
|
# top-left: item i's box spans y = 92 + i * 100 .. 142 + i * 100.
|
|
matched = image.getpixel((300, 217)) # inside #/texts/1
|
|
swept = image.getpixel((300, 117)) # inside #/texts/0
|
|
background = image.getpixel((300, 30)) # outside all boxes
|
|
|
|
assert matched != background
|
|
assert swept != background
|
|
assert matched != swept
|
|
|
|
|
|
async def test_client_visualize_chunk_uses_given_refs(temp_db_path):
|
|
"""Explicit refs (the citation's doc_item_refs) restrict the visualization
|
|
to exactly those items, instead of re-expanding the chunk's context."""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import (
|
|
DoclingDocument,
|
|
ImageRef,
|
|
ProvenanceItem,
|
|
)
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
from PIL import Image as PilImageModule
|
|
|
|
docling_doc = DoclingDocument(name="refs-test")
|
|
page_size = Size(width=612.0, height=792.0)
|
|
for page_no in (1, 2):
|
|
docling_doc.add_page(
|
|
page_no=page_no,
|
|
size=page_size,
|
|
image=ImageRef.from_pil(
|
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
|
),
|
|
)
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text="Content on page one.",
|
|
prov=ProvenanceItem(
|
|
page_no=1, bbox=BoundingBox(l=50, t=700, r=550, b=650), charspan=(0, 20)
|
|
),
|
|
)
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text="Content on page two.",
|
|
prov=ProvenanceItem(
|
|
page_no=2, bbox=BoundingBox(l=50, t=700, r=550, b=650), charspan=(0, 20)
|
|
),
|
|
)
|
|
|
|
chunks = [
|
|
Chunk(
|
|
content="Content on page one.\nContent on page two.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0", "#/texts/1"],
|
|
"page_numbers": [1, 2],
|
|
"labels": ["paragraph", "paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://refs")
|
|
chunk = (await client.chunk_repository.get_by_document_id(doc.id))[0]
|
|
|
|
# No refs: re-expands the chunk's own refs → both pages.
|
|
assert len(await client.visualize_chunk(chunk)) == 2
|
|
# Given only the page-one ref → only page one is rendered.
|
|
assert len(await client.visualize_chunk(chunk, refs=["#/texts/0"])) == 1
|
|
|
|
|
|
async def test_client_visualize_chunk_no_expand_shows_only_chunk(temp_db_path):
|
|
"""expand=False draws only the chunk's own items, not the expanded section."""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import (
|
|
DoclingDocument,
|
|
ImageRef,
|
|
ProvenanceItem,
|
|
)
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
from PIL import Image as PilImageModule
|
|
|
|
docling_doc = DoclingDocument(name="no-expand-test")
|
|
page_size = Size(width=612.0, height=792.0)
|
|
for page_no in (1, 2):
|
|
docling_doc.add_page(
|
|
page_no=page_no,
|
|
size=page_size,
|
|
image=ImageRef.from_pil(
|
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
|
),
|
|
)
|
|
for page_no in (1, 2):
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text=f"Short paragraph on page {page_no}.",
|
|
prov=ProvenanceItem(
|
|
page_no=page_no,
|
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
|
charspan=(0, 20),
|
|
),
|
|
)
|
|
|
|
chunks = [
|
|
Chunk(
|
|
content="Short paragraph on page 1.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0"],
|
|
"page_numbers": [1],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://no-expand")
|
|
chunk = (await client.chunk_repository.get_by_document_id(doc.id))[0]
|
|
|
|
# Default expands the chunk's context outward → reaches page two.
|
|
assert len(await client.visualize_chunk(chunk)) == 2
|
|
# expand=False draws only the chunk's own page-one item.
|
|
assert len(await client.visualize_chunk(chunk, expand=False)) == 1
|
|
|
|
|
|
# =============================================================================
|
|
# convert() method tests
|
|
# =============================================================================
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_convert_text(temp_db_path):
|
|
"""Test convert() with plain text content."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
text = "This is some test content for conversion."
|
|
docling_doc = await client.convert(text)
|
|
|
|
assert isinstance(docling_doc, DoclingDocument)
|
|
# Check the content is preserved in markdown export
|
|
markdown = docling_doc.export_to_markdown()
|
|
assert "test content" in markdown
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_convert_file(temp_db_path):
|
|
"""Test convert() with a file path."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test.txt"
|
|
temp_path.write_text("File content for conversion test.")
|
|
|
|
docling_doc = await client.convert(temp_path)
|
|
|
|
assert isinstance(docling_doc, DoclingDocument)
|
|
markdown = docling_doc.export_to_markdown()
|
|
assert "File content" in markdown
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_convert_file_not_found(temp_db_path):
|
|
"""Test convert() raises ValueError for non-existent file."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with pytest.raises(ValueError, match="File does not exist"):
|
|
await client.convert(Path("/nonexistent/path/file.txt"))
|
|
|
|
|
|
async def test_client_convert_from_url(temp_db_path):
|
|
"""convert() with an http(s) URL downloads to a tempfile and converts."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
mock_response = AsyncMock()
|
|
mock_response.content = (
|
|
b"<html><body><p>URL convert path content.</p></body></html>"
|
|
)
|
|
mock_response.headers = {"content-type": "text/html"}
|
|
mock_response.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_response):
|
|
docling_doc = await client.convert("https://example.com/page.html")
|
|
|
|
assert isinstance(docling_doc, DoclingDocument)
|
|
markdown = docling_doc.export_to_markdown()
|
|
assert "URL convert path content" in markdown
|
|
|
|
|
|
async def test_client_convert_from_url_unsupported_content_type(temp_db_path):
|
|
"""convert() rejects URLs whose content type isn't supported by the converter."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
mock_response = AsyncMock()
|
|
mock_response.content = b"\x00\x01\x02binary"
|
|
mock_response.headers = {"content-type": "application/octet-stream"}
|
|
mock_response.raise_for_status = AsyncMock()
|
|
|
|
with patch("httpx.AsyncClient.get", return_value=mock_response):
|
|
with pytest.raises(ValueError, match="Unsupported content type"):
|
|
await client.convert("https://example.com/blob.bin")
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_convert_unsupported_extension(temp_db_path):
|
|
"""Test convert() raises ValueError for unsupported file extension."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test.xyz"
|
|
temp_path.write_text("content")
|
|
|
|
with pytest.raises(ValueError, match="Unsupported file extension"):
|
|
await client.convert(temp_path)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_convert_file_uri(temp_db_path):
|
|
"""Test convert() with a file:// URI string."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
temp_path = Path(temp_dir) / "test.txt"
|
|
temp_path.write_text("URI file content.")
|
|
file_uri = temp_path.as_uri()
|
|
|
|
docling_doc = await client.convert(file_uri)
|
|
|
|
assert isinstance(docling_doc, DoclingDocument)
|
|
markdown = docling_doc.export_to_markdown()
|
|
assert "URI file content" in markdown
|
|
|
|
|
|
# =============================================================================
|
|
# chunk() method tests
|
|
# =============================================================================
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_chunk_basic(temp_db_path):
|
|
"""Test chunk() produces Chunk objects from DoclingDocument."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# First convert some text
|
|
docling_doc = await client.convert("This is test content for chunking.")
|
|
|
|
# Then chunk it
|
|
chunks = await client.chunk(docling_doc)
|
|
|
|
assert isinstance(chunks, list)
|
|
assert len(chunks) > 0
|
|
assert all(isinstance(c, Chunk) for c in chunks)
|
|
# Chunks should have content but no embedding yet
|
|
assert all(c.content for c in chunks)
|
|
assert all(c.embedding is None for c in chunks)
|
|
# Chunks should not have document_id yet (not stored)
|
|
assert all(c.document_id is None for c in chunks)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_chunk_preserves_metadata(temp_db_path):
|
|
"""Test chunk() preserves structured metadata from DoclingDocument."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Convert structured markdown
|
|
markdown = """# Chapter 1
|
|
|
|
This is the first paragraph.
|
|
|
|
## Section 1.1
|
|
|
|
This is a subsection.
|
|
"""
|
|
docling_doc = await client.convert(markdown)
|
|
chunks = await client.chunk(docling_doc)
|
|
|
|
assert len(chunks) > 0
|
|
|
|
# Check that at least some chunks have metadata
|
|
has_metadata = False
|
|
for chunk in chunks:
|
|
meta = chunk.get_chunk_metadata()
|
|
if meta.doc_item_refs or meta.headings:
|
|
has_metadata = True
|
|
break
|
|
|
|
assert has_metadata, "Chunks should have structured metadata"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_chunk_empty_document(temp_db_path):
|
|
"""Test chunk() with empty DoclingDocument."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create an empty DoclingDocument
|
|
empty_doc = DoclingDocument(name="empty")
|
|
|
|
chunks = await client.chunk(empty_doc)
|
|
|
|
assert isinstance(chunks, list)
|
|
assert len(chunks) == 0
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_import_document_embeds_chunks_without_embeddings(temp_db_path):
|
|
"""Test that import_document embeds chunks that don't have embeddings."""
|
|
from docling_core.types.doc.document import DoclingDocument
|
|
from docling_core.types.doc.labels import DocItemLabel
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create a DoclingDocument
|
|
docling_doc = DoclingDocument(name="test")
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.TEXT, text="Document with unembedded chunks"
|
|
)
|
|
|
|
# Create chunks without embeddings
|
|
chunks = [
|
|
Chunk(content="First chunk without embedding", order=0),
|
|
Chunk(content="Second chunk without embedding", order=1),
|
|
]
|
|
|
|
# Import document with chunks that have no embeddings
|
|
doc = await client.import_document(
|
|
docling_document=docling_doc,
|
|
chunks=chunks,
|
|
)
|
|
assert doc.id is not None
|
|
|
|
# Verify chunks were stored
|
|
stored_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(stored_chunks) == 2
|
|
|
|
# Verify vector search works (proves embeddings were generated)
|
|
results = await client.search("First chunk", search_type="vector")
|
|
assert len(results) > 0
|
|
assert results[0].content == "First chunk without embedding"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_update_document_embeds_chunks_without_embeddings(temp_db_path):
|
|
"""Test that update_document embeds chunks that don't have embeddings."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create initial document
|
|
doc = await client.create_document(content="Initial content")
|
|
assert doc.id is not None
|
|
|
|
# Update with chunks that have no embeddings
|
|
new_chunks = [
|
|
Chunk(content="Updated chunk without embedding", order=0),
|
|
]
|
|
await client.update_document(
|
|
document_id=doc.id,
|
|
content="Updated content",
|
|
chunks=new_chunks,
|
|
)
|
|
|
|
# Verify chunks were stored
|
|
stored_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
|
assert len(stored_chunks) == 1
|
|
assert stored_chunks[0].content == "Updated chunk without embedding"
|
|
|
|
# Verify vector search works (proves embeddings were generated)
|
|
results = await client.search("Updated chunk", search_type="vector")
|
|
assert len(results) > 0
|
|
assert results[0].content == "Updated chunk without embedding"
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_create_document_with_html_format(temp_db_path):
|
|
"""Test create_document with HTML format preserves document structure."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
html_content = """
|
|
<h1>Main Title</h1>
|
|
<p>Introduction paragraph.</p>
|
|
<h2>Section Header</h2>
|
|
<ul>
|
|
<li>Item 1</li>
|
|
<li>Item 2</li>
|
|
</ul>
|
|
"""
|
|
|
|
doc = await client.create_document(
|
|
content=html_content,
|
|
uri="test://html-doc",
|
|
format="html",
|
|
)
|
|
|
|
assert doc.id is not None
|
|
assert doc.docling_document is not None
|
|
|
|
# Verify the DoclingDocument has proper structure
|
|
docling_doc = doc.get_docling_document()
|
|
assert docling_doc is not None
|
|
|
|
items = list(docling_doc.iterate_items())
|
|
labels = [str(getattr(item, "label", "")) for item, _ in items]
|
|
|
|
# HTML format should preserve headers and list items
|
|
assert "title" in labels or "section_header" in labels
|
|
assert "list_item" in labels
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_client_convert_with_html_format(temp_db_path):
|
|
"""Test convert with HTML format."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
html_content = "<h1>Title</h1><p>Text</p>"
|
|
|
|
docling_doc = await client.convert(html_content, format="html")
|
|
|
|
items = list(docling_doc.iterate_items())
|
|
labels = [str(getattr(item, "label", "")) for item, _ in items]
|
|
|
|
assert "title" in labels
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.vcr()
|
|
async def test_sql_injection_is_blocked_with_escaping(temp_db_path):
|
|
"""SQL injection is blocked when using escape_sql_string.
|
|
|
|
This test verifies that escape_sql_string properly prevents SQL injection
|
|
by escaping single quotes in user input.
|
|
"""
|
|
from haiku.rag.utils import escape_sql_string
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
# Create documents
|
|
await client.create_document(
|
|
content="Secret classified data XYZ",
|
|
uri="secret://doc",
|
|
title="Secret",
|
|
)
|
|
await client.create_document(
|
|
content="Public report about weather",
|
|
uri="public://report",
|
|
title="Weather Report",
|
|
)
|
|
|
|
# Without escaping, this injection would match all documents
|
|
# by breaking out of the string literal: title = 'x' OR title LIKE '%'
|
|
injection_payload = "x' OR title LIKE '%"
|
|
|
|
# With proper escaping, single quotes become double quotes
|
|
# so the filter becomes: title = 'x'' OR title LIKE ''%'
|
|
# which searches for a literal title containing the injection string
|
|
safe_payload = escape_sql_string(injection_payload)
|
|
docs = await client.list_documents(filter=f"title = '{safe_payload}'")
|
|
|
|
# Should find 0 documents (injection is escaped, searching for literal string)
|
|
assert len(docs) == 0
|
|
|
|
# Verify the escaping works correctly
|
|
assert safe_payload == "x'' OR title LIKE ''%"
|
|
|
|
# Verify unescaped injection would have matched documents (for test validity)
|
|
# This demonstrates that the injection works without escaping
|
|
docs_unescaped = await client.list_documents(
|
|
filter=f"title = '{injection_payload}'"
|
|
)
|
|
assert len(docs_unescaped) == 2 # SQL injection succeeds without escaping
|
|
|
|
|
|
# =============================================================================
|
|
# URL-prefixed content regression tests
|
|
# =============================================================================
|
|
|
|
|
|
def _patch_embed_chunks(monkeypatch):
|
|
async def fake_embed_chunks(chunks, embedder, config):
|
|
for chunk in chunks:
|
|
chunk.embedding = [0.0] * 2560
|
|
return chunks
|
|
|
|
monkeypatch.setattr("haiku.rag.embeddings.embed_chunks", fake_embed_chunks)
|
|
|
|
|
|
async def test_create_document_with_url_prefixed_content(temp_db_path, monkeypatch):
|
|
"""Text whose first line is a URL must be stored as text, not fetched."""
|
|
_patch_embed_chunks(monkeypatch)
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
content = "https://example.com/foo\n\n# Heading\n\nBody text here."
|
|
doc = await client.create_document(content=content, uri="test://url-prefixed")
|
|
|
|
assert doc.id is not None
|
|
assert "example.com" in doc.content
|
|
assert "Heading" in doc.content
|
|
|
|
|
|
async def test_update_document_with_url_prefixed_content(temp_db_path, monkeypatch):
|
|
"""update_document(content=...) with URL-prefixed text must not fetch it."""
|
|
_patch_embed_chunks(monkeypatch)
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.create_document(
|
|
content="initial body", uri="test://update-url"
|
|
)
|
|
assert doc.id is not None
|
|
|
|
url_prefixed = "https://example.com/bar\n\n# New heading\n\nReplacement body."
|
|
updated = await client.update_document(doc.id, content=url_prefixed)
|
|
|
|
assert "example.com" in updated.content
|
|
assert "New heading" in updated.content
|
|
|
|
|
|
async def test_update_document_with_chunks_keeps_page_images(temp_db_path, monkeypatch):
|
|
"""Replacing content and chunks without a docling document writes the stored
|
|
record back as-is, so its page rasters must survive the round trip."""
|
|
_patch_embed_chunks(monkeypatch)
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.create_document(content="initial body", uri="test://pages")
|
|
assert doc.id is not None
|
|
|
|
sentinel_pages = b"\x80SENTINEL_PAGE_BYTES"
|
|
await client.store.documents_table.update(
|
|
{"docling_pages": sentinel_pages}, where=f"id = '{doc.id}'"
|
|
)
|
|
|
|
await client.update_document(
|
|
doc.id,
|
|
content="replacement body",
|
|
chunks=[Chunk(content="replacement body")],
|
|
)
|
|
|
|
stored = await client.document_repository.get_by_id(doc.id, include_blobs=True)
|
|
assert stored is not None
|
|
assert stored.content == "replacement body"
|
|
assert stored.docling_pages == sentinel_pages
|
|
|
|
|
|
async def test_rebuild_rechunk_with_url_prefixed_stored_content(
|
|
temp_db_path, monkeypatch
|
|
):
|
|
"""RECHUNK rebuild must handle stored markdown whose first line is a URL."""
|
|
from haiku.rag.client import RebuildMode
|
|
|
|
_patch_embed_chunks(monkeypatch)
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.create_document(
|
|
content="plain seed content", uri="file:///nonexistent/path.txt"
|
|
)
|
|
assert doc.id is not None
|
|
|
|
# Overwrite stored content to simulate markdown that starts with a URL,
|
|
# bypassing the (also-affected) create_document path so this test
|
|
# specifically exercises the rebuild path.
|
|
doc.content = "https://example.com/baz\n\n# Stored\n\nStored body text."
|
|
await client.document_repository.update(doc)
|
|
|
|
processed_ids = [
|
|
doc_id async for doc_id in client.rebuild_database(mode=RebuildMode.RECHUNK)
|
|
]
|
|
assert doc.id in processed_ids
|
|
|
|
doc_after = await client.document_repository.get_by_id(doc.id)
|
|
assert doc_after is not None
|
|
assert "example.com" in doc_after.content
|
|
assert "Stored" in doc_after.content
|
|
|
|
|
|
async def test_metadata_only_update_waits_for_write_lock(temp_db_path):
|
|
"""The metadata-only update path serializes with other writers so it
|
|
cannot land inside another writer's critical section (e.g. between
|
|
create_tag's version snapshot and its per-table tag creation)."""
|
|
import asyncio
|
|
|
|
dim = Config.embeddings.model.vector_dim
|
|
docling_doc = DoclingDocument(name="d")
|
|
docling_doc.add_text(label=DocItemLabel.TEXT, text="body")
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(
|
|
docling_doc,
|
|
[Chunk(content="body", embedding=[0.1] * dim, order=0)],
|
|
uri="mem://meta",
|
|
)
|
|
|
|
async with client.store._write_lock:
|
|
task = asyncio.create_task(
|
|
client.update_document(document_id=doc.id, metadata={"k": "v"})
|
|
)
|
|
await asyncio.sleep(0.1)
|
|
assert not task.done()
|
|
updated = await task
|
|
assert updated.metadata == {"k": "v"}
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"uri,expected",
|
|
[
|
|
("https://example.com/doc.pdf", True),
|
|
("s3://bucket/key", True),
|
|
("mem://not-a-source", False),
|
|
# urlparse rejects a malformed IPv6 host; a stored URI that no longer
|
|
# parses must not abort the caller's rebuild sweep.
|
|
("http://[::1", False),
|
|
],
|
|
ids=["https", "s3", "unknown_scheme", "unparseable"],
|
|
)
|
|
def test_check_source_accessible(uri, expected):
|
|
assert check_source_accessible(uri) is expected
|
|
|
|
|
|
def test_check_source_accessible_file_uri(tmp_path):
|
|
existing = tmp_path / "there.txt"
|
|
existing.write_text("x")
|
|
|
|
assert check_source_accessible(existing.as_uri()) is True
|
|
assert check_source_accessible((tmp_path / "gone.txt").as_uri()) is False
|
|
|
|
|
|
def _bbox_doc(*, with_page_image: bool, pages: tuple[int, ...] = (1,)):
|
|
"""DoclingDocument with one paragraph per page, each carrying a bbox.
|
|
|
|
``with_page_image=False`` produces pages with no raster, so bounding boxes
|
|
resolve but there is nothing to draw them on.
|
|
"""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import ImageRef, ProvenanceItem
|
|
from PIL import Image as PilImageModule
|
|
|
|
doc = DoclingDocument(name="bbox-doc")
|
|
size = Size(width=612.0, height=792.0)
|
|
for page_no in pages:
|
|
image = (
|
|
ImageRef.from_pil(
|
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
|
)
|
|
if with_page_image
|
|
else None
|
|
)
|
|
doc.add_page(page_no=page_no, size=size, image=image)
|
|
doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text=f"Content on page {page_no}.",
|
|
prov=ProvenanceItem(
|
|
page_no=page_no,
|
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
|
charspan=(0, 20),
|
|
),
|
|
)
|
|
return doc
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_visualize_chunk_returns_empty_without_page_rasters(temp_db_path):
|
|
"""Boxes resolve, but a document ingested without page images has nothing
|
|
to render them onto."""
|
|
docling_doc = _bbox_doc(with_page_image=False)
|
|
chunks = [
|
|
Chunk(
|
|
content="Content on page 1.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0"],
|
|
"page_numbers": [1],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://no-raster")
|
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
assert await client.visualize_chunk(stored[0]) == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_visualize_chunk_skips_pages_without_a_raster(temp_db_path):
|
|
"""A document where only some pages carry a raster renders just those."""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import ImageRef, ProvenanceItem
|
|
from PIL import Image as PilImageModule
|
|
|
|
docling_doc = DoclingDocument(name="mixed-rasters")
|
|
size = Size(width=612.0, height=792.0)
|
|
docling_doc.add_page(
|
|
page_no=1,
|
|
size=size,
|
|
image=ImageRef.from_pil(
|
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
|
),
|
|
)
|
|
docling_doc.add_page(page_no=2, size=size, image=None)
|
|
for page_no in (1, 2):
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text=f"Content on page {page_no}.",
|
|
prov=ProvenanceItem(
|
|
page_no=page_no,
|
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
|
charspan=(0, 20),
|
|
),
|
|
)
|
|
|
|
chunks = [
|
|
Chunk(
|
|
content="Content on page 1.\nContent on page 2.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0", "#/texts/1"],
|
|
"page_numbers": [1, 2],
|
|
"labels": ["paragraph", "paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(
|
|
docling_doc, chunks, uri="test://mixed-rasters"
|
|
)
|
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
images = await client.visualize_chunk(stored[0])
|
|
|
|
assert len(images) == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_visualize_chunk_returns_empty_when_pages_row_missing(temp_db_path):
|
|
docling_doc = _bbox_doc(with_page_image=True)
|
|
chunks = [
|
|
Chunk(
|
|
content="Content on page 1.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0"],
|
|
"page_numbers": [1],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://no-row")
|
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
async def no_pages_row(document_id):
|
|
return None
|
|
|
|
client.document_repository.get_pages_data = no_pages_row # type: ignore[method-assign]
|
|
|
|
assert await client.visualize_chunk(stored[0]) == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_visualize_chunk_skips_box_on_unstored_page(temp_db_path):
|
|
"""A bounding box referencing a page the document never registered is
|
|
skipped rather than raising."""
|
|
from docling_core.types.doc.base import BoundingBox, Size
|
|
from docling_core.types.doc.document import ImageRef, ProvenanceItem
|
|
from PIL import Image as PilImageModule
|
|
|
|
docling_doc = DoclingDocument(name="orphan-page-box")
|
|
docling_doc.add_page(
|
|
page_no=1,
|
|
size=Size(width=612.0, height=792.0),
|
|
image=ImageRef.from_pil(
|
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
|
),
|
|
)
|
|
docling_doc.add_text(
|
|
label=DocItemLabel.PARAGRAPH,
|
|
text="Content attributed to a page with no raster.",
|
|
prov=ProvenanceItem(
|
|
page_no=3,
|
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
|
charspan=(0, 20),
|
|
),
|
|
)
|
|
|
|
chunks = [
|
|
Chunk(
|
|
content="Content attributed to a page with no raster.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0"],
|
|
"page_numbers": [3],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(
|
|
docling_doc, chunks, uri="test://orphan-page"
|
|
)
|
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
assert await client.visualize_chunk(stored[0]) == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_visualize_chunk_without_refs_falls_back_to_chunk_metadata(temp_db_path):
|
|
"""A chunk carrying no doc_item_refs has nothing to expand from."""
|
|
docling_doc = _bbox_doc(with_page_image=True)
|
|
chunks = [
|
|
Chunk(
|
|
content="Content on page 1.",
|
|
metadata={"page_numbers": [1], "labels": ["paragraph"]},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://no-refs")
|
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
assert await client.visualize_chunk(stored[0]) == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_visualize_chunk_falls_back_when_expansion_drops_refs(temp_db_path):
|
|
"""If expansion returns results carrying no refs, the original search
|
|
results' refs are used instead."""
|
|
from haiku.rag.client import search as search_module
|
|
|
|
docling_doc = _bbox_doc(with_page_image=True)
|
|
chunks = [
|
|
Chunk(
|
|
content="Content on page 1.",
|
|
metadata={
|
|
"doc_item_refs": ["#/texts/0"],
|
|
"page_numbers": [1],
|
|
"labels": ["paragraph"],
|
|
},
|
|
order=0,
|
|
embedding=[0.1] * 2560,
|
|
)
|
|
]
|
|
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
doc = await client.import_document(docling_doc, chunks, uri="test://drops-refs")
|
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
|
|
|
async def expansion_without_refs(_client, results):
|
|
return [r.model_copy(update={"doc_item_refs": []}) for r in results]
|
|
|
|
with patch.object(search_module, "expand_context", expansion_without_refs):
|
|
images = await client.visualize_chunk(stored[0])
|
|
|
|
assert len(images) == 1
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
@pytest.mark.parametrize("auto_vacuum", [True, False])
|
|
async def test_import_documents_schedules_vacuum_per_config(temp_db_path, auto_vacuum):
|
|
"""A batch import runs a background vacuum only when auto_vacuum is on."""
|
|
from haiku.rag.config import AppConfig
|
|
|
|
config = AppConfig()
|
|
config.storage.auto_vacuum = auto_vacuum
|
|
|
|
async with HaikuRAG(temp_db_path, config=config, create=True) as client:
|
|
docling_doc = await client.convert("Batch imported body.")
|
|
chunks = await client.chunk(docling_doc)
|
|
|
|
# Spy rather than inspecting _vacuum_tasks: the scheduling code
|
|
# discards each task on completion, so the set races to empty. The
|
|
# spy's count does not race, and draining the scheduled task keeps
|
|
# the assertion deterministic without pulling in the close-time pass.
|
|
with patch.object(client.store, "vacuum", new=AsyncMock()) as vacuum:
|
|
await client.import_documents(
|
|
[
|
|
DocumentImport(
|
|
docling_document=docling_doc,
|
|
chunks=chunks,
|
|
uri="test://batch-vacuum",
|
|
)
|
|
]
|
|
)
|
|
await asyncio.gather(*client._vacuum_tasks)
|
|
|
|
assert vacuum.await_count == (1 if auto_vacuum else 0)
|
|
|
|
|
|
@pytest.mark.vcr()
|
|
async def test_reingesting_a_source_applies_an_explicit_title(temp_db_path):
|
|
"""Re-adding a changed source with a title updates both, in place."""
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
source = Path(temp_dir) / "retitled.txt"
|
|
source.write_text("stable content")
|
|
|
|
first = await client.create_document_from_source(source)
|
|
assert not isinstance(first, list)
|
|
source.write_text("changed content")
|
|
|
|
second = await client.create_document_from_source(
|
|
source, title="Explicit Title"
|
|
)
|
|
|
|
assert not isinstance(second, list)
|
|
assert second.id == first.id
|
|
assert second.title == "Explicit Title"
|
|
assert second.content == "changed content"
|
|
|
|
|
|
async def test_update_document_rejects_unknown_id(temp_db_path):
|
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
|
with pytest.raises(ValueError, match="not found"):
|
|
await client.update_document("no-such-document", content="x")
|