haiku.rag/tests/test_document.py

140 lines
4.6 KiB
Python

import pytest
from datasets import Dataset
from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config
from haiku.rag.store.engine import Store
from haiku.rag.store.models.document import Document
from haiku.rag.store.repositories.document import DocumentRepository
@pytest.mark.asyncio
async def test_create_document_with_chunks(qa_corpus: Dataset, temp_db_path):
"""Test creating a document with chunks from the qa_corpus using repository."""
# Create client
client = HaikuRAG(db_path=temp_db_path, config=Config, create=True)
# Get the first document from the corpus
first_doc = qa_corpus[0]
document_text = first_doc["document_extracted"]
# Create the document with chunks in the database
created_document = await client.create_document(
content=document_text,
metadata={"source": "qa_corpus", "topic": first_doc.get("document_topic", "")},
)
# Verify the document was created
assert created_document.id is not None
assert created_document.content == document_text
# Check that chunks were created using repository
chunks = await client.chunk_repository.get_by_document_id(created_document.id)
assert len(chunks) > 0
# Verify chunk order is set correctly
for i, chunk in enumerate(chunks):
assert chunk.order == i
client.close()
@pytest.mark.asyncio
async def test_document_repository_crud(qa_corpus: Dataset, temp_db_path):
"""Test CRUD operations in DocumentRepository."""
# Create a store and repository
store = Store(temp_db_path, create=True)
doc_repo = DocumentRepository(store)
# Get the first document from the corpus
first_doc = qa_corpus[0]
document_text = first_doc["document_extracted"]
# Create a document with URI
test_uri = "file:///path/to/test.txt"
document = Document(
content=document_text, uri=test_uri, metadata={"source": "test"}
)
created_document = await doc_repo.create(document)
# Test get_by_id
assert created_document.id is not None
retrieved_document = await doc_repo.get_by_id(created_document.id)
assert retrieved_document is not None
assert retrieved_document.content == document_text
assert retrieved_document.uri == test_uri
# Test get_by_uri
retrieved_by_uri = await doc_repo.get_by_uri(test_uri)
assert retrieved_by_uri is not None
assert retrieved_by_uri.id == created_document.id
assert retrieved_by_uri.content == document_text
assert retrieved_by_uri.uri == test_uri
# Test get_by_uri with non-existent URI
non_existent = await doc_repo.get_by_uri("file:///non/existent.txt")
assert non_existent is None
# Test update (should regenerate chunks)
retrieved_document.content = "Updated content for testing"
updated_document = await doc_repo.update(retrieved_document)
assert updated_document.content == "Updated content for testing"
# Test list_all
all_documents = await doc_repo.list_all()
assert len(all_documents) == 1
assert all_documents[0].id == created_document.id
# Test delete
deleted = await doc_repo.delete(created_document.id)
assert deleted is True
# Verify document is gone
retrieved_document = await doc_repo.get_by_id(created_document.id)
assert retrieved_document is None
store.close()
@pytest.mark.asyncio
async def test_document_list_with_filter(qa_corpus: Dataset, temp_db_path):
"""Test listing documents with filter clause."""
store = Store(temp_db_path, create=True)
doc_repo = DocumentRepository(store)
first_doc = qa_corpus[0]
document_text = first_doc["document_extracted"]
doc1 = Document(
content=document_text,
uri="https://example.com/doc1.txt",
metadata={"source": "test", "category": "A"},
)
doc2 = Document(
content=document_text,
uri="https://arxiv.org/paper.pdf",
metadata={"source": "test", "category": "B"},
)
doc3 = Document(
content=document_text,
uri="https://example.com/doc3.txt",
metadata={"source": "test", "category": "A"},
)
created_doc1 = await doc_repo.create(doc1)
created_doc2 = await doc_repo.create(doc2)
created_doc3 = await doc_repo.create(doc3)
all_documents = await doc_repo.list_all()
assert len(all_documents) == 3
arxiv_documents = await doc_repo.list_all(filter="uri LIKE '%arxiv%'")
assert len(arxiv_documents) == 1
assert arxiv_documents[0].id == created_doc2.id
example_documents = await doc_repo.list_all(filter="uri LIKE '%example.com%'")
assert len(example_documents) == 2
assert {doc.id for doc in example_documents} == {created_doc1.id, created_doc3.id}
store.close()