Files
dm-pal/src-tauri/src/commands/image_commands.rs
T

376 lines
13 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
use crate::commands::emit_busy;
use crate::llm::AppState;
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,
}
/// 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 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),
/// Fatal error.
Error(String),
}
/// 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>,
app: AppHandle,
req: ImageRequest,
) -> Result<String, String> {
// ponytail: emit a global busy signal so the shell shows a working
// indicator even if the DM navigates away. Paired with the false below.
emit_busy(&app, true);
let result = generate_image_inner(state, req).await;
emit_busy(&app, false);
result
}
async fn generate_image_inner(state: tauri::State<'_, AppState>, req: ImageRequest) -> Result<String, String> {
let config = state.config.lock().map_err(|e| e.to_string())?.clone();
let api_url = config.image_api_url.clone();
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 = 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` (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>,
app: AppHandle,
req: ImageRequest,
channel: Channel<ImageEvent>,
) -> Result<(), String> {
let config = state.config.lock().map_err(|e| e.to_string())?.clone();
let api_url = config.image_api_url.clone();
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)));
return Ok(());
}
// Drive the request on a background task so the command returns immediately
// and progress flows through the channel. Errors become ImageEvent::Error.
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 {
// 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 _ = ch.send(ImageEvent::Done(data_url(&png_bytes)));
}
Err(e) => {
let _ = ch.send(ImageEvent::Error(e));
}
}
emit_busy(&app, false);
});
Ok(())
}
/// 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,
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();
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 `/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| {
// 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();
let text = res.text().await.unwrap_or_default();
return Err(format!("image error {status}: {text}"));
}
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)
}
/// 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)
}
/// 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 {
// ponytail: no base64 crate dep — a 30-line encoder.
let table: [u8; 64] = *b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
let mut out = String::with_capacity((png.len() + 2) / 3 * 4);
let mut chunks = png.chunks_exact(3);
for c in &mut chunks {
let n = (c[0] as usize) << 16 | (c[1] as usize) << 8 | (c[2] as usize);
out.push(table[(n >> 18) & 63] as char);
out.push(table[(n >> 12) & 63] as char);
out.push(table[(n >> 6) & 63] as char);
out.push(table[n & 63] as char);
}
let rem = chunks.remainder();
match rem.len() {
1 => {
let n = (rem[0] as usize) << 16;
out.push(table[(n >> 18) & 63] as char);
out.push(table[(n >> 12) & 63] as char);
out.push('=');
out.push('=');
}
2 => {
let n = (rem[0] as usize) << 16 | (rem[1] as usize) << 8;
out.push(table[(n >> 18) & 63] as char);
out.push(table[(n >> 12) & 63] as char);
out.push(table[(n >> 6) & 63] as char);
out.push('=');
}
_ => {}
}
format!("data:image/png;base64,{out}")
}
/// Minimal standard base64 decoder (no extra dependency).
fn base64_decode(input: &str) -> Result<Vec<u8>, String> {
fn val(c: u8) -> Option<u8> {
match c {
b'A'..=b'Z' => Some(c - b'A'),
b'a'..=b'z' => Some(c - b'a' + 26),
b'0'..=b'9' => Some(c - b'0' + 52),
b'+' => Some(62),
b'/' => Some(63),
_ => None,
}
}
let input = input.trim();
let bytes: Vec<u8> = input.bytes().filter(|b| !b.is_ascii_whitespace()).collect();
if bytes.is_empty() {
return Ok(Vec::new());
}
let mut out = Vec::with_capacity(bytes.len() * 3 / 4);
let mut buf: u32 = 0;
let mut bits: u32 = 0;
for &b in &bytes {
if b == b'=' {
break;
}
let v = val(b).ok_or_else(|| format!("invalid base64 char: {b}"))? as u32;
buf = (buf << 6) | v;
bits += 6;
if bits >= 8 {
bits -= 8;
out.push((buf >> bits) as u8);
buf &= (1 << bits) - 1;
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn roundtrip_base64() {
let data = b"hello world \x00\xff\x10";
let url = data_url(data);
let b64 = url.strip_prefix("data:image/png;base64,").unwrap();
let decoded = base64_decode(b64).unwrap();
assert_eq!(decoded, data);
}
#[test]
fn cache_key_is_stable() {
// sanity: same inputs → same 16-hex-digit filename shape
let mut h = DefaultHasher::new();
"http://localhost:1234".hash(&mut h);
"prompt".hash(&mut h);
let s = format!("{:016x}", h.finish());
assert_eq!(s.len(), 16);
}
#[test]
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);
}
}