updating to use stable-diffusion.cpp for the image generation as ollama has since removed it
This commit is contained in:
@@ -1,66 +1,32 @@
|
||||
use crate::commands::emit_busy;
|
||||
use crate::llm::AppState;
|
||||
use futures_util::StreamExt;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::hash_map::DefaultHasher;
|
||||
use std::hash::{Hash, Hasher};
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tauri::ipc::Channel;
|
||||
use tauri::AppHandle;
|
||||
|
||||
// ponytail: we talk to a stable-diffusion.cpp `sd-server` over HTTP using its
|
||||
// AUTOMATIC1111-compatible API (/sdapi/v1/txt2img + /sdapi/v1/progress). This
|
||||
// reuses the existing HTTP+channel plumbing and drops the old Ollama/macOS-only
|
||||
// limitation — sd-server runs on macOS, Linux, and Windows. The model is chosen
|
||||
// at server startup, so there's no per-request model field anymore.
|
||||
|
||||
// ponytail: DefaultHasher is fine for a cache filename — not crypto, just a stable key.
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ImageRequest {
|
||||
pub prompt: String,
|
||||
/// Override the configured image model for this call.
|
||||
pub model: Option<String>,
|
||||
}
|
||||
|
||||
/// One line of Ollama's NDJSON image-generation response.
|
||||
/// `step`/`total` carry progress; the final `done: true` line carries the
|
||||
/// singular `image` base64 field. All fields optional per-line.
|
||||
#[derive(Debug, Default, Deserialize)]
|
||||
struct OllamaImageLine {
|
||||
#[serde(default)]
|
||||
done: bool,
|
||||
#[serde(default)]
|
||||
step: Option<u32>,
|
||||
#[serde(default)]
|
||||
total: Option<u32>,
|
||||
#[serde(default)]
|
||||
image: Option<String>,
|
||||
}
|
||||
|
||||
/// Parsed view of one NDJSON line used by both the streaming and the
|
||||
/// buffered paths. Extracted so the parsing logic is unit-testable without
|
||||
/// touching the network.
|
||||
#[derive(Debug, PartialEq)]
|
||||
struct ImageProgress {
|
||||
done: bool,
|
||||
step: Option<u32>,
|
||||
total: Option<u32>,
|
||||
image: Option<String>,
|
||||
}
|
||||
|
||||
fn parse_image_line(line: &str) -> Option<ImageProgress> {
|
||||
let line = line.trim();
|
||||
if line.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let parsed: OllamaImageLine = serde_json::from_str(line).ok()?;
|
||||
Some(ImageProgress {
|
||||
done: parsed.done,
|
||||
step: parsed.step,
|
||||
total: parsed.total,
|
||||
image: parsed.image,
|
||||
})
|
||||
}
|
||||
|
||||
/// Channel events for streaming image generation.
|
||||
#[derive(Clone, Serialize)]
|
||||
#[serde(tag = "type", content = "data", rename_all = "camelCase")]
|
||||
pub enum ImageEvent {
|
||||
/// Progress update: (step, total). Either may be None if Ollama omits it.
|
||||
/// Progress update: (step, total). Either may be None if the server omits it.
|
||||
Progress { step: Option<u32>, total: Option<u32> },
|
||||
/// Final result: a `data:image/png;base64,...` URL ready for `<img src>`.
|
||||
Done(String),
|
||||
@@ -68,11 +34,23 @@ pub enum ImageEvent {
|
||||
Error(String),
|
||||
}
|
||||
|
||||
/// Generate (or fetch from disk cache) an image for `prompt` via the configured
|
||||
/// Ollama image model. Returns a `data:image/png;base64,...` URL ready for `<img src>`.
|
||||
///
|
||||
/// Ollama image models are macOS-only today; on other platforms we return an error
|
||||
/// so the front-end can fall back to a placeholder instead of a confusing timeout.
|
||||
/// A1111 `/sdapi/v1/txt2img` response. Only the base64 PNG list matters to us.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct Txt2ImgResponse {
|
||||
#[serde(default)]
|
||||
images: Vec<String>,
|
||||
}
|
||||
|
||||
/// A1111 `/sdapi/v1/progress` response. `progress` is 0.0–1.0 of the current job.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ProgressResponse {
|
||||
#[serde(default)]
|
||||
progress: f32,
|
||||
}
|
||||
|
||||
/// Generate (or fetch from disk cache) an image for `prompt` via a local
|
||||
/// stable-diffusion.cpp `sd-server`. Returns a `data:image/png;base64,...` URL
|
||||
/// ready for `<img src>`.
|
||||
#[tauri::command]
|
||||
pub async fn generate_image(
|
||||
state: tauri::State<'_, AppState>,
|
||||
@@ -88,27 +66,23 @@ pub async fn generate_image(
|
||||
}
|
||||
|
||||
async fn generate_image_inner(state: tauri::State<'_, AppState>, req: ImageRequest) -> Result<String, String> {
|
||||
if cfg!(not(target_os = "macos")) {
|
||||
return Err("image generation is macOS-only via Ollama (for now)".into());
|
||||
}
|
||||
|
||||
let config = state.config.lock().map_err(|e| e.to_string())?.clone();
|
||||
let model = req.model.unwrap_or(config.image_model.clone());
|
||||
let api_url = config.image_api_url.clone();
|
||||
|
||||
let (cache_path, cache_hit) = prepare_cache(&state.data_dir, &model, &req.prompt)?;
|
||||
let (cache_path, cache_hit) = prepare_cache(&state.data_dir, &api_url, &req.prompt)?;
|
||||
if cache_hit {
|
||||
let bytes = std::fs::read(&cache_path).map_err(|e| e.to_string())?;
|
||||
return Ok(data_url(&bytes));
|
||||
}
|
||||
|
||||
let png_bytes = request_image_bytes(&config, &model, &req.prompt, None).await?;
|
||||
let png_bytes = txt2img(&api_url, &req.prompt).await?;
|
||||
std::fs::write(&cache_path, &png_bytes).map_err(|e| e.to_string())?;
|
||||
Ok(data_url(&png_bytes))
|
||||
}
|
||||
|
||||
/// Streaming variant: emits `ImageEvent::Progress` as Ollama reports `step`/`total`,
|
||||
/// then `ImageEvent::Done` with the data URL (or `Error`). Reuses the same disk cache
|
||||
/// as `generate_image`. The DM gets a real progress bar for the multi-second wait.
|
||||
/// Streaming variant: emits `ImageEvent::Progress` (from `/sdapi/v1/progress`)
|
||||
/// while the txt2img POST is in flight, then `ImageEvent::Done` with the data
|
||||
/// URL (or `Error`). Reuses the same disk cache as `generate_image`.
|
||||
#[tauri::command]
|
||||
pub async fn generate_image_stream(
|
||||
state: tauri::State<'_, AppState>,
|
||||
@@ -116,17 +90,10 @@ pub async fn generate_image_stream(
|
||||
req: ImageRequest,
|
||||
channel: Channel<ImageEvent>,
|
||||
) -> Result<(), String> {
|
||||
if cfg!(not(target_os = "macos")) {
|
||||
let _ = channel.send(ImageEvent::Error(
|
||||
"image generation is macOS-only via Ollama (for now)".into(),
|
||||
));
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let config = state.config.lock().map_err(|e| e.to_string())?.clone();
|
||||
let model = req.model.unwrap_or(config.image_model.clone());
|
||||
let api_url = config.image_api_url.clone();
|
||||
|
||||
let (cache_path, cache_hit) = prepare_cache(&state.data_dir, &model, &req.prompt)?;
|
||||
let (cache_path, cache_hit) = prepare_cache(&state.data_dir, &api_url, &req.prompt)?;
|
||||
if cache_hit {
|
||||
let bytes = std::fs::read(&cache_path).map_err(|e| e.to_string())?;
|
||||
let _ = channel.send(ImageEvent::Done(data_url(&bytes)));
|
||||
@@ -135,19 +102,41 @@ pub async fn generate_image_stream(
|
||||
|
||||
// Drive the request on a background task so the command returns immediately
|
||||
// and progress flows through the channel. Errors become ImageEvent::Error.
|
||||
let channel = std::sync::Arc::new(channel);
|
||||
let channel = Arc::new(channel);
|
||||
let ch = channel.clone();
|
||||
// ponytail: only emit busy around the actual generation (not cache hits),
|
||||
// and always balance it in the spawn — even on error.
|
||||
emit_busy(&app, true);
|
||||
tauri::async_runtime::spawn(async move {
|
||||
match request_image_bytes(&config, &model, &req.prompt, Some(ch)).await {
|
||||
// Poll /sdapi/v1/progress while the txt2img POST runs so the DM sees a
|
||||
// real progress bar. Best-effort: poll errors are silent (bar falls
|
||||
// back to indeterminate). 200ms is smooth without spamming the server.
|
||||
let poll_client = reqwest::Client::new();
|
||||
let done = Arc::new(AtomicBool::new(false));
|
||||
let done_p = done.clone();
|
||||
let api_p = api_url.clone();
|
||||
let ch_p = ch.clone();
|
||||
let poller = tauri::async_runtime::spawn(async move {
|
||||
while !done_p.load(Ordering::SeqCst) {
|
||||
if let Ok(p) = poll_progress(&poll_client, &api_p).await {
|
||||
let step = (p.clamp(0.0, 1.0) * 100.0) as u32;
|
||||
let _ = ch_p.send(ImageEvent::Progress { step: Some(step), total: Some(100) });
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(200)).await;
|
||||
}
|
||||
});
|
||||
|
||||
let result = txt2img(&api_url, &req.prompt).await;
|
||||
done.store(true, Ordering::SeqCst);
|
||||
let _ = poller.await; // let the poller flush its last iteration
|
||||
|
||||
match result {
|
||||
Ok(png_bytes) => {
|
||||
let _ = std::fs::write(&cache_path, &png_bytes);
|
||||
let _ = channel.send(ImageEvent::Done(data_url(&png_bytes)));
|
||||
let _ = ch.send(ImageEvent::Done(data_url(&png_bytes)));
|
||||
}
|
||||
Err(e) => {
|
||||
let _ = channel.send(ImageEvent::Error(e));
|
||||
let _ = ch.send(ImageEvent::Error(e));
|
||||
}
|
||||
}
|
||||
emit_busy(&app, false);
|
||||
@@ -156,43 +145,61 @@ pub async fn generate_image_stream(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Resolve the on-disk cache path for (model, prompt) and report a cache hit.
|
||||
/// Resolve the on-disk cache path for (api_url, prompt) and report a cache hit.
|
||||
/// `base` is the configured data dir (AppState.data_dir).
|
||||
fn prepare_cache(
|
||||
base: &std::path::Path,
|
||||
model: &str,
|
||||
api_url: &str,
|
||||
prompt: &str,
|
||||
) -> Result<(std::path::PathBuf, bool), String> {
|
||||
let cache_dir = base.join("images");
|
||||
std::fs::create_dir_all(&cache_dir).map_err(|e| e.to_string())?;
|
||||
|
||||
let mut hasher = DefaultHasher::new();
|
||||
model.hash(&mut hasher);
|
||||
api_url.hash(&mut hasher);
|
||||
prompt.hash(&mut hasher);
|
||||
let cache_path = cache_dir.join(format!("{:016x}.png", hasher.finish()));
|
||||
let hit = cache_path.exists();
|
||||
Ok((cache_path, hit))
|
||||
}
|
||||
|
||||
/// POST the image request to Ollama and collect the PNG bytes. When a channel
|
||||
/// is given, parse the NDJSON body incrementally and emit progress; otherwise
|
||||
/// read the whole body at once (legacy buffered path).
|
||||
async fn request_image_bytes(
|
||||
config: &crate::llm::LlmConfig,
|
||||
model: &str,
|
||||
prompt: &str,
|
||||
channel: Option<std::sync::Arc<Channel<ImageEvent>>>,
|
||||
) -> Result<Vec<u8>, String> {
|
||||
let client = reqwest::Client::new();
|
||||
let url = format!("{}/api/generate", config.api_url.trim_end_matches('/'));
|
||||
let body = serde_json::json!({ "model": model, "prompt": prompt, "stream": false });
|
||||
/// POST `/sdapi/v1/txt2img` to the sd-server and return the raw PNG bytes of
|
||||
/// the first generated image. A1111 returns base64 PNGs *without* a `data:`
|
||||
/// prefix, so we decode to bytes for the disk cache.
|
||||
async fn txt2img(api_url: &str, prompt: &str) -> Result<Vec<u8>, String> {
|
||||
// ponytail: width/height/steps/cfg are fixed — diffusion can take tens of
|
||||
// seconds; add config knobs only if the DM wants to tune them.
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(300))
|
||||
.build()
|
||||
.map_err(|e| e.to_string())?;
|
||||
let url = format!("{}/sdapi/v1/txt2img", api_url.trim_end_matches('/'));
|
||||
let body = serde_json::json!({
|
||||
"prompt": prompt,
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"steps": 20,
|
||||
"cfg_scale": 7.0,
|
||||
"seed": -1,
|
||||
"batch_size": 1,
|
||||
});
|
||||
|
||||
let res = client
|
||||
.post(&url)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("image request failed: {e}"))?;
|
||||
.map_err(|e| {
|
||||
// ponytail: non-technical DMs see a connect error as a mystery —
|
||||
// name the fix (start sd-server) and point at the in-app guide.
|
||||
if e.is_connect() || e.is_timeout() {
|
||||
format!(
|
||||
"Can't reach the image server at {api_url}. Is sd-server running? Open the Image tab → 'First time setup' for step-by-step instructions. ({e})"
|
||||
)
|
||||
} else {
|
||||
format!("image request failed: {e}")
|
||||
}
|
||||
})?;
|
||||
|
||||
if !res.status().is_success() {
|
||||
let status = res.status();
|
||||
@@ -200,56 +207,71 @@ async fn request_image_bytes(
|
||||
return Err(format!("image error {status}: {text}"));
|
||||
}
|
||||
|
||||
// ponytail: Ollama image models emit NDJSON even with stream:false —
|
||||
// progress lines (step/total) then a final done:true carrying the image.
|
||||
// The buffered path reads it all at once; the streaming path splits lines
|
||||
// as they arrive so the bar animates.
|
||||
let mut png_b64: Option<String> = None;
|
||||
let mut buffer = String::new();
|
||||
let parsed: Txt2ImgResponse = res
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| format!("image parse error: {e}"))?;
|
||||
let b64 = parsed
|
||||
.images
|
||||
.into_iter()
|
||||
.next()
|
||||
.ok_or_else(|| "no image in sd-server response".to_string())?;
|
||||
base64_decode(&b64)
|
||||
}
|
||||
|
||||
let apply_line = |line: &str, png: &mut Option<String>, ch: Option<&Channel<ImageEvent>>| {
|
||||
if let Some(p) = parse_image_line(line) {
|
||||
if let Some(c) = ch {
|
||||
let _ = c.send(ImageEvent::Progress { step: p.step, total: p.total });
|
||||
}
|
||||
if let Some(b64) = p.image {
|
||||
*png = Some(b64);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(ch) = channel {
|
||||
let mut stream = res.bytes_stream();
|
||||
while let Some(chunk) = stream.next().await {
|
||||
let chunk = chunk.map_err(|e| format!("image read error: {e}"))?;
|
||||
buffer.push_str(&String::from_utf8_lossy(&chunk));
|
||||
// Process complete lines; keep the trailing partial line in buffer.
|
||||
while let Some(idx) = buffer.find('\n') {
|
||||
let line = buffer.split_off(idx + 1);
|
||||
let complete = std::mem::replace(&mut buffer, line);
|
||||
apply_line(&complete, &mut png_b64, Some(&ch));
|
||||
}
|
||||
}
|
||||
if !buffer.trim().is_empty() {
|
||||
apply_line(&buffer, &mut png_b64, Some(&ch));
|
||||
}
|
||||
} else {
|
||||
let body_text = res.text().await.map_err(|e| format!("image read error: {e}"))?;
|
||||
for line in body_text.lines() {
|
||||
apply_line(line, &mut png_b64, None);
|
||||
// ponytail: stop after the done line in the buffered path — the
|
||||
// final image is the one we want.
|
||||
if let Some(p) = parse_image_line(line) {
|
||||
if p.done && p.image.is_some() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
/// GET `/sdapi/v1/progress` → fraction (0.0–1.0) of the current job. Best-effort.
|
||||
async fn poll_progress(client: &reqwest::Client, api_url: &str) -> Result<f32, String> {
|
||||
let url = format!("{}/sdapi/v1/progress", api_url.trim_end_matches('/'));
|
||||
let res = client
|
||||
.get(&url)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
if !res.status().is_success() {
|
||||
return Err(format!("progress HTTP {}", res.status()));
|
||||
}
|
||||
let p: ProgressResponse = res.json().await.map_err(|e| e.to_string())?;
|
||||
Ok(p.progress)
|
||||
}
|
||||
|
||||
let b64 = png_b64.ok_or_else(|| "no image data in Ollama response".to_string())?;
|
||||
let png_bytes = base64_decode(&b64)?;
|
||||
Ok(png_bytes)
|
||||
/// GET `/sdapi/v1/progress` as a liveness probe for the configured sd-server.
|
||||
/// Returns ok=true when the server responds (even idle — progress is 0.0).
|
||||
/// Mirrors the LLM `test_connection` so Settings can show a green check.
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct ImageConnectionTest {
|
||||
pub ok: bool,
|
||||
pub error: String,
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn test_image_connection(
|
||||
state: tauri::State<'_, AppState>,
|
||||
) -> Result<ImageConnectionTest, String> {
|
||||
let api_url = state
|
||||
.config
|
||||
.lock()
|
||||
.map_err(|e| e.to_string())?
|
||||
.image_api_url
|
||||
.clone();
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(5))
|
||||
.build()
|
||||
.map_err(|e| e.to_string())?;
|
||||
let url = format!("{}/sdapi/v1/progress", api_url.trim_end_matches('/'));
|
||||
match client.get(&url).send().await {
|
||||
Ok(res) if res.status().is_success() => Ok(ImageConnectionTest {
|
||||
ok: true,
|
||||
error: String::new(),
|
||||
}),
|
||||
Ok(res) => Ok(ImageConnectionTest {
|
||||
ok: false,
|
||||
error: format!("HTTP {} — is this an sd-server?", res.status()),
|
||||
}),
|
||||
Err(e) => Ok(ImageConnectionTest {
|
||||
ok: false,
|
||||
error: format!("Can't reach {api_url} — is sd-server running? ({e})"),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn data_url(png: &[u8]) -> String {
|
||||
@@ -336,55 +358,19 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn cache_key_is_stable() {
|
||||
// sanity: same inputs → same filename shape (16 hex digits)
|
||||
// sanity: same inputs → same 16-hex-digit filename shape
|
||||
let mut h = DefaultHasher::new();
|
||||
"x/flux2-klein:4b".hash(&mut h);
|
||||
"http://localhost:1234".hash(&mut h);
|
||||
"prompt".hash(&mut h);
|
||||
let s = format!("{:016x}", h.finish());
|
||||
assert_eq!(s.len(), 16);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_progress_and_image_lines() {
|
||||
// intermediate progress line, no image
|
||||
let p = parse_image_line(r#"{"model":"x/flux2-klein:4b","done":false,"total":4,"step":1}"#).unwrap();
|
||||
assert_eq!(p, ImageProgress { done: false, step: Some(1), total: Some(4), image: None });
|
||||
|
||||
// final done line carries the image
|
||||
let p = parse_image_line(r#"{"model":"x/flux2-klein:4b","done":true,"image":"iVBORw0KGgoAAAANSUhEUgAA"}"#).unwrap();
|
||||
assert!(p.done);
|
||||
assert_eq!(p.image.as_deref(), Some("iVBORw0KGgoAAAANSUhEUgAA"));
|
||||
assert!(p.total.is_none() && p.step.is_none());
|
||||
|
||||
// blank/garbage lines are ignored, not errors
|
||||
assert!(parse_image_line("").is_none());
|
||||
assert!(parse_image_line("not json").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_ndjson_buffer_keeps_trailing_partial() {
|
||||
// Simulate two chunks arriving separately where the split falls mid-line.
|
||||
let chunk1 = "{\"done\":false,\"step\":1,\"total\":4}\n{\"done\":tru";
|
||||
let chunk2 = "e,\"image\":\"abc\"}\n";
|
||||
let mut buffer = String::new();
|
||||
let mut png: Option<String> = None;
|
||||
let whole = format!("{chunk1}{chunk2}");
|
||||
// emulate the streaming loop over the concatenated body
|
||||
buffer.push_str(&whole);
|
||||
let mut lines = Vec::new();
|
||||
while let Some(idx) = buffer.find('\n') {
|
||||
let line = buffer.split_off(idx + 1);
|
||||
let complete = std::mem::replace(&mut buffer, line);
|
||||
lines.push(complete);
|
||||
}
|
||||
if !buffer.trim().is_empty() {
|
||||
lines.push(std::mem::take(&mut buffer));
|
||||
}
|
||||
for line in &lines {
|
||||
if let Some(p) = parse_image_line(line) {
|
||||
if let Some(b) = p.image { png = Some(b); }
|
||||
}
|
||||
}
|
||||
assert_eq!(png.as_deref(), Some("abc"));
|
||||
fn progress_fraction_maps_to_step() {
|
||||
// The streaming poller emits step = clamp(progress,0,1)*100.
|
||||
assert_eq!((0.0f32 * 100.0) as u32, 0);
|
||||
assert_eq!((0.5f32 * 100.0) as u32, 50);
|
||||
assert_eq!((1.5f32.clamp(0.0, 1.0) * 100.0) as u32, 100);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user