Turn rebuild into a generator, track progress in command line
This commit is contained in:
parent
11d9f2701e
commit
85443ee963
3 changed files with 33 additions and 4 deletions
|
|
@ -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
|
||||||
|
|
@ -75,7 +76,23 @@ class HaikuRAGApp:
|
||||||
async def rebuild(self):
|
async def rebuild(self):
|
||||||
async with HaikuRAG(db_path=self.db_path) as client:
|
async with HaikuRAG(db_path=self.db_path) as client:
|
||||||
try:
|
try:
|
||||||
await client.rebuild_database()
|
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]")
|
self.console.print("[b]Database rebuild completed successfully.[/b]")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.console.print(f"[red]Error rebuilding database: {e}[/red]")
|
self.console.print(f"[red]Error rebuilding database: {e}[/red]")
|
||||||
|
|
|
||||||
|
|
@ -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,8 +271,12 @@ 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) -> None:
|
async def rebuild_database(self) -> AsyncGenerator[int, None]:
|
||||||
"""Rebuild the database by deleting all chunks and re-indexing all documents."""
|
"""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()
|
documents = await self.list_documents()
|
||||||
|
|
||||||
if not documents:
|
if not documents:
|
||||||
|
|
@ -284,6 +289,7 @@ class HaikuRAG:
|
||||||
await self.chunk_repository.create_chunks_for_document(
|
await self.chunk_repository.create_chunks_for_document(
|
||||||
doc.id, doc.content, commit=False
|
doc.id, doc.content, commit=False
|
||||||
)
|
)
|
||||||
|
yield doc.id
|
||||||
|
|
||||||
if self.store._connection:
|
if self.store._connection:
|
||||||
self.store._connection.commit()
|
self.store._connection.commit()
|
||||||
|
|
|
||||||
|
|
@ -29,7 +29,13 @@ async def test_rebuild_database(qa_corpus: Dataset):
|
||||||
assert len(chunks_before) > 0
|
assert len(chunks_before) > 0
|
||||||
|
|
||||||
# Perform rebuild
|
# Perform rebuild
|
||||||
await client.rebuild_database()
|
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()
|
documents_after = await client.list_documents()
|
||||||
assert len(documents_after) == 3
|
assert len(documents_after) == 3
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue