From a91249c575823e4b41a2d5578f34ba45d491fdff Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Tue, 17 Jun 2025 10:48:40 +0200 Subject: [PATCH] Basic client --- src/haiku/rag/client.py | 52 +++++++++++++++++++++++++++++ tests/test_client.py | 74 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 126 insertions(+) create mode 100644 src/haiku/rag/client.py create mode 100644 tests/test_client.py diff --git a/src/haiku/rag/client.py b/src/haiku/rag/client.py new file mode 100644 index 00000000..01e348e7 --- /dev/null +++ b/src/haiku/rag/client.py @@ -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() diff --git a/tests/test_client.py b/tests/test_client.py new file mode 100644 index 00000000..7d46d78b --- /dev/null +++ b/tests/test_client.py @@ -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()