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.
This commit is contained in:
@@ -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<u64>, total: Option<u64> },
|
||||
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<PullEvent>,
|
||||
) -> 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::<serde_json::Value>(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<LlmEvent>,
|
||||
parse: fn(&str) -> Result<Option<StreamLine>, String>,
|
||||
) -> Result<String, String> {
|
||||
let mut full = String::new();
|
||||
mut on_line: impl FnMut(&str) -> Result<bool, String>,
|
||||
) -> 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<LlmEvent>,
|
||||
parse: fn(&str) -> Result<Option<StreamLine>, String>,
|
||||
) -> Result<String, String> {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user