376 lines
13 KiB
Rust
376 lines
13 KiB
Rust
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);
|
||
}
|
||
} |