870 lines
33 KiB
Python
870 lines
33 KiB
Python
import asyncio
|
|
import json
|
|
import logging
|
|
from datetime import datetime
|
|
from pathlib import Path
|
|
from typing import TYPE_CHECKING
|
|
|
|
from rich.console import Console
|
|
from rich.markdown import Markdown
|
|
from rich.progress import (
|
|
BarColumn,
|
|
DownloadColumn,
|
|
Progress,
|
|
SpinnerColumn,
|
|
TaskID,
|
|
TextColumn,
|
|
TransferSpeedColumn,
|
|
)
|
|
from rich.syntax import Syntax
|
|
|
|
from haiku.rag.client import HaikuRAG, RebuildMode
|
|
from haiku.rag.config import AppConfig, Config
|
|
from haiku.rag.mcp import create_mcp_server
|
|
from haiku.rag.monitor import FileWatcher, S3Watcher
|
|
from haiku.rag.store.models.document import Document
|
|
|
|
if TYPE_CHECKING:
|
|
from haiku.rag.store.models import SearchResult
|
|
from haiku.rag.utils import format_bytes, format_citations_rich, get_package_versions
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class HaikuRAGApp: # pragma: no cover
|
|
def __init__(
|
|
self,
|
|
db_path: Path,
|
|
config: AppConfig = Config,
|
|
read_only: bool = False,
|
|
before: datetime | None = None,
|
|
):
|
|
self.db_path = db_path
|
|
self.config = config
|
|
self.read_only = read_only
|
|
self.before = before
|
|
self.console = Console()
|
|
|
|
from haiku.rag.store.engine import ConnectionMode
|
|
|
|
self._is_local = ConnectionMode.from_config(self.config) == ConnectionMode.LOCAL
|
|
self._display_path = self.db_path if self._is_local else self.config.lancedb.uri
|
|
|
|
async def init(self):
|
|
"""Initialize a new database."""
|
|
if self._is_local and self.db_path.exists():
|
|
self.console.print(
|
|
f"[yellow]Database already exists at {self.db_path}[/yellow]"
|
|
)
|
|
return
|
|
|
|
# Create the database
|
|
async with HaikuRAG(db_path=self.db_path, config=self.config, create=True):
|
|
pass
|
|
self.console.print(
|
|
f"[bold green]Database initialized at {self._display_path}[/bold green]"
|
|
)
|
|
|
|
async def info(self):
|
|
"""Display read-only information about the database without modifying it."""
|
|
|
|
from haiku.rag.store.engine import connect_lancedb, get_database_stats
|
|
from haiku.rag.store.upgrades import get_pending_upgrades
|
|
|
|
if self.before is not None:
|
|
self.console.print(
|
|
"[yellow]Note: --before is not supported by info; showing current state.[/yellow]"
|
|
)
|
|
|
|
# Basic: show path/URI
|
|
self.console.print("[bold]haiku.rag database info[/bold]")
|
|
self.console.print(
|
|
f" [repr.attrib_name]path[/repr.attrib_name]: {self._display_path}"
|
|
)
|
|
|
|
if self._is_local and not self.db_path.exists():
|
|
self.console.print("[red]Database path does not exist.[/red]")
|
|
return
|
|
|
|
# Connect directly. Don't go through Store so a database that is
|
|
# missing tables (e.g. pre-migration) still reports what it can.
|
|
db = await connect_lancedb(self.config, self.db_path)
|
|
stats = await get_database_stats(db)
|
|
|
|
if not any(entry["exists"] for entry in stats.values()):
|
|
self.console.print(
|
|
"[red]Database is empty. Use 'haiku-rag init' to initialize.[/red]"
|
|
)
|
|
return
|
|
|
|
versions = get_package_versions()
|
|
|
|
stored_version = "unknown"
|
|
embed_provider = "unknown"
|
|
embed_model = "unknown"
|
|
vector_dim = None
|
|
if stats["settings"]["exists"]:
|
|
settings_tbl = await db.open_table("settings")
|
|
rows = (
|
|
await settings_tbl.query().where("id = 'settings'").limit(1).to_arrow()
|
|
).to_pylist()
|
|
if rows:
|
|
raw = rows[0].get("settings") or "{}"
|
|
data = json.loads(raw) if isinstance(raw, str) else (raw or {})
|
|
stored_version = str(data.get("version", "unknown"))
|
|
embeddings = data.get("embeddings", {})
|
|
embed_model_obj = embeddings.get("model", {})
|
|
embed_provider = embed_model_obj.get("provider", "unknown")
|
|
embed_model = embed_model_obj.get("name", "unknown")
|
|
vector_dim = embed_model_obj.get("vector_dim")
|
|
|
|
self.console.print(
|
|
f" [repr.attrib_name]haiku.rag version (db)[/repr.attrib_name]: {stored_version}"
|
|
)
|
|
dim_part = f"{vector_dim}" if vector_dim is not None else "unknown"
|
|
self.console.print(
|
|
" [repr.attrib_name]embeddings[/repr.attrib_name]: "
|
|
f"{embed_provider}/{embed_model} (dim: {dim_part})"
|
|
)
|
|
|
|
# Per-table row counts and sizes. Missing required tables are
|
|
# reported as "absent" rather than raising.
|
|
for name in ("documents", "chunks", "document_items"):
|
|
entry = stats[name]
|
|
if entry["exists"]:
|
|
self.console.print(
|
|
f" [repr.attrib_name]{name}[/repr.attrib_name]: {entry['num_rows']} "
|
|
f"({format_bytes(entry['total_bytes'])})"
|
|
)
|
|
else:
|
|
self.console.print(
|
|
f" [repr.attrib_name]{name}[/repr.attrib_name]: [yellow]absent[/yellow]"
|
|
)
|
|
|
|
# Vector index information
|
|
if stats["chunks"]["exists"]:
|
|
num_chunks = stats["chunks"]["num_rows"]
|
|
if stats["chunks"].get("has_vector_index"):
|
|
num_indexed_rows = stats["chunks"].get("num_indexed_rows", 0)
|
|
num_unindexed_rows = stats["chunks"].get("num_unindexed_rows", 0)
|
|
self.console.print(
|
|
" [repr.attrib_name]vector index[/repr.attrib_name]: ✓ exists"
|
|
)
|
|
self.console.print(
|
|
f" [repr.attrib_name]indexed chunks[/repr.attrib_name]: {num_indexed_rows}"
|
|
)
|
|
if num_unindexed_rows > 0:
|
|
self.console.print(
|
|
f" [repr.attrib_name]unindexed chunks[/repr.attrib_name]: [yellow]{num_unindexed_rows}[/yellow] "
|
|
"(consider running: haiku-rag create-index)"
|
|
)
|
|
else:
|
|
self.console.print(
|
|
f" [repr.attrib_name]unindexed chunks[/repr.attrib_name]: {num_unindexed_rows}"
|
|
)
|
|
else:
|
|
if num_chunks >= 256:
|
|
self.console.print(
|
|
" [repr.attrib_name]vector index[/repr.attrib_name]: [yellow]✗ not created[/yellow] "
|
|
"(run: haiku-rag create-index)"
|
|
)
|
|
else:
|
|
self.console.print(
|
|
f" [repr.attrib_name]vector index[/repr.attrib_name]: ✗ not created "
|
|
f"(need {256 - num_chunks} more chunks)"
|
|
)
|
|
|
|
if stats["documents"]["exists"]:
|
|
self.console.print(
|
|
f" [repr.attrib_name]versions (documents)[/repr.attrib_name]: "
|
|
f"{stats['documents']['num_versions']}"
|
|
)
|
|
if stats["chunks"]["exists"]:
|
|
self.console.print(
|
|
f" [repr.attrib_name]versions (chunks)[/repr.attrib_name]: "
|
|
f"{stats['chunks']['num_versions']}"
|
|
)
|
|
|
|
# Migration status
|
|
pending = (
|
|
get_pending_upgrades(stored_version) if stored_version != "unknown" else []
|
|
)
|
|
self.console.rule()
|
|
if pending:
|
|
self.console.print(
|
|
f"[bold yellow]{len(pending)} migration(s) pending.[/bold yellow] "
|
|
"Run [cyan]haiku-rag migrate[/cyan] to upgrade."
|
|
)
|
|
for step in pending:
|
|
desc = step.description or ""
|
|
self.console.print(f" [yellow]→[/yellow] {step.version}: {desc}")
|
|
else:
|
|
self.console.print("[green]Database is up to date.[/green]")
|
|
|
|
self.console.rule()
|
|
self.console.print("[bold]Versions[/bold]")
|
|
self.console.print(
|
|
f" [repr.attrib_name]haiku.rag[/repr.attrib_name]: {versions['haiku_rag']}"
|
|
)
|
|
self.console.print(
|
|
f" [repr.attrib_name]lancedb[/repr.attrib_name]: {versions['lancedb']}"
|
|
)
|
|
self.console.print(
|
|
f" [repr.attrib_name]docling[/repr.attrib_name]: {versions['docling']}"
|
|
)
|
|
self.console.print(
|
|
f" [repr.attrib_name]pydantic-ai[/repr.attrib_name]: {versions['pydantic_ai']}"
|
|
)
|
|
self.console.print(
|
|
f" [repr.attrib_name]docling-document schema[/repr.attrib_name]: {versions['docling_document_schema']}"
|
|
)
|
|
|
|
async def history(self, table: str | None = None, limit: int | None = None):
|
|
"""Display version history for database tables.
|
|
|
|
Args:
|
|
table: Specific table to show history for (documents, chunks, settings).
|
|
If None, shows history for all tables.
|
|
limit: Maximum number of versions to show per table.
|
|
"""
|
|
from haiku.rag.store.engine import Store
|
|
|
|
if self._is_local and not self.db_path.exists():
|
|
self.console.print("[red]Database path does not exist.[/red]")
|
|
return
|
|
|
|
async with Store(
|
|
self.db_path,
|
|
config=self.config,
|
|
skip_validation=True,
|
|
read_only=True,
|
|
skip_migration_check=True,
|
|
before=self.before,
|
|
) as store:
|
|
tables = ["documents", "chunks", "settings"]
|
|
if table:
|
|
if table not in tables:
|
|
self.console.print(
|
|
f"[red]Unknown table: {table}. Must be one of: {', '.join(tables)}[/red]"
|
|
)
|
|
return
|
|
tables = [table]
|
|
|
|
self.console.print("[bold]Version History[/bold]")
|
|
|
|
for table_name in tables:
|
|
versions = await store.list_table_versions(table_name)
|
|
|
|
# Sort by version descending (newest first)
|
|
versions = sorted(versions, key=lambda v: v["version"], reverse=True)
|
|
|
|
if limit:
|
|
versions = versions[:limit]
|
|
|
|
self.console.print(f"\n[bold cyan]{table_name}[/bold cyan]")
|
|
|
|
if not versions:
|
|
self.console.print(" [dim]No versions found[/dim]")
|
|
continue
|
|
|
|
for v in versions:
|
|
version_num = v["version"]
|
|
timestamp = v["timestamp"]
|
|
self.console.print(
|
|
f" [repr.attrib_name]v{version_num}[/repr.attrib_name]: {timestamp}"
|
|
)
|
|
|
|
async def list_documents(self, filter: str | None = None):
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
documents = await self.client.list_documents(filter=filter)
|
|
for doc in documents:
|
|
self._rich_print_document(doc, truncate=True)
|
|
|
|
async def add_document_from_text(
|
|
self, text: str, title: str | None = None, metadata: dict | None = None
|
|
):
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
doc = await self.client.create_document(
|
|
text, title=title, metadata=metadata
|
|
)
|
|
self._rich_print_document(doc, truncate=True)
|
|
self.console.print(
|
|
f"[bold green]Document {doc.id} added successfully.[/bold green]"
|
|
)
|
|
|
|
async def add_document_from_source(
|
|
self, source: str, title: str | None = None, metadata: dict | None = None
|
|
):
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
result = await self.client.create_document_from_source(
|
|
source, title=title, metadata=metadata
|
|
)
|
|
if isinstance(result, list):
|
|
for doc in result:
|
|
self._rich_print_document(doc, truncate=True)
|
|
self.console.print(
|
|
f"[bold green]{len(result)} documents added successfully.[/bold green]"
|
|
)
|
|
else:
|
|
self._rich_print_document(result, truncate=True)
|
|
self.console.print(
|
|
f"[bold green]Document {result.id} added successfully.[/bold green]"
|
|
)
|
|
|
|
async def get_document(self, doc_id: str):
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
doc = await self.client.get_document_by_id(doc_id)
|
|
if doc is None:
|
|
self.console.print(f"[red]Document with id {doc_id} not found.[/red]")
|
|
return
|
|
self._rich_print_document(doc, truncate=False)
|
|
|
|
async def delete_document(self, doc_id: str):
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
deleted = await self.client.delete_document(doc_id)
|
|
if deleted:
|
|
self.console.print(
|
|
f"[bold green]Document {doc_id} deleted successfully.[/bold green]"
|
|
)
|
|
else:
|
|
self.console.print(
|
|
f"[yellow]Document with id {doc_id} not found.[/yellow]"
|
|
)
|
|
|
|
async def search(
|
|
self,
|
|
query: str | None = None,
|
|
limit: int | None = None,
|
|
filter: str | None = None,
|
|
image: Path | None = None,
|
|
):
|
|
if query is None and image is None:
|
|
self.console.print(
|
|
"[red]Provide either a query argument or --image PATH.[/red]"
|
|
)
|
|
return
|
|
if query is not None and image is not None:
|
|
self.console.print("[red]Pass either a query or --image, not both.[/red]")
|
|
return
|
|
|
|
search_input: str | bytes
|
|
if image is not None:
|
|
search_input = image.read_bytes()
|
|
else:
|
|
assert query is not None
|
|
search_input = query
|
|
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
results = await self.client.search(search_input, limit=limit, filter=filter)
|
|
if not results:
|
|
self.console.print("[yellow]No results found.[/yellow]")
|
|
return
|
|
for result in results:
|
|
self._rich_print_search_result(result)
|
|
|
|
async def visualize_chunk(self, chunk_id: str):
|
|
"""Display visual grounding images for a chunk."""
|
|
from textual_image.renderable import Image as RichImage
|
|
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
chunk = await self.client.get_chunk_by_id(chunk_id)
|
|
if not chunk:
|
|
self.console.print(f"[red]Chunk with id {chunk_id} not found.[/red]")
|
|
return
|
|
|
|
images = await self.client.visualize_chunk(chunk)
|
|
if not images:
|
|
self.console.print(
|
|
"[yellow]No visual grounding available for this chunk.[/yellow]"
|
|
)
|
|
self.console.print(
|
|
"This may be because the document was converted without page images."
|
|
)
|
|
return
|
|
|
|
self.console.print(f"[bold]Visual grounding for chunk {chunk_id}[/bold]")
|
|
if chunk.document_uri:
|
|
self.console.print(
|
|
f"[repr.attrib_name]document[/repr.attrib_name]: {chunk.document_uri}"
|
|
)
|
|
|
|
for i, img in enumerate(images):
|
|
self.console.print(
|
|
f"\n[bold cyan]Page {i + 1}/{len(images)}[/bold cyan]"
|
|
)
|
|
self.console.print(RichImage(img))
|
|
|
|
async def ask(
|
|
self,
|
|
question: str,
|
|
cite: bool = False,
|
|
filter: str | None = None,
|
|
):
|
|
"""Ask a question using the RAG system.
|
|
|
|
Args:
|
|
question: The question to ask
|
|
cite: Include citations in the answer
|
|
filter: SQL WHERE clause to filter documents
|
|
"""
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
answer, citations = await self.client.ask(question, filter=filter)
|
|
|
|
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
|
self.console.print()
|
|
self.console.print("[bold green]Answer:[/bold green]")
|
|
self.console.print(Markdown(answer))
|
|
if cite and citations:
|
|
for renderable in format_citations_rich(citations):
|
|
self.console.print(renderable)
|
|
|
|
async def analyze(
|
|
self,
|
|
question: str,
|
|
document: str | None = None,
|
|
filter: str | None = None,
|
|
):
|
|
"""Answer a question using the analysis agent with code execution.
|
|
|
|
Args:
|
|
question: The question to answer
|
|
document: Optional document ID or title to pre-load
|
|
filter: SQL WHERE clause to filter documents
|
|
"""
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as self.client:
|
|
documents = [document] if document else None
|
|
|
|
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
|
self.console.print()
|
|
self.console.print(
|
|
"[dim]Running analysis agent with code execution...[/dim]"
|
|
)
|
|
self.console.print()
|
|
|
|
result = await self.client.analyze(
|
|
question, documents=documents, filter=filter
|
|
)
|
|
|
|
self.console.print("[bold yellow]Program:[/bold yellow]")
|
|
self.console.print(Syntax(result.program, "python"))
|
|
self.console.print()
|
|
self.console.print("[bold green]Answer:[/bold green]")
|
|
self.console.print(Markdown(result.answer))
|
|
|
|
async def research(
|
|
self,
|
|
question: str,
|
|
filter: str | None = None,
|
|
):
|
|
"""Run research via the pydantic-graph pipeline.
|
|
|
|
Args:
|
|
question: The research question
|
|
filter: SQL WHERE clause to filter documents
|
|
"""
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as client:
|
|
self.console.print("[bold cyan]Starting research[/bold cyan]")
|
|
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
|
self.console.print()
|
|
|
|
report = await client.research(question=question, filter=filter)
|
|
|
|
if report is None:
|
|
self.console.print("[red]Research did not produce a report.[/red]")
|
|
return
|
|
|
|
# Display the report
|
|
self.console.print("[bold green]Research Report[/bold green]")
|
|
self.console.rule()
|
|
|
|
# Title and Executive Summary
|
|
self.console.print(f"[bold]{report.title}[/bold]")
|
|
self.console.print()
|
|
self.console.print("[bold cyan]Executive Summary:[/bold cyan]")
|
|
self.console.print(report.executive_summary)
|
|
self.console.print()
|
|
|
|
# Main Findings
|
|
if report.main_findings:
|
|
self.console.print("[bold cyan]Main Findings:[/bold cyan]")
|
|
for finding in report.main_findings:
|
|
self.console.print(f"• {finding}")
|
|
self.console.print()
|
|
|
|
# Conclusions
|
|
if report.conclusions:
|
|
self.console.print("[bold cyan]Conclusions:[/bold cyan]")
|
|
for conclusion in report.conclusions:
|
|
self.console.print(f"• {conclusion}")
|
|
self.console.print()
|
|
|
|
# Recommendations
|
|
if report.recommendations:
|
|
self.console.print("[bold cyan]Recommendations:[/bold cyan]")
|
|
for rec in report.recommendations:
|
|
self.console.print(f"• {rec}")
|
|
self.console.print()
|
|
|
|
# Limitations
|
|
if report.limitations:
|
|
self.console.print("[bold yellow]Limitations:[/bold yellow]")
|
|
for limitation in report.limitations:
|
|
self.console.print(f"• {limitation}")
|
|
self.console.print()
|
|
|
|
# Sources Summary
|
|
if report.sources_summary:
|
|
self.console.print("[bold cyan]Sources:[/bold cyan]")
|
|
self.console.print(report.sources_summary)
|
|
|
|
async def rebuild(self, mode: RebuildMode = RebuildMode.FULL):
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
skip_validation=True,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as client:
|
|
documents = await client.list_documents()
|
|
total_docs = len(documents)
|
|
|
|
if total_docs == 0:
|
|
self.console.print("[yellow]No documents found in database.[/yellow]")
|
|
return
|
|
|
|
mode_desc = {
|
|
RebuildMode.FULL: "full rebuild",
|
|
RebuildMode.RECHUNK: "rechunk",
|
|
RebuildMode.EMBED_ONLY: "embed only",
|
|
RebuildMode.TITLE_ONLY: "title only",
|
|
RebuildMode.DESCRIPTIONS: "picture descriptions",
|
|
}[mode]
|
|
|
|
self.console.print(
|
|
f"[bold cyan]Rebuilding database ({mode_desc}) with {total_docs} documents...[/bold cyan]"
|
|
)
|
|
with Progress() as progress:
|
|
task = progress.add_task("Rebuilding...", total=total_docs)
|
|
async for _ in client.rebuild_database(mode=mode):
|
|
progress.update(task, advance=1)
|
|
|
|
self.console.print(
|
|
"[bold green]Database rebuild completed successfully.[/bold green]"
|
|
)
|
|
|
|
async def vacuum(self):
|
|
"""Run database maintenance: optimize and cleanup table history."""
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
skip_validation=True,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as client:
|
|
await client.vacuum()
|
|
self.console.print("[bold green]Vacuum completed successfully.[/bold green]")
|
|
|
|
async def migrate(self) -> list[str]:
|
|
"""Run pending database migrations.
|
|
|
|
Returns:
|
|
List of descriptions of applied migrations.
|
|
"""
|
|
from haiku.rag.store.engine import Store
|
|
|
|
async with Store(
|
|
self.db_path,
|
|
config=self.config,
|
|
skip_validation=True,
|
|
skip_migration_check=True,
|
|
) as store:
|
|
return await store.migrate()
|
|
|
|
async def create_index(self):
|
|
"""Create vector index on the chunks table."""
|
|
async with HaikuRAG(
|
|
db_path=self.db_path,
|
|
config=self.config,
|
|
skip_validation=True,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as client:
|
|
row_count = await client.store.chunks_table.count_rows()
|
|
self.console.print(f"Chunks in database: {row_count}")
|
|
|
|
if row_count < 256:
|
|
self.console.print(
|
|
f"[yellow]Warning: Need at least 256 chunks to create an index (have {row_count})[/yellow]"
|
|
)
|
|
return
|
|
|
|
# Check if index already exists
|
|
indices = await client.store.chunks_table.list_indices()
|
|
has_vector_index = any("vector" in str(idx).lower() for idx in indices)
|
|
|
|
if has_vector_index:
|
|
self.console.print(
|
|
"[yellow]Rebuilding existing vector index...[/yellow]"
|
|
)
|
|
else:
|
|
self.console.print("[bold]Creating vector index...[/bold]")
|
|
|
|
await client.store._ensure_vector_index()
|
|
self.console.print(
|
|
"[bold green]Vector index created successfully.[/bold green]"
|
|
)
|
|
|
|
async def download_models(self):
|
|
"""Download Docling, HuggingFace tokenizer, and Ollama models per config."""
|
|
from haiku.rag.client.downloads import download_models
|
|
|
|
progress: Progress | None = None
|
|
task_id: TaskID | None = None
|
|
current_model = ""
|
|
current_digest = ""
|
|
|
|
async for event in download_models(self.config):
|
|
if event.status == "start":
|
|
self.console.print(
|
|
f"[bold blue]Downloading {event.model}...[/bold blue]"
|
|
)
|
|
elif event.status == "done":
|
|
if progress:
|
|
progress.stop()
|
|
progress = None
|
|
task_id = None
|
|
self.console.print(f"[green]✓[/green] {event.model}")
|
|
current_model = ""
|
|
current_digest = ""
|
|
elif event.status == "pulling":
|
|
self.console.print(f"[bold blue]Pulling {event.model}...[/bold blue]")
|
|
current_model = event.model
|
|
progress = Progress(
|
|
SpinnerColumn(),
|
|
TextColumn("[progress.description]{task.description}"),
|
|
BarColumn(),
|
|
DownloadColumn(),
|
|
TransferSpeedColumn(),
|
|
console=self.console,
|
|
transient=True,
|
|
auto_refresh=False,
|
|
)
|
|
progress.start()
|
|
task_id = progress.add_task(event.model, total=None)
|
|
elif event.status == "downloading" and progress and task_id is not None:
|
|
if event.digest != current_digest:
|
|
current_digest = event.digest
|
|
short_digest = event.digest[:19] if event.digest else ""
|
|
progress.update(
|
|
task_id,
|
|
description=f"{current_model} ({short_digest})",
|
|
total=event.total,
|
|
completed=0,
|
|
)
|
|
progress.update(task_id, completed=event.completed, refresh=True)
|
|
elif progress and task_id is not None:
|
|
progress.update(
|
|
task_id,
|
|
description=f"{current_model}: {event.status}",
|
|
refresh=True,
|
|
)
|
|
|
|
def show_settings(self):
|
|
"""Display current configuration settings."""
|
|
self.console.print("[bold]haiku.rag configuration[/bold]")
|
|
self.console.print()
|
|
|
|
# Get all config fields dynamically
|
|
for field_name, field_value in self.config.model_dump().items():
|
|
# Format the display value
|
|
if isinstance(field_value, str) and (
|
|
"key" in field_name.lower()
|
|
or "password" in field_name.lower()
|
|
or "token" in field_name.lower()
|
|
):
|
|
# Hide sensitive values but show if they're set
|
|
display_value = "✓ Set" if field_value else "✗ Not set"
|
|
else:
|
|
display_value = field_value
|
|
|
|
self.console.print(
|
|
f" [repr.attrib_name]{field_name}[/repr.attrib_name]: {display_value}"
|
|
)
|
|
|
|
def _rich_print_document(self, doc: Document, truncate: bool = False):
|
|
"""Format a document for display."""
|
|
if truncate:
|
|
content = doc.content.splitlines()
|
|
if len(content) > 3:
|
|
content = content[:3] + ["\n…"]
|
|
content = "\n".join(content)
|
|
content = Markdown(content)
|
|
else:
|
|
content = Markdown(doc.content)
|
|
title_part = (
|
|
f" [repr.attrib_name]title[/repr.attrib_name]: {doc.title}"
|
|
if doc.title
|
|
else ""
|
|
)
|
|
self.console.print(
|
|
f"[repr.attrib_name]id[/repr.attrib_name]: {doc.id} "
|
|
f"[repr.attrib_name]uri[/repr.attrib_name]: {doc.uri}"
|
|
+ title_part
|
|
+ f" [repr.attrib_name]meta[/repr.attrib_name]: {doc.metadata}"
|
|
)
|
|
self.console.print(
|
|
f"[repr.attrib_name]created at[/repr.attrib_name]: {doc.created_at} [repr.attrib_name]updated at[/repr.attrib_name]: {doc.updated_at}"
|
|
)
|
|
self.console.print("[repr.attrib_name]content[/repr.attrib_name]:")
|
|
self.console.print(content)
|
|
self.console.rule()
|
|
|
|
def _rich_print_search_result(self, result: "SearchResult"):
|
|
"""Format a search result for display."""
|
|
content = Markdown(result.content)
|
|
self.console.print(
|
|
f"[repr.attrib_name]document_id[/repr.attrib_name]: {result.document_id} "
|
|
f"[repr.attrib_name]chunk_id[/repr.attrib_name]: {result.chunk_id} "
|
|
f"[repr.attrib_name]score[/repr.attrib_name]: {result.score:.4f}"
|
|
)
|
|
if result.document_uri:
|
|
self.console.print(
|
|
f"[repr.attrib_name]document uri[/repr.attrib_name]: {result.document_uri}"
|
|
)
|
|
if result.document_title:
|
|
self.console.print("[repr.attrib_name]document title[/repr.attrib_name]:")
|
|
self.console.print(result.document_title)
|
|
if result.page_numbers:
|
|
self.console.print("[repr.attrib_name]pages[/repr.attrib_name]:")
|
|
self.console.print(", ".join(str(p) for p in result.page_numbers))
|
|
if result.headings:
|
|
self.console.print("[repr.attrib_name]headings[/repr.attrib_name]:")
|
|
self.console.print(" > ".join(result.headings))
|
|
self.console.print("[repr.attrib_name]content[/repr.attrib_name]:")
|
|
self.console.print(content)
|
|
self.console.rule()
|
|
|
|
async def serve(
|
|
self,
|
|
enable_monitor: bool = True,
|
|
enable_mcp: bool = True,
|
|
mcp_transport: str | None = None,
|
|
mcp_port: int = 8001,
|
|
):
|
|
"""Start the server with selected services."""
|
|
async with HaikuRAG(
|
|
self.db_path,
|
|
config=self.config,
|
|
read_only=self.read_only,
|
|
before=self.before,
|
|
) as client:
|
|
tasks = []
|
|
|
|
# Start file monitor if enabled (not available in read-only mode)
|
|
if enable_monitor:
|
|
if self.read_only:
|
|
logger.warning(
|
|
"File monitor disabled: cannot monitor files in read-only mode"
|
|
)
|
|
else:
|
|
monitor = FileWatcher(client=client, config=self.config)
|
|
monitor_task = asyncio.create_task(monitor.observe())
|
|
tasks.append(monitor_task)
|
|
|
|
if self.config.monitor.s3:
|
|
from haiku.rag.converters import get_converter
|
|
|
|
supported_extensions = get_converter(
|
|
self.config
|
|
).supported_extensions
|
|
for entry in self.config.monitor.s3:
|
|
s3_watcher = S3Watcher(
|
|
client=client,
|
|
entry=entry,
|
|
supported_extensions=supported_extensions,
|
|
)
|
|
tasks.append(asyncio.create_task(s3_watcher.observe()))
|
|
|
|
# Start MCP server if enabled
|
|
if enable_mcp:
|
|
server = create_mcp_server(
|
|
self.db_path, config=self.config, read_only=self.read_only
|
|
)
|
|
|
|
async def run_mcp():
|
|
if mcp_transport == "stdio":
|
|
await server.run_stdio_async()
|
|
else:
|
|
logger.info(f"Starting MCP server on port {mcp_port}")
|
|
await server.run_http_async(
|
|
transport="streamable-http", port=mcp_port
|
|
)
|
|
|
|
mcp_task = asyncio.create_task(run_mcp())
|
|
tasks.append(mcp_task)
|
|
|
|
if not tasks:
|
|
logger.warning("No services enabled")
|
|
return
|
|
|
|
try:
|
|
# Wait for any task to complete (or KeyboardInterrupt)
|
|
await asyncio.gather(*tasks)
|
|
except KeyboardInterrupt:
|
|
pass
|
|
finally:
|
|
# Cancel all tasks
|
|
for task in tasks:
|
|
task.cancel()
|
|
# Wait for cancellation
|
|
await asyncio.gather(*tasks, return_exceptions=True)
|