Add search to client
This commit is contained in:
parent
0b43c7dc7f
commit
6321cfb7a8
3 changed files with 100 additions and 3 deletions
39
README.md
39
README.md
|
|
@ -66,10 +66,44 @@ await client.update_document(doc)
|
||||||
# Delete document
|
# Delete document
|
||||||
await client.delete_document(doc.id)
|
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
|
# Clean up
|
||||||
client.close()
|
client.close()
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Search Functionality
|
||||||
|
|
||||||
|
`haiku.rag` provides hybrid search combining vector similarity and full-text search:
|
||||||
|
1. **Vector Search**: Uses embeddings to find semantically similar content
|
||||||
|
2. **Full-text Search**: Uses SQLite FTS5 for exact keyword matching
|
||||||
|
3. **Hybrid Ranking**: Combines both using Reciprocal Rank Fusion (RRF)
|
||||||
|
4. **Chunked Results**: Returns relevant document chunks with scores
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 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
|
||||||
|
)
|
||||||
|
|
||||||
|
# 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
|
## Smart Document Updates
|
||||||
|
|
||||||
The system automatically tracks file changes using MD5 hashes:
|
The system automatically tracks file changes using MD5 hashes:
|
||||||
|
|
@ -119,6 +153,5 @@ print(doc.metadata)
|
||||||
2. Create a feature branch
|
2. Create a feature branch
|
||||||
3. Add tests for new functionality
|
3. Add tests for new functionality
|
||||||
4. Ensure all tests pass: `pytest`
|
4. Ensure all tests pass: `pytest`
|
||||||
5. Run type checking: `pyright`
|
5. Run type checking & linting with `pyright` & `ruff check`
|
||||||
6. Run linting: `ruff check`
|
6. Submit a pull request
|
||||||
7. Submit a pull request
|
|
||||||
|
|
|
||||||
|
|
@ -9,7 +9,9 @@ import httpx
|
||||||
|
|
||||||
from haiku.rag.reader import FileReader
|
from haiku.rag.reader import FileReader
|
||||||
from haiku.rag.store.engine import Store
|
from haiku.rag.store.engine import Store
|
||||||
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
from haiku.rag.store.repositories.chunk import ChunkRepository
|
||||||
from haiku.rag.store.repositories.document import DocumentRepository
|
from haiku.rag.store.repositories.document import DocumentRepository
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -20,6 +22,7 @@ class HaikuRAG:
|
||||||
"""Initialize the RAG client with a database path."""
|
"""Initialize the RAG client with a database path."""
|
||||||
self.store = Store(db_path)
|
self.store = Store(db_path)
|
||||||
self.document_repository = DocumentRepository(self.store)
|
self.document_repository = DocumentRepository(self.store)
|
||||||
|
self.chunk_repository = ChunkRepository(self.store)
|
||||||
|
|
||||||
async def create_document(
|
async def create_document(
|
||||||
self, content: str, uri: str | None = None, metadata: dict | None = None
|
self, content: str, uri: str | None = None, metadata: dict | None = None
|
||||||
|
|
@ -221,6 +224,21 @@ class HaikuRAG:
|
||||||
"""List all documents with optional pagination."""
|
"""List all documents with optional pagination."""
|
||||||
return await self.document_repository.list_all(limit=limit, offset=offset)
|
return await self.document_repository.list_all(limit=limit, offset=offset)
|
||||||
|
|
||||||
|
async def search(
|
||||||
|
self, query: str, limit: int = 5, k: int = 60
|
||||||
|
) -> list[tuple[Chunk, float]]:
|
||||||
|
"""Search for relevant chunks using hybrid search (vector similarity + full-text search).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: The search query string
|
||||||
|
limit: Maximum number of results to return
|
||||||
|
k: Parameter for Reciprocal Rank Fusion (default: 60)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of (chunk, score) tuples ordered by relevance
|
||||||
|
"""
|
||||||
|
return await self.chunk_repository.search_chunks_hybrid(query, limit, k)
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
"""Close the underlying store connection."""
|
"""Close the underlying store connection."""
|
||||||
self.store.close()
|
self.store.close()
|
||||||
|
|
|
||||||
|
|
@ -426,3 +426,49 @@ async def test_client_url_create_update_no_op_behavior():
|
||||||
assert doc3.content == updated_content.decode() # Updated content
|
assert doc3.content == updated_content.decode() # Updated content
|
||||||
|
|
||||||
client.close()
|
client.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_client_search():
|
||||||
|
"""Test HaikuRAG search functionality."""
|
||||||
|
client = HaikuRAG(":memory:")
|
||||||
|
|
||||||
|
# Add multiple documents to search from
|
||||||
|
doc1_text = "Python is a high-level programming language known for its simplicity and readability."
|
||||||
|
doc2_text = "Machine learning algorithms help computers learn patterns from data without explicit programming."
|
||||||
|
doc3_text = "Data science combines statistics, programming, and domain expertise to extract insights."
|
||||||
|
|
||||||
|
# Create documents
|
||||||
|
doc1 = await client.create_document(
|
||||||
|
content=doc1_text, uri="doc1.txt", metadata={"topic": "python"}
|
||||||
|
)
|
||||||
|
doc2 = await client.create_document(
|
||||||
|
content=doc2_text, uri="doc2.txt", metadata={"topic": "ml"}
|
||||||
|
)
|
||||||
|
await client.create_document(
|
||||||
|
content=doc3_text, uri="doc3.txt", metadata={"topic": "data_science"}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Test search with keyword that should match doc1
|
||||||
|
results = await client.search("Python programming", limit=3)
|
||||||
|
|
||||||
|
assert len(results) > 0
|
||||||
|
assert all(len(result) == 2 for result in results)
|
||||||
|
|
||||||
|
# Verify first result is from the Python document (doc1)
|
||||||
|
first_chunk, _ = results[0]
|
||||||
|
assert first_chunk.document_id == doc1.id
|
||||||
|
|
||||||
|
# Test search with different query
|
||||||
|
ml_results = await client.search("machine learning data", limit=2)
|
||||||
|
assert len(ml_results) > 0
|
||||||
|
|
||||||
|
# Verify first result is from the machine learning document (doc2)
|
||||||
|
first_ml_chunk, _ = ml_results[0]
|
||||||
|
assert first_ml_chunk.document_id == doc2.id
|
||||||
|
|
||||||
|
# Test search with limit parameter
|
||||||
|
limited_results = await client.search("programming", limit=1)
|
||||||
|
assert len(limited_results) <= 1
|
||||||
|
|
||||||
|
client.close()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue