From 1e0b6b18aa0c1bba6bc3aedd9d591e5771d95329 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Tue, 17 Jun 2025 12:48:26 +0200 Subject: [PATCH] Add a context manager to the client --- README.md | 132 +++++++++++++++++++++------------------- src/haiku/rag/client.py | 9 +++ tests/test_client.py | 25 ++++++++ 3 files changed, 102 insertions(+), 64 deletions(-) diff --git a/README.md b/README.md index 939fb63f..34dfd829 100644 --- a/README.md +++ b/README.md @@ -34,48 +34,50 @@ source .venv/bin/activate from pathlib import Path from haiku.rag.client import HaikuRAG -# Initialize client with database path -client = HaikuRAG("path/to/database.db") -# Or use in-memory database for testing +# Use as async context manager (recommended) +async with HaikuRAG("path/to/database.db") as client: + # Create document from text + doc = await client.create_document( + content="Your document content here", + uri="doc://example", + metadata={"source": "manual", "topic": "example"} + ) + + # Create document from file (auto-parses content) + doc = await client.create_document_from_source("path/to/document.pdf") + + # Create document from URL + doc = await client.create_document_from_source("https://example.com/article.html") + + # Retrieve documents + doc = await client.get_document_by_id(1) + doc = await client.get_document_by_uri("file:///path/to/document.pdf") + + # List all documents with pagination + docs = await client.list_documents(limit=10, offset=0) + + # Update document content + doc.content = "Updated content" + await client.update_document(doc) + + # Delete document + await client.delete_document(doc.id) + + # Search documents using hybrid search (vector + full-text) + results = await client.search("machine learning algorithms", limit=5) + for chunk, score in results: + print(f"Score: {score:.3f}") + print(f"Content: {chunk.content}") + print(f"Document ID: {chunk.document_id}") + print("---") + + +# Or use without the context manager. client = HaikuRAG(":memory:") - -# Create document from text -doc = await client.create_document( - content="Your document content here", - uri="doc://example", - metadata={"source": "manual", "topic": "example"} -) - -# Create document from file (auto-parses content) -doc = await client.create_document_from_source("path/to/document.pdf") - -# Create document from URL -doc = await client.create_document_from_source("https://example.com/article.html") - -# Retrieve documents -doc = await client.get_document_by_id(1) -doc = await client.get_document_by_uri("file:///path/to/document.pdf") - -# List all documents with pagination -docs = await client.list_documents(limit=10, offset=0) - -# Update document content -doc.content = "Updated content" -await client.update_document(doc) - -# Delete document -await client.delete_document(doc.id) - -# Search documents using hybrid search (vector + full-text) -results = await client.search("machine learning algorithms", limit=5) -for chunk, score in results: - print(f"Score: {score:.3f}") - print(f"Content: {chunk.content}") - print(f"Document ID: {chunk.document_id}") - print("---") - -# Clean up -client.close() +try: + # ... operations ... +finally: + client.close() ``` ## Search Functionality @@ -87,21 +89,22 @@ client.close() 4. **Chunked Results**: Returns relevant document chunks with scores ```python -# Basic search -results = await client.search("your query here") +async with HaikuRAG("database.db") as client: + # Basic search + results = await client.search("your query here") -# Search with custom parameters -results = await client.search( - query="machine learning", - limit=10, # Maximum results to return - k=60 # RRF parameter for reciprocal rank fusion -) + # Search with custom parameters + results = await client.search( + query="machine learning", + limit=10, # Maximum results to return + k=60 # RRF parameter for reciprocal rank fusion + ) -# Process results -for chunk, relevance_score in results: - print(f"Relevance: {relevance_score:.3f}") - print(f"Content: {chunk.content}") - print(f"From document: {chunk.document_id}") + # Process results + for chunk, relevance_score in results: + print(f"Relevance: {relevance_score:.3f}") + print(f"Content: {chunk.content}") + print(f"From document: {chunk.document_id}") ``` ## Smart Document Updates @@ -109,18 +112,19 @@ for chunk, relevance_score in results: The system automatically tracks file changes using MD5 hashes: ```python -# First call - creates new document -doc1 = await client.create_document_from_source("document.txt") +async with HaikuRAG("database.db") as client: + # First call - creates new document + doc1 = await client.create_document_from_source("document.txt") -# Second call - no changes, returns existing document (no processing) -doc2 = await client.create_document_from_source("document.txt") -assert doc1.id == doc2.id + # Second call - no changes, returns existing document (no processing) + doc2 = await client.create_document_from_source("document.txt") + assert doc1.id == doc2.id -# After file modification - automatically updates existing document -# File content changed... -doc3 = await client.create_document_from_source("document.txt") -assert doc1.id == doc3.id # Same document -assert doc3.content != doc1.content # Updated content + # After file modification - automatically updates existing document + # File content changed... + doc3 = await client.create_document_from_source("document.txt") + assert doc1.id == doc3.id # Same document + assert doc3.content != doc1.content # Updated content ``` ## Supported File Formats diff --git a/src/haiku/rag/client.py b/src/haiku/rag/client.py index fef81586..77693542 100644 --- a/src/haiku/rag/client.py +++ b/src/haiku/rag/client.py @@ -24,6 +24,15 @@ class HaikuRAG: self.document_repository = DocumentRepository(self.store) self.chunk_repository = ChunkRepository(self.store) + async def __aenter__(self): + """Async context manager entry.""" + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb): + """Async context manager exit.""" + self.close() + return False + async def create_document( self, content: str, uri: str | None = None, metadata: dict | None = None ) -> Document: diff --git a/tests/test_client.py b/tests/test_client.py index 14ae3a30..77430e24 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -472,3 +472,28 @@ async def test_client_search(): assert len(limited_results) <= 1 client.close() + + +@pytest.mark.asyncio +async def test_client_async_context_manager(): + """Test HaikuRAG as async context manager.""" + + # Test that context manager works and auto-closes + async with HaikuRAG(":memory:") as client: + # Create a document to ensure the client works + doc = await client.create_document( + content="Test content for context manager", + uri="test://context", + metadata={"test": "context_manager"}, + ) + + assert doc.id is not None + assert doc.content == "Test content for context manager" + + # Test search works within context + results = await client.search("Test content", limit=1) + assert len(results) > 0 + + # Context manager should have automatically closed the connection + # We can't easily test that the connection is closed without accessing internals, + # but the test passing means the context manager methods work correctly