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::commands::emit_busy;
|
||||||
use crate::llm::{self, AppState, ChatMessage, ChatResponse, GenerateRequest, LlmEvent, OllamaChatResponse};
|
use crate::llm::{self, AppState, ChatMessage, ChatResponse, GenerateRequest, LlmEvent, OllamaChatResponse};
|
||||||
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
use tauri::ipc::Channel;
|
use tauri::ipc::Channel;
|
||||||
use tauri::AppHandle;
|
use tauri::AppHandle;
|
||||||
@@ -86,6 +87,91 @@ pub async fn generate_stream(
|
|||||||
Ok(())
|
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 ────────────────────────────────────
|
// ─── Get / Set LLM Config ────────────────────────────────────
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
@@ -379,14 +465,12 @@ async fn stream_chat(
|
|||||||
stream_from_response(res, channel, parse).await
|
stream_from_response(res, channel, parse).await
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Consume a streaming HTTP response: buffer bytes into lines, parse each
|
/// Consume a streaming HTTP response: buffer bytes into lines, hand each
|
||||||
/// with `parse`, emit Token per piece, Done via the caller's Ok(full text).
|
/// complete line to `on_line`. `Ok(true)` from the callback stops the loop.
|
||||||
async fn stream_from_response(
|
async fn for_each_ndjson_line(
|
||||||
res: reqwest::Response,
|
res: reqwest::Response,
|
||||||
channel: &Channel<LlmEvent>,
|
mut on_line: impl FnMut(&str) -> Result<bool, String>,
|
||||||
parse: fn(&str) -> Result<Option<StreamLine>, String>,
|
) -> Result<(), String> {
|
||||||
) -> Result<String, String> {
|
|
||||||
let mut full = String::new();
|
|
||||||
let mut buf = String::new();
|
let mut buf = String::new();
|
||||||
let mut stream = res.bytes_stream();
|
let mut stream = res.bytes_stream();
|
||||||
while let Some(chunk) = stream.next().await {
|
while let Some(chunk) = stream.next().await {
|
||||||
@@ -394,19 +478,34 @@ async fn stream_from_response(
|
|||||||
buf.push_str(&String::from_utf8_lossy(&chunk));
|
buf.push_str(&String::from_utf8_lossy(&chunk));
|
||||||
while let Some(nl) = buf.find('\n') {
|
while let Some(nl) = buf.find('\n') {
|
||||||
let line: String = buf.drain(..=nl).collect();
|
let line: String = buf.drain(..=nl).collect();
|
||||||
match parse(line.trim_end())? {
|
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 }) => {
|
Some(StreamLine { token: Some(t), done }) => {
|
||||||
full.push_str(&t);
|
full.push_str(&t);
|
||||||
let _ = channel.send(LlmEvent::Token(t));
|
let _ = channel.send(LlmEvent::Token(t));
|
||||||
if done {
|
Ok(done)
|
||||||
return Ok(full);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Some(StreamLine { token: None, done: true }) => return Ok(full),
|
|
||||||
_ => {}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
Some(StreamLine { token: None, done: true }) => Ok(true),
|
||||||
|
_ => Ok(false),
|
||||||
}
|
}
|
||||||
|
})
|
||||||
|
.await?;
|
||||||
// Stream closed without an explicit done — return what we got.
|
// Stream closed without an explicit done — return what we got.
|
||||||
Ok(full)
|
Ok(full)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -73,6 +73,7 @@ pub fn run() {
|
|||||||
commands::llm_commands::get_llm_config,
|
commands::llm_commands::get_llm_config,
|
||||||
commands::llm_commands::set_llm_config,
|
commands::llm_commands::set_llm_config,
|
||||||
commands::llm_commands::test_connection,
|
commands::llm_commands::test_connection,
|
||||||
|
commands::llm_commands::pull_model,
|
||||||
commands::image_commands::generate_image,
|
commands::image_commands::generate_image,
|
||||||
commands::image_commands::generate_image_stream,
|
commands::image_commands::generate_image_stream,
|
||||||
commands::image_commands::test_image_connection,
|
commands::image_commands::test_image_connection,
|
||||||
|
|||||||
@@ -1,13 +1,26 @@
|
|||||||
import { useEffect, useState } from "react";
|
import { useEffect, useState } from "react";
|
||||||
import { invoke } from "@tauri-apps/api/core";
|
import { invoke, Channel } from "@tauri-apps/api/core";
|
||||||
import { useToast } from "./Toast";
|
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 }) {
|
export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
||||||
const [step, setStep] = useState("welcome");
|
const [step, setStep] = useState("welcome");
|
||||||
const [loading, setLoading] = useState(false);
|
const [loading, setLoading] = useState(false);
|
||||||
const [error, setError] = useState("");
|
const [error, setError] = useState("");
|
||||||
const [models, setModels] = useState<string[]>([]);
|
const [models, setModels] = useState<string[]>([]);
|
||||||
const [selectedModel, setSelectedModel] = useState<string>("");
|
const [selectedModel, setSelectedModel] = useState<string>("");
|
||||||
|
// 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<number | null>(null);
|
||||||
|
const [pullStatus, setPullStatus] = useState("");
|
||||||
// ponytail: wizard defaults to local Ollama; no API-key UI here (Settings covers remote).
|
// ponytail: wizard defaults to local Ollama; no API-key UI here (Settings covers remote).
|
||||||
const apiUrl = "http://localhost:11434";
|
const apiUrl = "http://localhost:11434";
|
||||||
const apiKey = "";
|
const apiKey = "";
|
||||||
@@ -42,20 +55,18 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
|||||||
setLoading(false);
|
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);
|
setLoading(true);
|
||||||
setError("");
|
setError("");
|
||||||
try {
|
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", {
|
await invoke("set_llm_config", {
|
||||||
config: {
|
config: {
|
||||||
provider: "ollama",
|
provider: "ollama",
|
||||||
api_url: apiUrl,
|
api_url: apiUrl,
|
||||||
api_key: apiKey,
|
api_key: apiKey,
|
||||||
model: model,
|
model,
|
||||||
temperature: 0.7,
|
temperature: 0.7,
|
||||||
max_tokens: 512,
|
max_tokens: 512,
|
||||||
top_p: 0.9,
|
top_p: 0.9,
|
||||||
@@ -63,7 +74,6 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
|||||||
embed_model: "nomic-embed-text",
|
embed_model: "nomic-embed-text",
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
addToast(`Model ${model} selected. You may need to run 'ollama pull ${model}' if not already downloaded.`, "success");
|
|
||||||
setStep("done");
|
setStep("done");
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
setError(String(e));
|
setError(String(e));
|
||||||
@@ -72,6 +82,42 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
|||||||
setLoading(false);
|
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<PullEvent>();
|
||||||
|
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() {
|
function skip() {
|
||||||
// Skip wizard and go to settings
|
// Skip wizard and go to settings
|
||||||
onComplete();
|
onComplete();
|
||||||
@@ -93,11 +139,17 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
if (step === "model") {
|
if (step === "model") {
|
||||||
if (!selectedModel) {
|
const model = customModel.trim() || selectedModel;
|
||||||
setError("Please select a model");
|
if (!model) {
|
||||||
|
setError("Choose a model or type a tag to pull");
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
await pullModel(selectedModel);
|
// Installed → save directly; anything else gets pulled first.
|
||||||
|
if (models.includes(model)) {
|
||||||
|
await saveModelAndFinish(model);
|
||||||
|
} else {
|
||||||
|
await pullAndFinish(model);
|
||||||
|
}
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
if (step === "done") {
|
if (step === "done") {
|
||||||
@@ -193,7 +245,7 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
|||||||
<label className="text-[var(--color-text-secondary)] text-xs font-medium">Model</label>
|
<label className="text-[var(--color-text-secondary)] text-xs font-medium">Model</label>
|
||||||
<select
|
<select
|
||||||
value={selectedModel}
|
value={selectedModel}
|
||||||
onChange={(e) => setSelectedModel(e.target.value)}
|
onChange={(e) => { setSelectedModel(e.target.value); setCustomModel(""); }}
|
||||||
className="rounded-lg bg-[var(--color-bg-surface)] border border-[var(--color-border-glass)] px-3 py-2 text-[var(--color-text-primary)] focus:outline-none focus:border-[var(--color-gold-bright)] text-xs cursor-pointer"
|
className="rounded-lg bg-[var(--color-bg-surface)] border border-[var(--color-border-glass)] px-3 py-2 text-[var(--color-text-primary)] focus:outline-none focus:border-[var(--color-gold-bright)] text-xs cursor-pointer"
|
||||||
>
|
>
|
||||||
{models.map((m) => (
|
{models.map((m) => (
|
||||||
@@ -204,13 +256,37 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
|||||||
))}
|
))}
|
||||||
</select>
|
</select>
|
||||||
</div>
|
</div>
|
||||||
<div className="flex flex-col gap-3">
|
<div className="flex flex-col gap-1">
|
||||||
|
<label className="text-[var(--color-text-secondary)] text-xs font-medium">Or type any Ollama tag to pull</label>
|
||||||
|
<input
|
||||||
|
value={customModel}
|
||||||
|
onChange={(e) => { 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"
|
||||||
|
/>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
{pulling ? (
|
||||||
|
<div className="flex flex-col gap-2 mt-4">
|
||||||
|
<p className="text-[var(--color-text-secondary)] text-xs">Pulling {customModel.trim() || selectedModel} …</p>
|
||||||
|
<div className="w-full h-2 rounded-full bg-[var(--color-bg-surface)] overflow-hidden">
|
||||||
|
<div
|
||||||
|
className="h-full bg-[var(--color-gold-bright)] transition-all duration-300"
|
||||||
|
style={{ width: pullPct != null ? `${Math.round(pullPct * 100)}%` : "100%" }}
|
||||||
|
/>
|
||||||
|
</div>
|
||||||
|
<p className="text-[var(--color-text-dim)] text-xs">
|
||||||
|
{pullPct != null ? `${Math.round(pullPct * 100)}%` : pullStatus}
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
) : (
|
||||||
|
<div className="flex flex-col gap-3 mt-4">
|
||||||
<button
|
<button
|
||||||
onClick={handleSubmit}
|
onClick={handleSubmit}
|
||||||
disabled={loading || !selectedModel}
|
disabled={loading || (!customModel.trim() && !selectedModel)}
|
||||||
className="rounded-lg bg-[var(--color-gold-bright)] text-[var(--color-bg-deep)] px-4 py-2 text-sm font-semibold hover:bg-[var(--color-gold-muted)] transition-colors cursor-pointer disabled:opacity-50"
|
className="rounded-lg bg-[var(--color-gold-bright)] text-[var(--color-bg-deep)] px-4 py-2 text-sm font-semibold hover:bg-[var(--color-gold-muted)] transition-colors cursor-pointer disabled:opacity-50"
|
||||||
>
|
>
|
||||||
{loading ? "Setting model…" : "Set Model"}
|
{loading ? "Setting model…" : "Use Model"}
|
||||||
</button>
|
</button>
|
||||||
<button
|
<button
|
||||||
onClick={back}
|
onClick={back}
|
||||||
@@ -225,14 +301,18 @@ export function FirstRunWizard({ onComplete }: { onComplete: () => void }) {
|
|||||||
Skip for now (go to Settings)
|
Skip for now (go to Settings)
|
||||||
</button>
|
</button>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
)}
|
||||||
{error && (
|
{error && (
|
||||||
<p className="text-[var(--color-danger)] text-sm mt-4">{error}</p>
|
<p className="text-[var(--color-danger)] text-sm mt-4">{error}</p>
|
||||||
)}
|
)}
|
||||||
<p className="text-[var(--color-text-dim)] text-sm mt-4">
|
{!pulling && (
|
||||||
If you don't see your model, you may need to download it first. In a terminal, run:{' '}
|
<p className="text-[var(--color-text-dim)] text-xs mt-4">
|
||||||
<code className="font-mono text-[var(--color-gold-bright)]">ollama pull llama3.2</code>
|
Models not in the list are pulled automatically — type any Ollama tag
|
||||||
|
(e.g. <code className="font-mono text-[var(--color-gold-bright)]">llama3.2</code>, or
|
||||||
|
<code className="font-mono text-[var(--color-gold-bright)]">nomic-embed-text</code> for
|
||||||
|
lore search) and DM-Pal downloads it with a progress bar.
|
||||||
</p>
|
</p>
|
||||||
|
)}
|
||||||
</div>
|
</div>
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user