Add pagination to chunk repo's get_by_document_id
This commit is contained in:
parent
c43c8fab22
commit
448d00f9b2
2 changed files with 83 additions and 7 deletions
|
|
@ -281,13 +281,30 @@ class ChunkRepository:
|
||||||
results = results.limit(limit)
|
results = results.limit(limit)
|
||||||
return await self._process_search_results(results)
|
return await self._process_search_results(results)
|
||||||
|
|
||||||
async def get_by_document_id(self, document_id: str) -> list[Chunk]:
|
async def get_by_document_id(
|
||||||
"""Get all chunks for a specific document."""
|
self,
|
||||||
results = list(
|
document_id: str,
|
||||||
self.store.chunks_table.search()
|
limit: int | None = None,
|
||||||
.where(f"document_id = '{document_id}'")
|
offset: int | None = None,
|
||||||
.to_pydantic(self.store.ChunkRecord)
|
) -> list[Chunk]:
|
||||||
)
|
"""Get chunks for a specific document with optional pagination.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
document_id: The document ID to get chunks for.
|
||||||
|
limit: Maximum number of chunks to return. None for all.
|
||||||
|
offset: Number of chunks to skip. None for no offset.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of chunks ordered by their order field.
|
||||||
|
"""
|
||||||
|
query = self.store.chunks_table.search().where(f"document_id = '{document_id}'")
|
||||||
|
|
||||||
|
if offset is not None:
|
||||||
|
query = query.offset(offset)
|
||||||
|
if limit is not None:
|
||||||
|
query = query.limit(limit)
|
||||||
|
|
||||||
|
results = list(query.to_pydantic(self.store.ChunkRecord))
|
||||||
|
|
||||||
# Get document info
|
# Get document info
|
||||||
doc_results = list(
|
doc_results = list(
|
||||||
|
|
@ -320,6 +337,16 @@ class ChunkRepository:
|
||||||
chunks.sort(key=lambda c: c.order)
|
chunks.sort(key=lambda c: c.order)
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
|
async def count_by_document_id(self, document_id: str) -> int:
|
||||||
|
"""Count the number of chunks for a specific document."""
|
||||||
|
df = (
|
||||||
|
self.store.chunks_table.search()
|
||||||
|
.select(["id"])
|
||||||
|
.where(f"document_id = '{document_id}'")
|
||||||
|
.to_pandas()
|
||||||
|
)
|
||||||
|
return len(df)
|
||||||
|
|
||||||
async def get_adjacent_chunks(self, chunk: Chunk, num_adjacent: int) -> list[Chunk]:
|
async def get_adjacent_chunks(self, chunk: Chunk, num_adjacent: int) -> list[Chunk]:
|
||||||
"""Get adjacent chunks before and after the given chunk within the same document."""
|
"""Get adjacent chunks before and after the given chunk within the same document."""
|
||||||
assert chunk.document_id, "Document id is required for adjacent chunk finding"
|
assert chunk.document_id, "Document id is required for adjacent chunk finding"
|
||||||
|
|
|
||||||
|
|
@ -47,6 +47,55 @@ async def test_chunk_repository_operations(qa_corpus: Dataset, temp_db_path):
|
||||||
client.close()
|
client.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chunk_repository_pagination(qa_corpus: Dataset, temp_db_path):
|
||||||
|
"""Test ChunkRepository pagination with get_by_document_id and count_by_document_id."""
|
||||||
|
async with HaikuRAG(db_path=temp_db_path, config=Config, create=True) as client:
|
||||||
|
# Get the first document from the corpus (should produce multiple chunks)
|
||||||
|
first_doc = qa_corpus[0]
|
||||||
|
document_text = first_doc["document_extracted"]
|
||||||
|
|
||||||
|
# Create a document with chunks
|
||||||
|
created_document = await client.create_document(
|
||||||
|
content=document_text, metadata={"source": "test"}
|
||||||
|
)
|
||||||
|
assert created_document.id is not None
|
||||||
|
|
||||||
|
# Get total chunk count
|
||||||
|
total_count = await client.chunk_repository.count_by_document_id(
|
||||||
|
created_document.id
|
||||||
|
)
|
||||||
|
assert total_count > 0
|
||||||
|
|
||||||
|
# Get all chunks without pagination
|
||||||
|
all_chunks = await client.chunk_repository.get_by_document_id(
|
||||||
|
created_document.id
|
||||||
|
)
|
||||||
|
assert len(all_chunks) == total_count
|
||||||
|
|
||||||
|
# Test pagination with limit
|
||||||
|
limit = min(2, total_count)
|
||||||
|
first_batch = await client.chunk_repository.get_by_document_id(
|
||||||
|
created_document.id, limit=limit
|
||||||
|
)
|
||||||
|
assert len(first_batch) == limit
|
||||||
|
assert first_batch[0].id == all_chunks[0].id
|
||||||
|
|
||||||
|
# Test pagination with offset
|
||||||
|
if total_count > limit:
|
||||||
|
second_batch = await client.chunk_repository.get_by_document_id(
|
||||||
|
created_document.id, limit=limit, offset=limit
|
||||||
|
)
|
||||||
|
assert len(second_batch) <= limit
|
||||||
|
assert second_batch[0].id == all_chunks[limit].id
|
||||||
|
|
||||||
|
# Test offset beyond available chunks
|
||||||
|
empty_batch = await client.chunk_repository.get_by_document_id(
|
||||||
|
created_document.id, limit=10, offset=total_count + 100
|
||||||
|
)
|
||||||
|
assert len(empty_batch) == 0
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_chunking_pipeline(qa_corpus: Dataset, temp_db_path):
|
async def test_chunking_pipeline(qa_corpus: Dataset, temp_db_path):
|
||||||
"""Test document chunking using client primitives."""
|
"""Test document chunking using client primitives."""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue