Unify update_document, update_document_fields
This commit is contained in:
parent
16ba3a9962
commit
e97d235c98
3 changed files with 100 additions and 83 deletions
|
|
@ -198,13 +198,13 @@ class HaikuRAG:
|
||||||
document: Document,
|
document: Document,
|
||||||
chunks: list[Chunk],
|
chunks: list[Chunk],
|
||||||
) -> Document:
|
) -> Document:
|
||||||
"""Store a document with its pre-embedded chunks.
|
"""Store a document with chunks, embedding any that lack embeddings.
|
||||||
|
|
||||||
Handles versioning/rollback on failure.
|
Handles versioning/rollback on failure.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
document: The document to store (will be created).
|
document: The document to store (will be created).
|
||||||
chunks: Pre-embedded chunks to store with the document.
|
chunks: Chunks to store (will be embedded if lacking embeddings).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The created Document instance with ID set.
|
The created Document instance with ID set.
|
||||||
|
|
@ -243,13 +243,13 @@ class HaikuRAG:
|
||||||
document: Document,
|
document: Document,
|
||||||
chunks: list[Chunk],
|
chunks: list[Chunk],
|
||||||
) -> Document:
|
) -> Document:
|
||||||
"""Update a document and replace its chunks with pre-embedded chunks.
|
"""Update a document and replace its chunks, embedding any that lack embeddings.
|
||||||
|
|
||||||
Handles versioning/rollback on failure.
|
Handles versioning/rollback on failure.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
document: The document to update (must have ID set).
|
document: The document to update (must have ID set).
|
||||||
chunks: Pre-embedded chunks to replace existing chunks.
|
chunks: Chunks to replace existing (will be embedded if lacking embeddings).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The updated Document instance.
|
The updated Document instance.
|
||||||
|
|
@ -702,31 +702,7 @@ class HaikuRAG:
|
||||||
"""
|
"""
|
||||||
return await self.document_repository.get_by_uri(uri)
|
return await self.document_repository.get_by_uri(uri)
|
||||||
|
|
||||||
async def update_document(self, document: Document) -> Document:
|
async def update_document(
|
||||||
"""Update an existing document.
|
|
||||||
|
|
||||||
Reconverts content, rechunks, and regenerates embeddings.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
document: The document to update (must have ID set).
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
The updated Document instance.
|
|
||||||
"""
|
|
||||||
from haiku.rag.embeddings import embed_chunks
|
|
||||||
|
|
||||||
# Convert → Chunk → Embed using primitives
|
|
||||||
docling_document = await self.convert(document.content)
|
|
||||||
chunks = await self.chunk(docling_document)
|
|
||||||
embedded_chunks = await embed_chunks(chunks, self._config)
|
|
||||||
|
|
||||||
# Store DoclingDocument JSON
|
|
||||||
document.docling_document_json = docling_document.model_dump_json()
|
|
||||||
document.docling_version = docling_document.version
|
|
||||||
|
|
||||||
return await self._update_document_with_chunks(document, embedded_chunks)
|
|
||||||
|
|
||||||
async def update_document_fields(
|
|
||||||
self,
|
self,
|
||||||
document_id: str,
|
document_id: str,
|
||||||
content: str | None = None,
|
content: str | None = None,
|
||||||
|
|
@ -736,24 +712,27 @@ class HaikuRAG:
|
||||||
docling_document_json: str | None = None,
|
docling_document_json: str | None = None,
|
||||||
docling_version: str | None = None,
|
docling_version: str | None = None,
|
||||||
) -> Document:
|
) -> Document:
|
||||||
"""Update specific fields of a document by ID.
|
"""Update a document by ID.
|
||||||
|
|
||||||
|
Updates specified fields. When content or docling_document_json is provided,
|
||||||
|
the document is rechunked and re-embedded. Updates to only metadata or title
|
||||||
|
skip rechunking for efficiency.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
document_id: The ID of the document to update
|
document_id: The ID of the document to update.
|
||||||
content: New content for the document (mutually exclusive with docling_document_json)
|
content: New content (mutually exclusive with docling_document_json).
|
||||||
metadata: New metadata for the document
|
metadata: New metadata dict.
|
||||||
chunks: Custom chunks to use instead of auto-generating
|
chunks: Custom pre-embedded chunks (skips auto-chunking).
|
||||||
title: New title for the document
|
title: New title.
|
||||||
docling_document_json: Serialized DoclingDocument JSON (mutually exclusive with content)
|
docling_document_json: Serialized DoclingDocument JSON (mutually exclusive with content).
|
||||||
docling_version: DoclingDocument schema version (required with docling_document_json)
|
docling_version: DoclingDocument schema version (required with docling_document_json).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The updated Document instance.
|
The updated Document instance.
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
ValueError: If both content and docling_document_json are provided,
|
ValueError: If document not found, if both content and docling_document_json
|
||||||
if docling_document_json is provided without docling_version,
|
are provided, or if docling_document_json is provided without docling_version.
|
||||||
or if the JSON is invalid.
|
|
||||||
"""
|
"""
|
||||||
from docling_core.types.doc.document import DoclingDocument
|
from docling_core.types.doc.document import DoclingDocument
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -50,11 +50,11 @@ async def test_client_document_crud(qa_corpus: Dataset, temp_db_path):
|
||||||
assert non_existent is None
|
assert non_existent is None
|
||||||
|
|
||||||
# Test update_document
|
# Test update_document
|
||||||
retrieved_doc.content = "Updated content"
|
updated_doc = await client.update_document(
|
||||||
retrieved_doc.uri = "file:///updated/path.txt"
|
document_id=retrieved_doc.id, # type: ignore[arg-type]
|
||||||
updated_doc = await client.update_document(retrieved_doc)
|
content="Updated content",
|
||||||
|
)
|
||||||
assert updated_doc.content == "Updated content"
|
assert updated_doc.content == "Updated content"
|
||||||
assert updated_doc.uri == "file:///updated/path.txt"
|
|
||||||
|
|
||||||
# Test list_documents
|
# Test list_documents
|
||||||
all_docs = await client.list_documents()
|
all_docs = await client.list_documents()
|
||||||
|
|
@ -79,7 +79,7 @@ async def test_client_document_crud(qa_corpus: Dataset, temp_db_path):
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_update_document_fields(qa_corpus: Dataset, temp_db_path):
|
async def test_client_update_document(qa_corpus: Dataset, temp_db_path):
|
||||||
"""Test updating document with individual parameters."""
|
"""Test updating document with individual parameters."""
|
||||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
# Get test data
|
# Get test data
|
||||||
|
|
@ -99,7 +99,7 @@ async def test_client_update_document_fields(qa_corpus: Dataset, temp_db_path):
|
||||||
original_id = created_doc.id
|
original_id = created_doc.id
|
||||||
|
|
||||||
# Test updating only content
|
# Test updating only content
|
||||||
updated_doc = await client.update_document_fields(
|
updated_doc = await client.update_document(
|
||||||
document_id=original_id, content="Updated content only"
|
document_id=original_id, content="Updated content only"
|
||||||
)
|
)
|
||||||
assert updated_doc.id == original_id
|
assert updated_doc.id == original_id
|
||||||
|
|
@ -109,7 +109,7 @@ async def test_client_update_document_fields(qa_corpus: Dataset, temp_db_path):
|
||||||
|
|
||||||
# Test updating only metadata
|
# Test updating only metadata
|
||||||
new_metadata = {"source": "updated", "version": "2.0"}
|
new_metadata = {"source": "updated", "version": "2.0"}
|
||||||
updated_doc = await client.update_document_fields(
|
updated_doc = await client.update_document(
|
||||||
document_id=original_id, metadata=new_metadata
|
document_id=original_id, metadata=new_metadata
|
||||||
)
|
)
|
||||||
assert updated_doc.metadata == new_metadata
|
assert updated_doc.metadata == new_metadata
|
||||||
|
|
@ -118,7 +118,7 @@ async def test_client_update_document_fields(qa_corpus: Dataset, temp_db_path):
|
||||||
) # Should keep previous update
|
) # Should keep previous update
|
||||||
|
|
||||||
# Test updating only title
|
# Test updating only title
|
||||||
updated_doc = await client.update_document_fields(
|
updated_doc = await client.update_document(
|
||||||
document_id=original_id, title="New Title"
|
document_id=original_id, title="New Title"
|
||||||
)
|
)
|
||||||
assert updated_doc.title == "New Title"
|
assert updated_doc.title == "New Title"
|
||||||
|
|
@ -130,7 +130,7 @@ async def test_client_update_document_fields(qa_corpus: Dataset, temp_db_path):
|
||||||
Chunk(content="Custom chunk 1", order=0),
|
Chunk(content="Custom chunk 1", order=0),
|
||||||
Chunk(content="Custom chunk 2", order=1),
|
Chunk(content="Custom chunk 2", order=1),
|
||||||
]
|
]
|
||||||
updated_doc = await client.update_document_fields(
|
updated_doc = await client.update_document(
|
||||||
document_id=original_id,
|
document_id=original_id,
|
||||||
content="Content with custom chunks",
|
content="Content with custom chunks",
|
||||||
title="Final Title",
|
title="Final Title",
|
||||||
|
|
@ -1262,34 +1262,15 @@ async def test_client_create_document_from_file_stores_docling_json(temp_db_path
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_update_document_stores_docling_json(temp_db_path):
|
async def test_client_update_document_stores_docling_json(temp_db_path):
|
||||||
"""Test that update_document stores DoclingDocument JSON."""
|
"""Test that update_document stores DoclingDocument JSON when content changes."""
|
||||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
# Create initial document
|
# Create initial document
|
||||||
doc = await client.create_document(content="Initial content")
|
doc = await client.create_document(content="Initial content")
|
||||||
assert doc.id is not None
|
assert doc.id is not None
|
||||||
original_json = doc.docling_document_json
|
original_json = doc.docling_document_json
|
||||||
|
|
||||||
# Update the document
|
# Update content via update_document
|
||||||
doc.content = "Updated content"
|
updated_doc = await client.update_document(
|
||||||
updated_doc = await client.update_document(doc)
|
|
||||||
|
|
||||||
assert updated_doc.docling_document_json is not None
|
|
||||||
assert updated_doc.docling_version is not None
|
|
||||||
# JSON should be different because content changed
|
|
||||||
assert updated_doc.docling_document_json != original_json
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_client_update_document_fields_stores_docling_json(temp_db_path):
|
|
||||||
"""Test that update_document_fields 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_json
|
|
||||||
|
|
||||||
# Update content via update_document_fields
|
|
||||||
updated_doc = await client.update_document_fields(
|
|
||||||
document_id=doc.id, content="New content via fields update"
|
document_id=doc.id, content="New content via fields update"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -1300,10 +1281,10 @@ async def test_client_update_document_fields_stores_docling_json(temp_db_path):
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_update_document_fields_with_custom_chunks_no_docling_json(
|
async def test_client_update_document_with_custom_chunks_no_docling_json(
|
||||||
temp_db_path,
|
temp_db_path,
|
||||||
):
|
):
|
||||||
"""Test that update_document_fields with custom chunks does not update docling JSON."""
|
"""Test that update_document with custom chunks does not update docling JSON."""
|
||||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
# Create initial document
|
# Create initial document
|
||||||
doc = await client.create_document(content="Initial content")
|
doc = await client.create_document(content="Initial content")
|
||||||
|
|
@ -1312,7 +1293,7 @@ async def test_client_update_document_fields_with_custom_chunks_no_docling_json(
|
||||||
|
|
||||||
# Update with custom chunks
|
# Update with custom chunks
|
||||||
custom_chunks = [Chunk(content="Custom chunk", order=0)]
|
custom_chunks = [Chunk(content="Custom chunk", order=0)]
|
||||||
updated_doc = await client.update_document_fields(
|
updated_doc = await client.update_document(
|
||||||
document_id=doc.id, content="New content", chunks=custom_chunks
|
document_id=doc.id, content="New content", chunks=custom_chunks
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -1321,7 +1302,7 @@ async def test_client_update_document_fields_with_custom_chunks_no_docling_json(
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_update_document_fields_content_docling_mutually_exclusive(
|
async def test_client_update_document_content_docling_mutually_exclusive(
|
||||||
temp_db_path,
|
temp_db_path,
|
||||||
):
|
):
|
||||||
"""Test that content and docling_document_json cannot both be provided."""
|
"""Test that content and docling_document_json cannot both be provided."""
|
||||||
|
|
@ -1338,7 +1319,7 @@ async def test_client_update_document_fields_content_docling_mutually_exclusive(
|
||||||
docling_doc.add_text(label=DocItemLabel.TEXT, text="Some text")
|
docling_doc.add_text(label=DocItemLabel.TEXT, text="Some text")
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="mutually exclusive"):
|
with pytest.raises(ValueError, match="mutually exclusive"):
|
||||||
await client.update_document_fields(
|
await client.update_document(
|
||||||
document_id=doc.id,
|
document_id=doc.id,
|
||||||
content="New content",
|
content="New content",
|
||||||
docling_document_json=docling_doc.model_dump_json(),
|
docling_document_json=docling_doc.model_dump_json(),
|
||||||
|
|
@ -1347,7 +1328,7 @@ async def test_client_update_document_fields_content_docling_mutually_exclusive(
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_update_document_fields_with_docling_rechunks(temp_db_path):
|
async def test_client_update_document_with_docling_rechunks(temp_db_path):
|
||||||
"""Test that providing docling_document_json without chunks triggers rechunk."""
|
"""Test that providing docling_document_json without chunks triggers rechunk."""
|
||||||
from docling_core.types.doc.document import DoclingDocument
|
from docling_core.types.doc.document import DoclingDocument
|
||||||
|
|
||||||
|
|
@ -1367,7 +1348,7 @@ async def test_client_update_document_fields_with_docling_rechunks(temp_db_path)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Update with docling document only - should rechunk from it
|
# Update with docling document only - should rechunk from it
|
||||||
updated_doc = await client.update_document_fields(
|
updated_doc = await client.update_document(
|
||||||
document_id=doc.id,
|
document_id=doc.id,
|
||||||
docling_document_json=docling_doc.model_dump_json(),
|
docling_document_json=docling_doc.model_dump_json(),
|
||||||
docling_version=docling_doc.version,
|
docling_version=docling_doc.version,
|
||||||
|
|
@ -1386,7 +1367,7 @@ async def test_client_update_document_fields_with_docling_rechunks(temp_db_path)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_update_document_fields_docling_with_chunks(temp_db_path):
|
async def test_client_update_document_docling_with_chunks(temp_db_path):
|
||||||
"""Test that providing both docling_document_json and chunks stores both."""
|
"""Test that providing both docling_document_json and chunks stores both."""
|
||||||
from docling_core.types.doc.document import DoclingDocument
|
from docling_core.types.doc.document import DoclingDocument
|
||||||
|
|
||||||
|
|
@ -1407,7 +1388,7 @@ async def test_client_update_document_fields_docling_with_chunks(temp_db_path):
|
||||||
Chunk(content="Custom chunk 2", order=1),
|
Chunk(content="Custom chunk 2", order=1),
|
||||||
]
|
]
|
||||||
|
|
||||||
updated_doc = await client.update_document_fields(
|
updated_doc = await client.update_document(
|
||||||
document_id=doc.id,
|
document_id=doc.id,
|
||||||
chunks=custom_chunks,
|
chunks=custom_chunks,
|
||||||
docling_document_json=docling_doc.model_dump_json(),
|
docling_document_json=docling_doc.model_dump_json(),
|
||||||
|
|
@ -1708,3 +1689,59 @@ async def test_client_chunk_empty_document(temp_db_path):
|
||||||
|
|
||||||
assert isinstance(chunks, list)
|
assert isinstance(chunks, list)
|
||||||
assert len(chunks) == 0
|
assert len(chunks) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_import_document_embeds_chunks_without_embeddings(temp_db_path):
|
||||||
|
"""Test that import_document embeds chunks that don't have embeddings."""
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
# 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(
|
||||||
|
content="Document with unembedded chunks",
|
||||||
|
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.asyncio
|
||||||
|
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"
|
||||||
|
|
|
||||||
|
|
@ -46,10 +46,11 @@ async def test_version_rollback_on_update_failure(temp_db_path):
|
||||||
client.chunk_repository.create = succeed_then_fail # type: ignore[method-assign]
|
client.chunk_repository.create = succeed_then_fail # type: ignore[method-assign]
|
||||||
|
|
||||||
# Attempt update
|
# Attempt update
|
||||||
created.content = "Updated content"
|
|
||||||
|
|
||||||
with pytest.raises(RuntimeError):
|
with pytest.raises(RuntimeError):
|
||||||
await client.update_document(created)
|
await client.update_document(
|
||||||
|
document_id=created.id, # type: ignore[arg-type]
|
||||||
|
content="Updated content",
|
||||||
|
)
|
||||||
|
|
||||||
# Content and chunks should remain the original
|
# Content and chunks should remain the original
|
||||||
persisted = await client.get_document_by_id(created.id) # type: ignore[arg-type]
|
persisted = await client.get_document_by_id(created.id) # type: ignore[arg-type]
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue