From 00178e603d1356374f2a68f073791d55bbeebad3 Mon Sep 17 00:00:00 2001 From: itsamejms Date: Mon, 7 Sep 2026 10:27:41 +0100 Subject: [PATCH] feat: first-run wizard pulls models with a real progress bar New pull_model command streams Ollama /api/pull NDJSON over a Channel (status lines + completed/total bytes). The wizard's model step gains a 'type any tag to pull' input: uninstalled models download in-app with a progress bar instead of the old 'run ollama pull yourself' dead end. The NDJSON line-loop is factored out of the chat streaming path and shared. --- src-tauri/src/commands/llm_commands.rs | 133 +++++++++++++++++++++---- src-tauri/src/lib.rs | 1 + src/components/FirstRunWizard.tsx | 120 ++++++++++++++++++---- 3 files changed, 217 insertions(+), 37 deletions(-) diff --git a/src-tauri/src/commands/llm_commands.rs b/src-tauri/src/commands/llm_commands.rs index 18af168..8e55197 100644 --- a/src-tauri/src/commands/llm_commands.rs +++ b/src-tauri/src/commands/llm_commands.rs @@ -1,5 +1,6 @@ use crate::commands::emit_busy; use crate::llm::{self, AppState, ChatMessage, ChatResponse, GenerateRequest, LlmEvent, OllamaChatResponse}; +use serde::{Deserialize, Serialize}; use serde_json::json; use tauri::ipc::Channel; use tauri::AppHandle; @@ -86,6 +87,91 @@ pub async fn generate_stream( Ok(()) } +// ─── Model pull (Ollama) ───────────────────────────── + +/// Channel events for a streaming `ollama pull`. +#[derive(Clone, Serialize)] +#[serde(tag = "type", content = "data", rename_all = "camelCase")] +pub enum PullEvent { + /// Human-readable status line ("pulling manifest", "verifying sha256 digest"). + Status(String), + /// Download progress in bytes. + Progress { completed: Option, total: Option }, + Done, + Error(String), +} + +#[derive(Debug, Deserialize)] +pub struct PullRequest { + pub model: String, +} + +/// POST {api_url}/api/pull with stream:true, progress over the channel. +/// Ollama-only — other providers install models out of band, so the wizard +/// gates this. Progress fields vary per Ollama version: completed/total are +/// Option and the UI falls back to the status line. +#[tauri::command] +pub async fn pull_model( + state: tauri::State<'_, AppState>, + app: AppHandle, + req: PullRequest, + channel: Channel, +) -> Result<(), String> { + let config = state.config.lock().map_err(|e| e.to_string())?.clone(); + if !llm::is_ollama(&config.provider, &config.api_url) { + return Err("Model pull is only supported for Ollama endpoints".into()); + } + emit_busy(&app, true); + let url = format!("{}/api/pull", config.api_url.trim_end_matches('/')); + let body = json!({ "model": req.model, "stream": true }); + tauri::async_runtime::spawn(async move { + let client = reqwest::Client::new(); + let res = client.post(&url).json(&body).send().await; + let result = match res { + Ok(r) if r.status().is_success() => { + for_each_ndjson_line(r, |line| { + let line = line.trim(); + if line.is_empty() { + return Ok(false); + } + let Ok(v) = serde_json::from_str::(line) else { + return Ok(false); // skip unparseable progress lines + }; + if let Some(err) = v["error"].as_str() { + return Err(err.to_string()); + } + let status = v["status"].as_str().unwrap_or_default().to_string(); + let completed = v["completed"].as_u64(); + let total = v["total"].as_u64(); + if status == "success" { + return Ok(true); + } + if completed.is_some() || total.is_some() { + let _ = channel.send(PullEvent::Progress { completed, total }); + } else if !status.is_empty() { + let _ = channel.send(PullEvent::Status(status)); + } + Ok(false) + }) + .await + } + Ok(r) => Err(format!("pull error {}", r.status())), + Err(e) => Err(format!("pull failed: {e}")), + }; + match result { + Ok(()) => { + let _ = channel.send(PullEvent::Done); + } + Err(e) => { + let _ = channel.send(PullEvent::Error(e)); + } + } + // ponytail: always balance the busy counter, even on error. + emit_busy(&app, false); + }); + Ok(()) +} + // ─── Get / Set LLM Config ──────────────────────────────────── #[tauri::command] @@ -379,14 +465,12 @@ async fn stream_chat( stream_from_response(res, channel, parse).await } -/// Consume a streaming HTTP response: buffer bytes into lines, parse each -/// with `parse`, emit Token per piece, Done via the caller's Ok(full text). -async fn stream_from_response( +/// Consume a streaming HTTP response: buffer bytes into lines, hand each +/// complete line to `on_line`. `Ok(true)` from the callback stops the loop. +async fn for_each_ndjson_line( res: reqwest::Response, - channel: &Channel, - parse: fn(&str) -> Result, String>, -) -> Result { - let mut full = String::new(); + mut on_line: impl FnMut(&str) -> Result, +) -> Result<(), String> { let mut buf = String::new(); let mut stream = res.bytes_stream(); while let Some(chunk) = stream.next().await { @@ -394,19 +478,34 @@ async fn stream_from_response( buf.push_str(&String::from_utf8_lossy(&chunk)); while let Some(nl) = buf.find('\n') { let line: String = buf.drain(..=nl).collect(); - match parse(line.trim_end())? { - Some(StreamLine { token: Some(t), done }) => { - full.push_str(&t); - let _ = channel.send(LlmEvent::Token(t)); - if done { - return Ok(full); - } - } - Some(StreamLine { token: None, done: true }) => return Ok(full), - _ => {} + if on_line(line.trim_end())? { + return Ok(()); } } } + Ok(()) +} + +/// Consume a streaming chat response: parse each line with `parse`, emit +/// Token per piece, and return the full text when the stream finishes. +async fn stream_from_response( + res: reqwest::Response, + channel: &Channel, + parse: fn(&str) -> Result, String>, +) -> Result { + let mut full = String::new(); + for_each_ndjson_line(res, |line| { + match parse(line)? { + Some(StreamLine { token: Some(t), done }) => { + full.push_str(&t); + let _ = channel.send(LlmEvent::Token(t)); + Ok(done) + } + Some(StreamLine { token: None, done: true }) => Ok(true), + _ => Ok(false), + } + }) + .await?; // Stream closed without an explicit done — return what we got. Ok(full) } diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 536ddd6..773ae88 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -73,6 +73,7 @@ pub fn run() { commands::llm_commands::get_llm_config, commands::llm_commands::set_llm_config, commands::llm_commands::test_connection, + commands::llm_commands::pull_model, commands::image_commands::generate_image, commands::image_commands::generate_image_stream, commands::image_commands::test_image_connection, diff --git a/src/components/FirstRunWizard.tsx b/src/components/FirstRunWizard.tsx index ad54836..852834e 100644 --- a/src/components/FirstRunWizard.tsx +++ b/src/components/FirstRunWizard.tsx @@ -1,13 +1,26 @@ import { useEffect, useState } from "react"; -import { invoke } from "@tauri-apps/api/core"; +import { invoke, Channel } from "@tauri-apps/api/core"; import { useToast } from "./Toast"; +// ponytail: mirrors the Rust PullEvent enum (adjacently tagged, camelCase). +type PullEvent = + | { type: "progress"; data: { completed: number | null; total: number | null } } + | { type: "status"; data: string } + | { type: "done" } + | { type: "error"; data: string }; + export function FirstRunWizard({ onComplete }: { onComplete: () => void }) { const [step, setStep] = useState("welcome"); const [loading, setLoading] = useState(false); const [error, setError] = useState(""); const [models, setModels] = useState([]); const [selectedModel, setSelectedModel] = useState(""); + // ponytail: custom tag not in the installed list — pulled in-wizard with a + // progress bar instead of the old "run ollama pull yourself" dead end. + const [customModel, setCustomModel] = useState(""); + const [pulling, setPulling] = useState(false); + const [pullPct, setPullPct] = useState(null); + const [pullStatus, setPullStatus] = useState(""); // ponytail: wizard defaults to local Ollama; no API-key UI here (Settings covers remote). const apiUrl = "http://localhost:11434"; const apiKey = ""; @@ -42,20 +55,18 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) { setLoading(false); } - async function pullModel(model: string) { + // ponytail: save the chosen model and finish — no pull needed when it's + // already installed. + async function saveModelAndFinish(model: string) { setLoading(true); setError(""); try { - // We don't have a direct pull command, but we can suggest the user to run `ollama pull` in terminal. - // For now, we'll just set the model in config and hope it exists. - // Alternatively, we could invoke a generate command to trigger a pull? Not sure. - // We'll just set the config and complete. await invoke("set_llm_config", { config: { provider: "ollama", api_url: apiUrl, api_key: apiKey, - model: model, + model, temperature: 0.7, max_tokens: 512, top_p: 0.9, @@ -63,7 +74,6 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) { embed_model: "nomic-embed-text", } }); - addToast(`Model ${model} selected. You may need to run 'ollama pull ${model}' if not already downloaded.`, "success"); setStep("done"); } catch (e) { setError(String(e)); @@ -72,6 +82,42 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) { setLoading(false); } + // Real `ollama pull` over the Channel — progress bar driven by the + // completed/total bytes Ollama streams; falls back to the status line. + async function pullAndFinish(model: string) { + setPulling(true); + setError(""); + setPullPct(null); + setPullStatus("starting pull…"); + const channel = new Channel(); + let failed = false; + channel.onmessage = (ev) => { + if (ev.type === "progress") { + setPullStatus(""); + setPullPct(ev.data.completed != null && ev.data.total ? ev.data.completed / ev.data.total : null); + } else if (ev.type === "status") { + setPullStatus(ev.data); + setPullPct(null); + } else if (ev.type === "done") { + addToast(`Model ${model} pulled`, "success"); + saveModelAndFinish(model); + setPulling(false); + } else if (ev.type === "error") { + failed = true; + setError(`Pull failed: ${ev.data}`); + setPulling(false); + } + }; + try { + await invoke("pull_model", { req: { model }, channel }); + } catch (e) { + if (!failed) { + setError(`Pull failed: ${e}`); + setPulling(false); + } + } + } + function skip() { // Skip wizard and go to settings onComplete(); @@ -93,11 +139,17 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) { return; } if (step === "model") { - if (!selectedModel) { - setError("Please select a model"); + const model = customModel.trim() || selectedModel; + if (!model) { + setError("Choose a model or type a tag to pull"); return; } - await pullModel(selectedModel); + // Installed → save directly; anything else gets pulled first. + if (models.includes(model)) { + await saveModelAndFinish(model); + } else { + await pullAndFinish(model); + } return; } if (step === "done") { @@ -193,7 +245,7 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) { -
+
+ + { setCustomModel(e.target.value); setSelectedModel(""); }} + placeholder="e.g. llama3.2" + className="rounded-lg bg-[var(--color-bg-surface)] border border-[var(--color-border-glass)] px-3 py-2 text-[var(--color-text-primary)] placeholder-[var(--color-text-dim)] focus:outline-none focus:border-[var(--color-gold-bright)] text-xs" + /> +
+
+ {pulling ? ( +
+

Pulling {customModel.trim() || selectedModel} …

+
+
+
+

+ {pullPct != null ? `${Math.round(pullPct * 100)}%` : pullStatus} +

+
+ ) : ( +
-
+ )} {error && (

{error}

)} -

- If you don't see your model, you may need to download it first. In a terminal, run:{' '} - ollama pull llama3.2 -

+ {!pulling && ( +

+ Models not in the list are pulled automatically — type any Ollama tag + (e.g. llama3.2, or + nomic-embed-text for + lore search) and DM-Pal downloads it with a progress bar. +

+ )} ); }