Basic client

This commit is contained in:
Yiorgis Gozadinos 2025-06-17 10:48:40 +02:00
parent 1ff17c2f50
commit a91249c575
No known key found for this signature in database
2 changed files with 126 additions and 0 deletions

52
src/haiku/rag/client.py Normal file
View file

@ -0,0 +1,52 @@
from pathlib import Path
from typing import Literal
from haiku.rag.store.engine import Store
from haiku.rag.store.models.document import Document
from haiku.rag.store.repositories.document import DocumentRepository
class RAGClient:
"""High-level haiku-rag client."""
def __init__(self, db_path: Path | Literal[":memory:"]):
"""Initialize the RAG client with a database path."""
self.store = Store(db_path)
self.document_repository = DocumentRepository(self.store)
async def create_document(
self, content: str, uri: str | None = None, metadata: dict | None = None
) -> Document:
"""Create a new document with optional URI and metadata."""
document = Document(
content=content,
uri=uri,
metadata=metadata or {},
)
return await self.document_repository.create(document)
async def get_document_by_id(self, document_id: int) -> Document | None:
"""Get a document by its ID."""
return await self.document_repository.get_by_id(document_id)
async def get_document_by_uri(self, uri: str) -> Document | None:
"""Get a document by its URI."""
return await self.document_repository.get_by_uri(uri)
async def update_document(self, document: Document) -> Document:
"""Update an existing document."""
return await self.document_repository.update(document)
async def delete_document(self, document_id: int) -> bool:
"""Delete a document by its ID."""
return await self.document_repository.delete(document_id)
async def list_documents(
self, limit: int | None = None, offset: int | None = None
) -> list[Document]:
"""List all documents with optional pagination."""
return await self.document_repository.list_all(limit=limit, offset=offset)
def close(self):
"""Close the underlying store connection."""
self.store.close()

74
tests/test_client.py Normal file
View file

@ -0,0 +1,74 @@
import pytest
from datasets import Dataset
from haiku.rag.client import RAGClient
@pytest.mark.asyncio
async def test_client_document_crud(qa_corpus: Dataset):
"""Test RAGClient CRUD operations for documents."""
# Create client with in-memory database
client = RAGClient(":memory:")
# 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
assert created_doc.content == document_text
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 retrieved_doc.content == document_text
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 retrieved_by_uri.content == document_text
# 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
retrieved_doc.content = "Updated content"
retrieved_doc.uri = "file:///updated/path.txt"
updated_doc = await client.update_document(retrieved_doc)
assert updated_doc.content == "Updated content"
assert updated_doc.uri == "file:///updated/path.txt"
# 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
client.close()