Merge pull request #11 from ggozad/feat/rebuild
Rebuild database, client & cli command
This commit is contained in:
commit
7f862a057a
8 changed files with 152 additions and 0 deletions
|
|
@ -31,6 +31,9 @@ haiku-rag search "query"
|
||||||
# Ask questions
|
# Ask questions
|
||||||
haiku-rag ask "Who is the author of haiku.rag?"
|
haiku-rag ask "Who is the author of haiku.rag?"
|
||||||
|
|
||||||
|
# Rebuild database (re-chunk and re-embed all documents)
|
||||||
|
haiku-rag rebuild
|
||||||
|
|
||||||
# Start server with file monitoring
|
# Start server with file monitoring
|
||||||
export MONITOR_DIRECTORIES="/path/to/docs"
|
export MONITOR_DIRECTORIES="/path/to/docs"
|
||||||
haiku-rag serve
|
haiku-rag serve
|
||||||
|
|
|
||||||
10
docs/cli.md
10
docs/cli.md
|
|
@ -35,6 +35,16 @@ haiku-rag get 1
|
||||||
haiku-rag delete 1
|
haiku-rag delete 1
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### Rebuild Database
|
||||||
|
|
||||||
|
Rebuild the database by deleting all chunks & embeddings and re-indexing all documents:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
haiku-rag rebuild
|
||||||
|
```
|
||||||
|
|
||||||
|
Use this when you want to change things like the embedding model or chunk size for example.
|
||||||
|
|
||||||
## Search
|
## Search
|
||||||
|
|
||||||
Basic search:
|
Basic search:
|
||||||
|
|
|
||||||
|
|
@ -67,6 +67,13 @@ await client.update_document(doc)
|
||||||
await client.delete_document(doc.id)
|
await client.delete_document(doc.id)
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### Rebuilding the Database
|
||||||
|
|
||||||
|
```python
|
||||||
|
async for doc_id in client.rebuild_database():
|
||||||
|
print(f"Processed document {doc_id}")
|
||||||
|
```
|
||||||
|
|
||||||
## Searching Documents
|
## Searching Documents
|
||||||
|
|
||||||
Basic search:
|
Basic search:
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,7 @@ from pathlib import Path
|
||||||
|
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
from rich.markdown import Markdown
|
from rich.markdown import Markdown
|
||||||
|
from rich.progress import Progress
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
|
|
@ -72,6 +73,30 @@ class HaikuRAGApp:
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.console.print(f"[red]Error: {e}[/red]")
|
self.console.print(f"[red]Error: {e}[/red]")
|
||||||
|
|
||||||
|
async def rebuild(self):
|
||||||
|
async with HaikuRAG(db_path=self.db_path) as client:
|
||||||
|
try:
|
||||||
|
documents = await client.list_documents()
|
||||||
|
total_docs = len(documents)
|
||||||
|
|
||||||
|
if total_docs == 0:
|
||||||
|
self.console.print(
|
||||||
|
"[yellow]No documents found in database.[/yellow]"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
self.console.print(
|
||||||
|
f"[b]Rebuilding database with {total_docs} documents...[/b]"
|
||||||
|
)
|
||||||
|
with Progress() as progress:
|
||||||
|
task = progress.add_task("Rebuilding...", total=total_docs)
|
||||||
|
async for _ in client.rebuild_database():
|
||||||
|
progress.update(task, advance=1)
|
||||||
|
|
||||||
|
self.console.print("[b]Database rebuild completed successfully.[/b]")
|
||||||
|
except Exception as e:
|
||||||
|
self.console.print(f"[red]Error rebuilding database: {e}[/red]")
|
||||||
|
|
||||||
def _rich_print_document(self, doc: Document, truncate: bool = False):
|
def _rich_print_document(self, doc: Document, truncate: bool = False):
|
||||||
"""Format a document for display."""
|
"""Format a document for display."""
|
||||||
if truncate:
|
if truncate:
|
||||||
|
|
|
||||||
|
|
@ -128,6 +128,21 @@ def ask(
|
||||||
event_loop.run_until_complete(app.ask(question=question))
|
event_loop.run_until_complete(app.ask(question=question))
|
||||||
|
|
||||||
|
|
||||||
|
@cli.command(
|
||||||
|
"rebuild",
|
||||||
|
help="Rebuild the database by deleting all chunks and re-indexing all documents",
|
||||||
|
)
|
||||||
|
def rebuild(
|
||||||
|
db: Path = typer.Option(
|
||||||
|
get_default_data_dir() / "haiku.rag.sqlite",
|
||||||
|
"--db",
|
||||||
|
help="Path to the SQLite database file",
|
||||||
|
),
|
||||||
|
):
|
||||||
|
app = HaikuRAGApp(db_path=db)
|
||||||
|
event_loop.run_until_complete(app.rebuild())
|
||||||
|
|
||||||
|
|
||||||
@cli.command(
|
@cli.command(
|
||||||
"serve", help="Start the haiku.rag MCP server (by default in streamable HTTP mode)"
|
"serve", help="Start the haiku.rag MCP server (by default in streamable HTTP mode)"
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
import hashlib
|
import hashlib
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import tempfile
|
import tempfile
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
@ -270,6 +271,29 @@ class HaikuRAG:
|
||||||
qa_agent = get_qa_agent(self)
|
qa_agent = get_qa_agent(self)
|
||||||
return await qa_agent.answer(question)
|
return await qa_agent.answer(question)
|
||||||
|
|
||||||
|
async def rebuild_database(self) -> AsyncGenerator[int, None]:
|
||||||
|
"""Rebuild the database by deleting all chunks and re-indexing all documents.
|
||||||
|
|
||||||
|
Yields:
|
||||||
|
int: The ID of the document currently being processed
|
||||||
|
"""
|
||||||
|
documents = await self.list_documents()
|
||||||
|
|
||||||
|
if not documents:
|
||||||
|
return
|
||||||
|
|
||||||
|
await self.chunk_repository.delete_all()
|
||||||
|
|
||||||
|
for doc in documents:
|
||||||
|
if doc.id is not None:
|
||||||
|
await self.chunk_repository.create_chunks_for_document(
|
||||||
|
doc.id, doc.content, commit=False
|
||||||
|
)
|
||||||
|
yield doc.id
|
||||||
|
|
||||||
|
if self.store._connection:
|
||||||
|
self.store._connection.commit()
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
"""Close the underlying store connection."""
|
"""Close the underlying store connection."""
|
||||||
self.store.close()
|
self.store.close()
|
||||||
|
|
|
||||||
|
|
@ -208,6 +208,22 @@ class ChunkRepository(BaseRepository[Chunk]):
|
||||||
|
|
||||||
return created_chunks
|
return created_chunks
|
||||||
|
|
||||||
|
async def delete_all(self, commit: bool = True) -> bool:
|
||||||
|
"""Delete all chunks from the database."""
|
||||||
|
if self.store._connection is None:
|
||||||
|
raise ValueError("Store connection is not available")
|
||||||
|
|
||||||
|
cursor = self.store._connection.cursor()
|
||||||
|
|
||||||
|
cursor.execute("DELETE FROM chunks_fts")
|
||||||
|
cursor.execute("DELETE FROM chunk_embeddings")
|
||||||
|
cursor.execute("DELETE FROM chunks")
|
||||||
|
|
||||||
|
deleted = cursor.rowcount > 0
|
||||||
|
if commit:
|
||||||
|
self.store._connection.commit()
|
||||||
|
return deleted
|
||||||
|
|
||||||
async def delete_by_document_id(
|
async def delete_by_document_id(
|
||||||
self, document_id: int, commit: bool = True
|
self, document_id: int, commit: bool = True
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
|
|
||||||
52
tests/test_rebuild.py
Normal file
52
tests/test_rebuild.py
Normal file
|
|
@ -0,0 +1,52 @@
|
||||||
|
import pytest
|
||||||
|
from datasets import Dataset
|
||||||
|
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_rebuild_database(qa_corpus: Dataset):
|
||||||
|
"""Test rebuild functionality with existing documents."""
|
||||||
|
client = HaikuRAG(":memory:")
|
||||||
|
|
||||||
|
created_docs: list[Document] = []
|
||||||
|
for content in qa_corpus["document_extracted"][:3]:
|
||||||
|
doc = await client.create_document(
|
||||||
|
content=content,
|
||||||
|
)
|
||||||
|
created_docs.append(doc)
|
||||||
|
|
||||||
|
documents_before = await client.list_documents()
|
||||||
|
assert len(documents_before) == 3
|
||||||
|
|
||||||
|
chunks_before = []
|
||||||
|
for doc in created_docs:
|
||||||
|
assert doc.id is not None
|
||||||
|
doc_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
chunks_before.extend(doc_chunks)
|
||||||
|
|
||||||
|
assert len(chunks_before) > 0
|
||||||
|
|
||||||
|
# Perform rebuild
|
||||||
|
processed_doc_ids = []
|
||||||
|
async for doc_id in client.rebuild_database():
|
||||||
|
processed_doc_ids.append(doc_id)
|
||||||
|
|
||||||
|
# Verify all documents were processed
|
||||||
|
expected_doc_ids = [doc.id for doc in created_docs]
|
||||||
|
assert set(processed_doc_ids) == set(expected_doc_ids)
|
||||||
|
|
||||||
|
documents_after = await client.list_documents()
|
||||||
|
assert len(documents_after) == 3
|
||||||
|
|
||||||
|
# Verify chunks were recreated
|
||||||
|
chunks_after = []
|
||||||
|
for doc in documents_after:
|
||||||
|
if doc.id is not None:
|
||||||
|
doc_chunks = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
chunks_after.extend(doc_chunks)
|
||||||
|
|
||||||
|
assert len(chunks_after) > 0
|
||||||
|
|
||||||
|
client.close()
|
||||||
Loading…
Reference in a new issue