Remove create_chunks_for_document from repositories
This commit is contained in:
parent
17c7147a49
commit
16ba3a9962
3 changed files with 87 additions and 160 deletions
|
|
@ -180,11 +180,18 @@ class HaikuRAG:
|
|||
|
||||
Returns:
|
||||
List of Chunk objects (without embeddings, without document_id).
|
||||
Each chunk has its `order` field set to its position in the list.
|
||||
"""
|
||||
from haiku.rag.chunkers import get_chunker
|
||||
|
||||
chunker = get_chunker(self._config)
|
||||
return await chunker.chunk(docling_document)
|
||||
chunks = await chunker.chunk(docling_document)
|
||||
|
||||
# Set order for each chunk
|
||||
for i, chunk in enumerate(chunks):
|
||||
chunk.order = i
|
||||
|
||||
return chunks
|
||||
|
||||
async def _store_document_with_chunks(
|
||||
self,
|
||||
|
|
@ -1298,37 +1305,34 @@ class HaikuRAG:
|
|||
self, documents: list[Document]
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Re-embed all chunks without changing chunk boundaries."""
|
||||
from haiku.rag.embeddings import contextualize
|
||||
|
||||
for doc in documents:
|
||||
assert doc.id is not None
|
||||
|
||||
# Get raw chunk records directly from LanceDB
|
||||
chunk_records = list(
|
||||
self.store.chunks_table.search()
|
||||
.where(f"document_id = '{doc.id}'")
|
||||
.to_pydantic(self.store.ChunkRecord)
|
||||
)
|
||||
if not chunk_records:
|
||||
# Get existing chunks
|
||||
chunks = await self.chunk_repository.get_by_document_id(doc.id)
|
||||
if not chunks:
|
||||
continue
|
||||
|
||||
# Batch embed all chunk contents
|
||||
contents = [rec.content for rec in chunk_records]
|
||||
embeddings = await self.chunk_repository.embedder.embed(contents)
|
||||
# Generate new embeddings using contextualize for consistency
|
||||
texts = contextualize(chunks)
|
||||
embeddings = await self.chunk_repository.embedder.embed(texts)
|
||||
|
||||
# Build updated records only for chunks with changed embeddings
|
||||
# Build updated records
|
||||
updated_records = [
|
||||
self.store.ChunkRecord(
|
||||
id=rec.id,
|
||||
document_id=rec.document_id,
|
||||
content=rec.content,
|
||||
metadata=rec.metadata,
|
||||
order=rec.order,
|
||||
id=chunk.id, # type: ignore[arg-type]
|
||||
document_id=chunk.document_id, # type: ignore[arg-type]
|
||||
content=chunk.content,
|
||||
metadata=json.dumps(chunk.metadata),
|
||||
order=chunk.order,
|
||||
vector=embedding,
|
||||
)
|
||||
for rec, embedding in zip(chunk_records, embeddings)
|
||||
if rec.vector != embedding
|
||||
for chunk, embedding in zip(chunks, embeddings)
|
||||
]
|
||||
|
||||
# Batch update chunks with changed embeddings
|
||||
# Batch update all chunks
|
||||
if updated_records:
|
||||
self.store.chunks_table.merge_insert(
|
||||
"id"
|
||||
|
|
@ -1340,76 +1344,68 @@ class HaikuRAG:
|
|||
self, documents: list[Document]
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Re-chunk and re-embed from existing document content."""
|
||||
converter = get_converter(self._config)
|
||||
from haiku.rag.embeddings import embed_chunks
|
||||
|
||||
for doc in documents:
|
||||
assert doc.id is not None
|
||||
docling_document = await converter.convert_text(doc.content)
|
||||
|
||||
# Update document with docling JSON
|
||||
# Convert content to DoclingDocument
|
||||
docling_document = await self.convert(doc.content)
|
||||
|
||||
# Chunk and embed
|
||||
chunks = await self.chunk(docling_document)
|
||||
embedded_chunks = await embed_chunks(chunks, self._config)
|
||||
|
||||
# Update document with docling JSON and store new chunks
|
||||
doc.docling_document_json = docling_document.model_dump_json()
|
||||
doc.docling_version = docling_document.version
|
||||
await self.document_repository.update(doc)
|
||||
await self._update_document_with_chunks(doc, embedded_chunks)
|
||||
|
||||
await self.chunk_repository.create_chunks_for_document(
|
||||
doc.id, docling_document
|
||||
)
|
||||
yield doc.id
|
||||
|
||||
async def _rebuild_full(
|
||||
self, documents: list[Document]
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Full rebuild: re-convert from source, re-chunk, re-embed."""
|
||||
converter = get_converter(self._config)
|
||||
from haiku.rag.embeddings import embed_chunks
|
||||
|
||||
for doc in documents:
|
||||
assert doc.id is not None
|
||||
|
||||
# Try to rebuild from source if available
|
||||
if doc.uri and self._check_source_accessible(doc.uri):
|
||||
try:
|
||||
await self.delete_document(doc.id)
|
||||
new_doc = await self.create_document_from_source(
|
||||
source=doc.uri, metadata=doc.metadata or {}
|
||||
)
|
||||
assert isinstance(new_doc, Document)
|
||||
assert new_doc.id is not None
|
||||
yield new_doc.id
|
||||
continue
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error recreating document from source %s: %s",
|
||||
doc.uri,
|
||||
e,
|
||||
)
|
||||
continue
|
||||
|
||||
# Fallback: rebuild from stored content
|
||||
if doc.uri:
|
||||
source_accessible = self._check_source_accessible(doc.uri)
|
||||
|
||||
if source_accessible:
|
||||
try:
|
||||
await self.delete_document(doc.id)
|
||||
new_doc = await self.create_document_from_source(
|
||||
source=doc.uri, metadata=doc.metadata or {}
|
||||
)
|
||||
assert isinstance(new_doc, Document)
|
||||
assert new_doc.id is not None
|
||||
yield new_doc.id
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error recreating document from source %s: %s",
|
||||
doc.uri,
|
||||
e,
|
||||
)
|
||||
continue
|
||||
else:
|
||||
logger.warning(
|
||||
"Source missing for %s, re-embedding from content", doc.uri
|
||||
)
|
||||
docling_document = await converter.convert_text(doc.content)
|
||||
|
||||
# Update document with docling JSON
|
||||
doc.docling_document_json = docling_document.model_dump_json()
|
||||
doc.docling_version = docling_document.version
|
||||
await self.document_repository.update(doc)
|
||||
|
||||
await self.chunk_repository.create_chunks_for_document(
|
||||
doc.id, docling_document
|
||||
)
|
||||
yield doc.id
|
||||
else:
|
||||
docling_document = await converter.convert_text(doc.content)
|
||||
|
||||
# Update document with docling JSON
|
||||
doc.docling_document_json = docling_document.model_dump_json()
|
||||
doc.docling_version = docling_document.version
|
||||
await self.document_repository.update(doc)
|
||||
|
||||
await self.chunk_repository.create_chunks_for_document(
|
||||
doc.id, docling_document
|
||||
logger.warning(
|
||||
"Source missing for %s, re-embedding from content", doc.uri
|
||||
)
|
||||
yield doc.id
|
||||
|
||||
docling_document = await self.convert(doc.content)
|
||||
chunks = await self.chunk(docling_document)
|
||||
embedded_chunks = await embed_chunks(chunks, self._config)
|
||||
|
||||
doc.docling_document_json = docling_document.model_dump_json()
|
||||
doc.docling_version = docling_document.version
|
||||
await self._update_document_with_chunks(doc, embedded_chunks)
|
||||
|
||||
yield doc.id
|
||||
|
||||
def _check_source_accessible(self, uri: str) -> bool:
|
||||
"""Check if a document's source URI is accessible."""
|
||||
|
|
|
|||
|
|
@ -16,9 +16,6 @@ from lancedb.rerankers import RRFReranker
|
|||
from haiku.rag.store.engine import DocumentRecord, Store
|
||||
from haiku.rag.store.models.chunk import Chunk
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from docling_core.types.doc.document import DoclingDocument
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
|
|
@ -190,55 +187,6 @@ class ChunkRepository:
|
|||
)
|
||||
return chunks
|
||||
|
||||
async def create_chunks_for_document(
|
||||
self, document_id: str, document: "DoclingDocument"
|
||||
) -> list[Chunk]:
|
||||
"""Create chunks and embeddings for a document from DoclingDocument."""
|
||||
from haiku.rag.chunkers import get_chunker
|
||||
|
||||
chunker = get_chunker(self.store._config)
|
||||
chunks = await chunker.chunk(document)
|
||||
|
||||
# Build embedding texts with headings prepended for better semantic search
|
||||
# The stored content stays raw, but embeddings capture section context
|
||||
embedding_texts = []
|
||||
for chunk in chunks:
|
||||
chunk_meta = chunk.get_chunk_metadata()
|
||||
if chunk_meta.headings:
|
||||
embedding_text = "\n".join(chunk_meta.headings) + "\n" + chunk.content
|
||||
else:
|
||||
embedding_text = chunk.content
|
||||
embedding_texts.append(embedding_text)
|
||||
embeddings = await self.embedder.embed(embedding_texts)
|
||||
|
||||
# Prepare all chunk records for batch insertion
|
||||
chunk_records = []
|
||||
created_chunks = []
|
||||
|
||||
for order, (chunk, embedding) in enumerate(zip(chunks, embeddings)):
|
||||
chunk_id = str(uuid4())
|
||||
|
||||
chunk_record = self.store.ChunkRecord(
|
||||
id=chunk_id,
|
||||
document_id=document_id,
|
||||
content=chunk.content,
|
||||
metadata=json.dumps(chunk.metadata),
|
||||
order=order,
|
||||
vector=embedding,
|
||||
)
|
||||
chunk_records.append(chunk_record)
|
||||
|
||||
chunk.id = chunk_id
|
||||
chunk.document_id = document_id
|
||||
chunk.order = order
|
||||
created_chunks.append(chunk)
|
||||
|
||||
# Batch insert all chunks at once
|
||||
if chunk_records:
|
||||
self.store.chunks_table.add(chunk_records)
|
||||
|
||||
return created_chunks
|
||||
|
||||
async def delete_all(self) -> None:
|
||||
"""Delete all chunks from the database."""
|
||||
# Drop and recreate table to clear all data
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ from datasets import Dataset
|
|||
|
||||
from haiku.rag.client import HaikuRAG
|
||||
from haiku.rag.config import Config
|
||||
from haiku.rag.converters import get_converter
|
||||
from haiku.rag.store.engine import Store
|
||||
from haiku.rag.store.models.chunk import Chunk, ChunkMetadata
|
||||
from haiku.rag.store.models.document import Document
|
||||
|
|
@ -53,45 +52,29 @@ async def test_chunk_repository_operations(qa_corpus: Dataset, temp_db_path):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_chunks_for_document(qa_corpus: Dataset, temp_db_path):
|
||||
"""Test creating chunks for a document."""
|
||||
# Create a store and repositories
|
||||
store = Store(temp_db_path, create=True)
|
||||
chunk_repo = ChunkRepository(store)
|
||||
doc_repo = DocumentRepository(store)
|
||||
async def test_chunking_pipeline(qa_corpus: Dataset, temp_db_path):
|
||||
"""Test document chunking using client primitives."""
|
||||
from haiku.rag.client import HaikuRAG
|
||||
from haiku.rag.embeddings import embed_chunks
|
||||
|
||||
# Get the first document from the corpus
|
||||
first_doc = qa_corpus[0]
|
||||
document_text = first_doc["document_extracted"]
|
||||
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
|
||||
# Get the first document from the corpus
|
||||
first_doc = qa_corpus[0]
|
||||
document_text = first_doc["document_extracted"]
|
||||
|
||||
# Create a document first (without chunks)
|
||||
document = Document(content=document_text, metadata={"source": "test"})
|
||||
created_document = await doc_repo.create(document)
|
||||
document_id = created_document.id
|
||||
# Use client primitives: convert → chunk → embed
|
||||
docling_document = await client.convert(document_text)
|
||||
chunks = await client.chunk(docling_document)
|
||||
embedded_chunks = await embed_chunks(chunks)
|
||||
|
||||
assert document_id is not None, "Document ID should not be None"
|
||||
# Verify chunks were created with embeddings
|
||||
assert len(chunks) > 0
|
||||
assert all(chunk.embedding is None for chunk in chunks) # Before embedding
|
||||
assert all(chunk.embedding is not None for chunk in embedded_chunks) # After
|
||||
|
||||
# Convert text to DoclingDocument
|
||||
converter = get_converter(Config)
|
||||
docling_document = await converter.convert_text(document_text, name="test.md")
|
||||
|
||||
# Test creating chunks for the document
|
||||
chunks = await chunk_repo.create_chunks_for_document(document_id, docling_document)
|
||||
|
||||
# Verify chunks were created
|
||||
assert len(chunks) > 0
|
||||
assert all(chunk.document_id == document_id for chunk in chunks)
|
||||
assert all(chunk.id is not None for chunk in chunks)
|
||||
|
||||
# Verify chunk order
|
||||
for i, chunk in enumerate(chunks):
|
||||
assert chunk.order == i
|
||||
|
||||
# Verify chunks exist in database
|
||||
db_chunks = await chunk_repo.get_by_document_id(document_id)
|
||||
assert len(db_chunks) == len(chunks)
|
||||
|
||||
store.close()
|
||||
# Verify chunk order
|
||||
for i, chunk in enumerate(chunks):
|
||||
assert chunk.order == i
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Reference in a new issue