From 80447c9024fdf77550f11943e81100bb53056041 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 4 Dec 2025 16:30:26 +0200 Subject: [PATCH] Make download_models() a HaikuRAG client method. Update app to show progress of what happens. --- CHANGELOG.md | 5 ++ docs/cli.md | 10 +++- haiku_rag_slim/haiku/rag/app.py | 67 +++++++++++++++++++++- haiku_rag_slim/haiku/rag/cli.py | 6 +- haiku_rag_slim/haiku/rag/client.py | 89 ++++++++++++++++++++++++++++++ haiku_rag_slim/haiku/rag/utils.py | 48 ---------------- 6 files changed, 169 insertions(+), 56 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 999ade56..947e6c51 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,11 @@ # Changelog ## [Unreleased] +### Changed + +- **Download Models Progress**: `haiku-rag download-models` now shows real-time progress with Rich progress bars for Ollama model downloads +- **Refactored Download Models**: Moved core download logic to `HaikuRAG.download_models()` async generator that yields `DownloadProgress` events, separating business logic from UI + ## [0.19.6] - 2025-12-03 ## [0.19.6] - 2025-12-03 diff --git a/docs/cli.md b/docs/cli.md index 6373aa13..da3f6493 100644 --- a/docs/cli.md +++ b/docs/cli.md @@ -298,6 +298,10 @@ Download required runtime models: haiku-rag download-models ``` -This command: -- Downloads Docling OCR/conversion models (no-op if already present). -- Pulls Ollama models referenced in your configuration (embeddings, QA, research, rerank). +This command downloads: + +- Docling OCR/conversion models +- HuggingFace tokenizer (for chunking) +- Ollama models referenced in your configuration (embeddings, QA, research, rerank) + +Progress is displayed in real-time with download status and progress bars for Ollama model pulls. diff --git a/haiku_rag_slim/haiku/rag/app.py b/haiku_rag_slim/haiku/rag/app.py index 0ca69fbc..d40622a2 100644 --- a/haiku_rag_slim/haiku/rag/app.py +++ b/haiku_rag_slim/haiku/rag/app.py @@ -6,7 +6,15 @@ from pathlib import Path from rich.console import Console from rich.markdown import Markdown -from rich.progress import Progress +from rich.progress import ( + BarColumn, + DownloadColumn, + Progress, + SpinnerColumn, + TaskID, + TextColumn, + TransferSpeedColumn, +) from haiku.rag.client import HaikuRAG, RebuildMode from haiku.rag.config import AppConfig, Config @@ -491,6 +499,63 @@ class HaikuRAGApp: except Exception as e: self.console.print(f"[red]Error creating index: {e}[/red]") + async def download_models(self): + """Download Docling, HuggingFace tokenizer, and Ollama models per config.""" + from haiku.rag.client import HaikuRAG + + client = HaikuRAG(db_path=None, config=self.config) + + progress: Progress | None = None + task_id: TaskID | None = None + current_model = "" + current_digest = "" + + async for event in client.download_models(): + 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]") diff --git a/haiku_rag_slim/haiku/rag/cli.py b/haiku_rag_slim/haiku/rag/cli.py index 49ca5f7b..6c165ae9 100644 --- a/haiku_rag_slim/haiku/rag/cli.py +++ b/haiku_rag_slim/haiku/rag/cli.py @@ -436,11 +436,9 @@ def info( @cli.command("download-models", help="Download Docling and Ollama models per config") def download_models_cmd(): - from haiku.rag.utils import prefetch_models - + app = HaikuRAGApp(db_path=Path(), config=get_config()) try: - asyncio.run(prefetch_models()) - typer.echo("Models downloaded successfully.") + asyncio.run(app.download_models()) except Exception as e: typer.echo(f"Error downloading models: {e}") raise typer.Exit(1) diff --git a/haiku_rag_slim/haiku/rag/client.py b/haiku_rag_slim/haiku/rag/client.py index 79e05907..d3c15513 100644 --- a/haiku_rag_slim/haiku/rag/client.py +++ b/haiku_rag_slim/haiku/rag/client.py @@ -1,8 +1,11 @@ +import asyncio import hashlib +import json import logging import mimetypes import tempfile from collections.abc import AsyncGenerator +from dataclasses import dataclass from enum import Enum from pathlib import Path from urllib.parse import urlparse @@ -30,6 +33,17 @@ class RebuildMode(Enum): EMBED_ONLY = "embed_only" # Keep chunks, only regenerate embeddings +@dataclass +class DownloadProgress: + """Progress event for model downloads.""" + + model: str + status: str + completed: int = 0 + total: int = 0 + digest: str = "" + + class HaikuRAG: """High-level haiku-rag client.""" @@ -856,6 +870,81 @@ class HaikuRAG: """Optimize and clean up old versions across all tables.""" await self.store.vacuum() + async def download_models(self) -> AsyncGenerator[DownloadProgress, None]: + """Download required models, yielding progress events. + + Yields DownloadProgress events for: + - Docling models (status="docling_start", "docling_done") + - HuggingFace tokenizer (status="tokenizer_start", "tokenizer_done") + - Ollama models (status="pulling", "downloading", "done", or other Ollama statuses) + """ + # Docling models + try: + from docling.utils.model_downloader import download_models + + yield DownloadProgress(model="docling", status="start") + await asyncio.to_thread(download_models) + yield DownloadProgress(model="docling", status="done") + except ImportError: + pass + + # HuggingFace tokenizer + from transformers import AutoTokenizer + + tokenizer_name = self._config.processing.chunking_tokenizer + yield DownloadProgress(model=tokenizer_name, status="start") + await asyncio.to_thread(AutoTokenizer.from_pretrained, tokenizer_name) + yield DownloadProgress(model=tokenizer_name, status="done") + + # Collect Ollama models from config + required_models: set[str] = set() + if self._config.embeddings.model.provider == "ollama": + required_models.add(self._config.embeddings.model.name) + if self._config.qa.model.provider == "ollama": + required_models.add(self._config.qa.model.name) + if self._config.research.model.provider == "ollama": + required_models.add(self._config.research.model.name) + if ( + self._config.reranking.model + and self._config.reranking.model.provider == "ollama" + ): + required_models.add(self._config.reranking.model.name) + + if not required_models: + return + + base_url = self._config.providers.ollama.base_url + + async with httpx.AsyncClient(timeout=None) as client: + for model in sorted(required_models): + yield DownloadProgress(model=model, status="pulling") + + async with client.stream( + "POST", f"{base_url}/api/pull", json={"model": model} + ) as r: + async for line in r.aiter_lines(): + if not line: + continue + try: + data = json.loads(line) + status = data.get("status", "") + digest = data.get("digest", "") + + if digest and "total" in data: + yield DownloadProgress( + model=model, + status="downloading", + total=data.get("total", 0), + completed=data.get("completed", 0), + digest=digest, + ) + elif status: + yield DownloadProgress(model=model, status=status) + except json.JSONDecodeError: + pass + + yield DownloadProgress(model=model, status="done") + def close(self): """Close the underlying store connection.""" self.store.close() diff --git a/haiku_rag_slim/haiku/rag/utils.py b/haiku_rag_slim/haiku/rag/utils.py index b2c03e11..2eec29b6 100644 --- a/haiku_rag_slim/haiku/rag/utils.py +++ b/haiku_rag_slim/haiku/rag/utils.py @@ -365,51 +365,3 @@ def load_callable(path: str): f"Attribute '{func_name}' in module '{module_part}' is not callable" ) return func - - -async def prefetch_models(): - """Prefetch runtime models (Docling + Ollama + HuggingFace tokenizer as configured).""" - import asyncio - - import httpx - - from haiku.rag.config import Config - - try: - from docling.utils.model_downloader import download_models - - await asyncio.to_thread(download_models) - except ImportError: - # Docling not installed, skip downloading docling models - pass - - # Download HuggingFace tokenizer - from transformers import AutoTokenizer - - await asyncio.to_thread( - AutoTokenizer.from_pretrained, Config.processing.chunking_tokenizer - ) - - # Collect Ollama models from config - required_models: set[str] = set() - if Config.embeddings.model.provider == "ollama": - required_models.add(Config.embeddings.model.name) - if Config.qa.model.provider == "ollama": - required_models.add(Config.qa.model.name) - if Config.research.model.provider == "ollama": - required_models.add(Config.research.model.name) - if Config.reranking.model and Config.reranking.model.provider == "ollama": - required_models.add(Config.reranking.model.name) - - if not required_models: - return - - base_url = Config.providers.ollama.base_url - - async with httpx.AsyncClient(timeout=None) as client: - for model in sorted(required_models): - async with client.stream( - "POST", f"{base_url}/api/pull", json={"model": model} - ) as r: - async for _ in r.aiter_lines(): - pass