Skip to main content

mold_server/
queue.rs

1use std::collections::VecDeque;
2use std::sync::Arc;
3
4use base64::Engine as _;
5use mold_core::{
6    ImageData, OutputFormat, OutputMetadata, SseCompleteEvent, SseErrorEvent, SseProgressEvent,
7};
8use mold_db::{MetadataDb, RecordSource};
9use sha2::{Digest, Sha256};
10use std::sync::atomic::{AtomicBool, Ordering};
11use std::time::Instant;
12use tokio::sync::Notify;
13
14use crate::gpu_pool::GpuJob;
15use crate::model_manager;
16use crate::state::{
17    ActiveGenerationSnapshot, AppState, GenerationJob, GenerationJobResult, SseCompletionPayload,
18    SseMessage,
19};
20
21/// Convert an inference-crate progress event to an SSE wire event.
22fn progress_to_sse(event: mold_inference::ProgressEvent) -> SseProgressEvent {
23    event.into()
24}
25
26/// Strips backtrace frames from candle error messages.
27///
28/// Renders the full anyhow cause chain (`{:#}`) so wrappers like
29/// `with_context("mmap single-file checkpoint at …")` carry their root cause
30/// through to the wire — otherwise users see the outer wrapper only.
31pub(crate) fn clean_error_message(e: &anyhow::Error) -> String {
32    let full = format!("{e:#}");
33    let mut lines: Vec<&str> = Vec::new();
34    for line in full.lines() {
35        let trimmed = line.trim_start();
36        if (trimmed.starts_with("0:") || trimmed.starts_with("1:"))
37            && trimmed.len() > 3
38            && trimmed
39                .as_bytes()
40                .first()
41                .is_some_and(|b| b.is_ascii_digit())
42        {
43            break;
44        }
45        if trimmed.len() > 2
46            && trimmed.as_bytes()[0].is_ascii_digit()
47            && trimmed.contains("::")
48            && trimmed.contains("at ")
49        {
50            break;
51        }
52        lines.push(line);
53    }
54    let msg = lines.join("\n").trim().to_string();
55    if msg.is_empty() {
56        format!("{}", e.root_cause())
57    } else {
58        msg
59    }
60}
61
62fn set_active_generation(state: &AppState, model: &str, prompt: &str) {
63    let prompt_sha256 = format!("{:x}", Sha256::digest(prompt.as_bytes()));
64    let started_at_unix_ms = mold_core::time::now_epoch_ms_u64();
65
66    let mut active = state
67        .active_generation
68        .write()
69        .unwrap_or_else(|e| e.into_inner());
70    *active = Some(ActiveGenerationSnapshot {
71        model: model.to_string(),
72        prompt_sha256,
73        started_at_unix_ms,
74        started_at: Instant::now(),
75    });
76}
77
78fn clear_active_generation(state: &AppState) {
79    let mut active = state
80        .active_generation
81        .write()
82        .unwrap_or_else(|e| e.into_inner());
83    *active = None;
84}
85
86/// Test-facing single-image wrapper around the shared output persistence path.
87///
88/// Errors writing to disk are logged and skipped. DB errors are also logged
89/// but do not fail the save — the file is the source of truth.
90///
91/// Shared between the legacy single-GPU `process_job` (this file) and the
92/// per-GPU worker (`gpu_worker.rs`). Keep these on one helper so the DB
93/// upsert can never silently regress on one path while the other keeps
94/// working.
95#[allow(clippy::too_many_arguments)]
96#[cfg(test)]
97pub(crate) fn save_image_to_dir(
98    dir: &std::path::Path,
99    img: &mold_core::ImageData,
100    model: &str,
101    batch_size: u32,
102    metadata: Option<&OutputMetadata>,
103    generation_time_ms: Option<i64>,
104    db: Option<&MetadataDb>,
105    events: Option<&crate::events::EventBroadcaster>,
106) {
107    save_image_to_dir_with_suffix(
108        dir,
109        img,
110        model,
111        batch_size,
112        None,
113        metadata,
114        generation_time_ms,
115        db,
116        events,
117    );
118}
119
120#[allow(clippy::too_many_arguments)]
121fn save_image_to_dir_with_suffix(
122    dir: &std::path::Path,
123    img: &mold_core::ImageData,
124    model: &str,
125    batch_size: u32,
126    suffix: Option<&str>,
127    metadata: Option<&OutputMetadata>,
128    generation_time_ms: Option<i64>,
129    db: Option<&MetadataDb>,
130    events: Option<&crate::events::EventBroadcaster>,
131) -> Option<String> {
132    if let Err(e) = std::fs::create_dir_all(dir) {
133        tracing::warn!("failed to create output dir {}: {e}", dir.display());
134        return None;
135    }
136    let timestamp_ms = mold_core::time::now_epoch_ms_u64();
137    let ext = img.format.to_string();
138    let mut filename =
139        mold_core::default_output_filename(model, timestamp_ms, &ext, batch_size, img.index);
140    if let Some(suffix) = suffix {
141        filename = format!(
142            "{}-{suffix}.{ext}",
143            filename.trim_end_matches(&format!(".{ext}"))
144        );
145    }
146    let path = dir.join(&filename);
147    match std::fs::write(&path, &img.data) {
148        Ok(()) => tracing::info!("saved image to {}", path.display()),
149        Err(e) => {
150            tracing::warn!("failed to save image to {}: {e}", path.display());
151            return None;
152        }
153    }
154    let mut image_row = None;
155    if let (Some(db), Some(meta)) = (db, metadata) {
156        image_row = mold_db::persist::record_saved_output_returning(
157            db,
158            dir,
159            &filename,
160            &path,
161            &mold_db::persist::OutputRecordParams {
162                format: img.format,
163                metadata: meta,
164                source: RecordSource::Server,
165                generation_time_ms,
166                backend: Some(mold_inference::compiled_backend_label()),
167            },
168        )
169        .map(|rec| Box::new(rec.to_gallery_image()));
170    }
171    // Emit even without a DB row — `image: None` tells clients to refetch
172    // `/api/gallery` instead of inserting in place.
173    if let Some(events) = events {
174        events.publish(mold_core::ServerEvent::GalleryAdded {
175            filename: filename.clone(),
176            image: image_row,
177        });
178    }
179    Some(filename)
180}
181
182/// Gallery filenames a generation's outputs were saved under, threaded into
183/// the SSE complete event so mirroring clients keep the same identity.
184#[derive(Debug, Default, Clone)]
185pub(crate) struct SavedOutputNames {
186    /// The payload the complete event carries (upscaled when upscaling ran).
187    pub output: Option<String>,
188    /// The pre-upscale original, when one was saved separately.
189    pub original: Option<String>,
190}
191
192#[allow(clippy::too_many_arguments)]
193pub(crate) fn save_generated_image_outputs(
194    dir: &std::path::Path,
195    original: Option<&ImageData>,
196    output: &ImageData,
197    model: &str,
198    batch_size: u32,
199    metadata: &OutputMetadata,
200    generation_time_ms: Option<i64>,
201    db: Option<&MetadataDb>,
202    events: Option<&crate::events::EventBroadcaster>,
203) -> SavedOutputNames {
204    let mut names = SavedOutputNames::default();
205    if let Some(original) = original {
206        let mut original_metadata = metadata.clone();
207        apply_output_dimensions_to_metadata(&mut original_metadata, original);
208        names.original = save_image_to_dir_with_suffix(
209            dir,
210            original,
211            model,
212            batch_size,
213            Some("original"),
214            Some(&original_metadata),
215            generation_time_ms,
216            db,
217            events,
218        );
219    }
220    let mut output_metadata = metadata.clone();
221    apply_output_dimensions_to_metadata(&mut output_metadata, output);
222    names.output = save_image_to_dir_with_suffix(
223        dir,
224        output,
225        model,
226        batch_size,
227        original.map(|_| "upscaled"),
228        Some(&output_metadata),
229        generation_time_ms,
230        db,
231        events,
232    );
233    names
234}
235
236/// Save a video file to disk and (best-effort) record its metadata row.
237/// Mirrors `save_image_to_dir` for the video-output path. See that helper
238/// for the multi-path-callers note.
239///
240/// When `gif_preview` is non-empty, also persists
241/// `$MOLD_HOME/cache/previews/<filename>.preview.gif`. The gallery preview
242/// endpoint (`GET /api/gallery/preview/:filename`) streams from that path
243/// so remote TUI clients can animate the detail pane without re-fetching
244/// the full MP4.
245#[allow(clippy::too_many_arguments)]
246pub(crate) fn save_video_to_dir(
247    dir: &std::path::Path,
248    bytes: &[u8],
249    gif_preview: &[u8],
250    format: OutputFormat,
251    model: &str,
252    metadata: &OutputMetadata,
253    generation_time_ms: Option<i64>,
254    db: Option<&MetadataDb>,
255    events: Option<&crate::events::EventBroadcaster>,
256) -> Option<String> {
257    if let Err(e) = std::fs::create_dir_all(dir) {
258        tracing::warn!("failed to create output dir {}: {e}", dir.display());
259        return None;
260    }
261    let ts = mold_core::time::now_epoch_ms_u64();
262    let ext = format.extension();
263    let filename = mold_core::default_output_filename(model, ts, ext, 1, 0);
264    let path = dir.join(&filename);
265    if let Err(e) = std::fs::write(&path, bytes) {
266        tracing::error!("failed to save video to {}: {e}", path.display());
267        return None;
268    }
269    if !gif_preview.is_empty() {
270        save_video_preview_gif(&filename, gif_preview);
271    }
272    let mut image_row = None;
273    if let Some(db) = db {
274        image_row = mold_db::persist::record_saved_output_returning(
275            db,
276            dir,
277            &filename,
278            &path,
279            &mold_db::persist::OutputRecordParams {
280                format,
281                metadata,
282                source: RecordSource::Server,
283                generation_time_ms,
284                backend: Some(mold_inference::compiled_backend_label()),
285            },
286        )
287        .map(|rec| Box::new(rec.to_gallery_image()));
288    }
289    if let Some(events) = events {
290        events.publish(mold_core::ServerEvent::GalleryAdded {
291            filename: filename.clone(),
292            image: image_row,
293        });
294    }
295    Some(filename)
296}
297
298fn requested_post_upscale_model(req: &mold_core::GenerateRequest) -> Option<&str> {
299    req.upscale_model
300        .as_deref()
301        .map(str::trim)
302        .filter(|m| !m.is_empty())
303}
304
305fn post_upscale_model_to_pull(
306    config: &mold_core::Config,
307    req: &mold_core::GenerateRequest,
308) -> Result<Option<String>, String> {
309    let Some(requested) = requested_post_upscale_model(req) else {
310        return Ok(None);
311    };
312    let model_name = mold_core::manifest::resolve_model_name(requested);
313    if model_manager::configured_upscaler_weights_exist(config, &model_name) {
314        return Ok(None);
315    }
316    if mold_core::manifest::find_manifest(&model_name).is_none() {
317        return Err(format!("unknown upscaler model '{model_name}'"));
318    }
319    Ok(Some(model_name))
320}
321
322async fn ensure_post_upscale_model_downloaded(
323    state: &AppState,
324    req: &mold_core::GenerateRequest,
325    progress_tx: Option<&tokio::sync::mpsc::UnboundedSender<SseMessage>>,
326) -> Result<(), String> {
327    let model_to_pull = {
328        let config = state.config.read().await;
329        post_upscale_model_to_pull(&config, req)?
330    };
331    let Some(model_name) = model_to_pull else {
332        return Ok(());
333    };
334
335    if let Some(tx) = progress_tx {
336        let _ = tx.send(SseMessage::Progress(SseProgressEvent::StageStart {
337            name: format!("Downloading upscaler {model_name}"),
338        }));
339    }
340    let progress = progress_tx.cloned().map(|tx| {
341        Arc::new(move |event: mold_core::download::DownloadProgressEvent| {
342            let event = match event {
343                mold_core::download::DownloadProgressEvent::Status { message } => {
344                    SseProgressEvent::Info { message }
345                }
346                mold_core::download::DownloadProgressEvent::FileStart {
347                    filename,
348                    file_index,
349                    total_files,
350                    size_bytes,
351                    batch_bytes_downloaded,
352                    batch_bytes_total,
353                    batch_elapsed_ms,
354                } => SseProgressEvent::DownloadProgress {
355                    filename,
356                    file_index,
357                    total_files,
358                    bytes_downloaded: 0,
359                    bytes_total: size_bytes,
360                    batch_bytes_downloaded,
361                    batch_bytes_total,
362                    batch_elapsed_ms,
363                },
364                mold_core::download::DownloadProgressEvent::FileProgress {
365                    filename,
366                    file_index,
367                    bytes_downloaded,
368                    bytes_total,
369                    batch_bytes_downloaded,
370                    batch_bytes_total,
371                    batch_elapsed_ms,
372                } => SseProgressEvent::DownloadProgress {
373                    filename,
374                    file_index,
375                    total_files: 0,
376                    bytes_downloaded,
377                    bytes_total,
378                    batch_bytes_downloaded,
379                    batch_bytes_total,
380                    batch_elapsed_ms,
381                },
382                mold_core::download::DownloadProgressEvent::FileDone {
383                    filename,
384                    file_index,
385                    total_files,
386                    batch_bytes_downloaded,
387                    batch_bytes_total,
388                    batch_elapsed_ms,
389                } => SseProgressEvent::DownloadDone {
390                    filename,
391                    file_index,
392                    total_files,
393                    batch_bytes_downloaded,
394                    batch_bytes_total,
395                    batch_elapsed_ms,
396                },
397            };
398            let _ = tx.send(SseMessage::Progress(event));
399        }) as model_manager::DownloadProgressCallback
400    });
401    model_manager::pull_model(state, &model_name, progress)
402        .await
403        .map_err(|e| format!("failed to pull upscaler model: {}", e.error))?;
404    if let Some(tx) = progress_tx {
405        let _ = tx.send(SseMessage::Progress(SseProgressEvent::PullComplete {
406            model: model_name,
407        }));
408    }
409    Ok(())
410}
411
412pub(crate) fn apply_output_dimensions_to_metadata(metadata: &mut OutputMetadata, img: &ImageData) {
413    metadata.apply_output_dimensions(img.width, img.height);
414}
415
416pub(crate) fn apply_upscale_response_to_image_generation(
417    req: &mold_core::GenerateRequest,
418    response: &mut mold_core::GenerateResponse,
419    original: ImageData,
420    upscaled: mold_core::UpscaleResponse,
421) -> anyhow::Result<ImageData> {
422    if response.video.is_some() || requested_post_upscale_model(req).is_none() {
423        return Ok(original);
424    }
425    if upscaled.image.data.is_empty() {
426        anyhow::bail!("upscaler returned an empty image");
427    }
428    response.generation_time_ms = response
429        .generation_time_ms
430        .saturating_add(upscaled.upscale_time_ms);
431    Ok(ImageData {
432        index: original.index,
433        ..upscaled.image
434    })
435}
436
437pub(crate) fn settle_post_generation_upscale(
438    original: ImageData,
439    result: Result<ImageData, String>,
440) -> (ImageData, Option<ImageData>, Option<String>) {
441    match result {
442        Ok(upscaled) => (upscaled, Some(original), None),
443        Err(error) => (original, None, Some(error)),
444    }
445}
446
447async fn upscale_generated_image_on_single_worker(
448    state: &AppState,
449    req: &mold_core::GenerateRequest,
450    seed_used: u64,
451    img: ImageData,
452    progress_tx: Option<&tokio::sync::mpsc::UnboundedSender<SseMessage>>,
453) -> Result<ImageData, String> {
454    let Some(upscale_model) = requested_post_upscale_model(req).map(str::to_string) else {
455        return Ok(img);
456    };
457    let model_name = mold_core::manifest::resolve_model_name(&upscale_model);
458    if let Some(tx) = progress_tx {
459        let _ = tx.send(SseMessage::Progress(SseProgressEvent::StageStart {
460            name: format!("Loading upscaler {model_name}"),
461        }));
462    }
463
464    let needs_pull = {
465        let config = state.config.read().await;
466        config
467            .models
468            .get(&model_name)
469            .and_then(|c| c.transformer.as_ref())
470            .is_none()
471    };
472    if needs_pull {
473        if mold_core::manifest::find_manifest(&model_name).is_none() {
474            return Err(format!("unknown upscaler model '{model_name}'"));
475        }
476        model_manager::pull_model(state, &model_name, None)
477            .await
478            .map_err(|e| format!("failed to pull upscaler model: {}", e.error))?;
479    }
480
481    let weights_path = {
482        let config = state.config.read().await;
483        config
484            .models
485            .get(&model_name)
486            .and_then(|c| c.transformer.as_ref())
487            .map(std::path::PathBuf::from)
488    }
489    .ok_or_else(|| format!("upscaler model '{model_name}' not configured after pull"))?;
490
491    let upscale_req = mold_core::UpscaleRequest {
492        model: model_name.clone(),
493        image: img.data.clone(),
494        output_format: img.format,
495        tile_size: None,
496        metadata: Some(OutputMetadata::from_generate_request(
497            req,
498            seed_used,
499            None,
500            mold_core::build_info::version_string(),
501        )),
502    };
503    let upscaler_cache = state.upscaler_cache.clone();
504    let progress_tx_for_blocking = progress_tx.cloned();
505    let upscaled =
506        tokio::task::spawn_blocking(move || -> anyhow::Result<mold_core::UpscaleResponse> {
507            let mut cache = upscaler_cache.lock().unwrap_or_else(|e| e.into_inner());
508            let needs_new = cache.as_ref().is_none_or(|e| e.model_name() != model_name);
509            if needs_new {
510                let new_engine = mold_inference::create_upscale_engine(
511                    model_name.clone(),
512                    weights_path,
513                    mold_inference::LoadStrategy::Eager,
514                    0,
515                )?;
516                *cache = Some(new_engine);
517            }
518            let engine = cache.as_mut().unwrap();
519            if let Some(tx) = progress_tx_for_blocking {
520                engine.set_on_progress(Box::new(move |event| {
521                    let _ = tx.send(SseMessage::Progress(progress_to_sse(event)));
522                }));
523            }
524            let result = engine.upscale(&upscale_req);
525            engine.clear_on_progress();
526            result
527        })
528        .await
529        .map_err(|e| format!("upscale task failed: {e}"))?
530        .map_err(|e| format!("upscale failed: {e}"))?;
531
532    let mut response = mold_core::GenerateResponse {
533        images: vec![],
534        video: None,
535        generation_time_ms: 0,
536        model: req.model.clone(),
537        seed_used: req.seed.unwrap_or(0),
538        gpu: None,
539    };
540    apply_upscale_response_to_image_generation(req, &mut response, img, upscaled)
541        .map_err(|e| format!("upscale failed: {e}"))
542}
543
544/// Persist a video's `.preview.gif` sidecar to the server's preview cache
545/// (`$MOLD_HOME/cache/previews/<filename>.preview.gif`). Best-effort —
546/// warnings log and return so a failure here never fails the save path.
547///
548/// Shared with the multi-GPU worker path (`gpu_worker::process_job`) so
549/// video outputs land a preview regardless of which save flow wrote the
550/// MP4; otherwise `/api/gallery/preview/:filename` would 404 whenever the
551/// server is running with GPU workers enabled.
552pub(crate) fn save_video_preview_gif(filename: &str, gif_bytes: &[u8]) {
553    let preview_dir = mold_core::Config::mold_dir()
554        .unwrap_or_else(|| std::path::PathBuf::from(".mold"))
555        .join("cache")
556        .join("previews");
557    save_video_preview_gif_to(&preview_dir, filename, gif_bytes);
558}
559
560/// Testable inner of [`save_video_preview_gif`] that accepts an explicit
561/// preview directory (lets unit tests exercise the write path without
562/// racing on the `MOLD_HOME` env var).
563fn save_video_preview_gif_to(preview_dir: &std::path::Path, filename: &str, gif_bytes: &[u8]) {
564    if let Err(e) = std::fs::create_dir_all(preview_dir) {
565        tracing::warn!(
566            "failed to create preview cache dir {}: {e}",
567            preview_dir.display()
568        );
569        return;
570    }
571    let preview_path = preview_dir.join(mold_core::media_paths::preview_gif_filename(filename));
572    if let Err(e) = std::fs::write(&preview_path, gif_bytes) {
573        tracing::warn!(
574            "failed to write preview gif {}: {e}",
575            preview_path.display()
576        );
577    }
578}
579
580/// Build the SSE `complete` wire event from a finished generation response.
581///
582/// Video responses encode the actual video bytes (MP4/GIF/APNG/WebP) as the
583/// payload and populate every `video_*` metadata field; image responses
584/// encode the image bytes with the video fields cleared. `img` is the
585/// `ImageData` chosen by the caller — either the first generated image or an
586/// `ImageData` synthesized from the video thumbnail (the single-primary-image
587/// shape that the internal `GenerationJobResult` still expects).
588///
589/// Shared between the single-GPU path (`process_job` in this file) and the
590/// multi-GPU path (`gpu_worker::process_job`) so the two can never drift on
591/// which `video_*` fields are populated. Before this helper existed the
592/// multi-GPU worker always encoded the thumbnail PNG as the payload and
593/// hard-coded every `video_*` field to `None`, which silently degraded every
594/// LTX-Video / LTX-2 generation into an image response on hosts with at
595/// least one GPU worker.
596pub(crate) fn build_sse_complete_event(
597    response: &mold_core::GenerateResponse,
598    img: &mold_core::ImageData,
599    original: Option<&mold_core::ImageData>,
600    metadata: Option<&OutputMetadata>,
601    saved: &SavedOutputNames,
602    payload: SseCompletionPayload,
603) -> SseCompleteEvent {
604    let b64 = base64::engine::general_purpose::STANDARD;
605    let include_media = payload == SseCompletionPayload::Full;
606    // Mirror exactly what the save path records: video metadata is used
607    // as-built, image metadata gets the payload's actual dimensions.
608    let event_metadata = metadata.map(|meta| {
609        let mut meta = meta.clone();
610        if response.video.is_none() {
611            apply_output_dimensions_to_metadata(&mut meta, img);
612        }
613        Box::new(meta)
614    });
615    if let Some(ref video) = response.video {
616        SseCompleteEvent {
617            image: if include_media {
618                b64.encode(&video.data)
619            } else {
620                String::new()
621            },
622            format: video.format,
623            width: video.width,
624            height: video.height,
625            original_image: None,
626            original_width: None,
627            original_height: None,
628            seed_used: response.seed_used,
629            generation_time_ms: response.generation_time_ms,
630            model: response.model.clone(),
631            video_frames: Some(video.frames),
632            video_fps: Some(video.fps),
633            video_thumbnail: include_media.then(|| b64.encode(&video.thumbnail)),
634            video_gif_preview: if !include_media || video.gif_preview.is_empty() {
635                None
636            } else {
637                Some(b64.encode(&video.gif_preview))
638            },
639            video_has_audio: video.has_audio,
640            video_duration_ms: video.duration_ms,
641            video_audio_sample_rate: video.audio_sample_rate,
642            video_audio_channels: video.audio_channels,
643            gpu: response.gpu,
644            filename: saved.output.clone(),
645            original_filename: None,
646            metadata: event_metadata,
647        }
648    } else {
649        SseCompleteEvent {
650            image: if include_media {
651                b64.encode(&img.data)
652            } else {
653                String::new()
654            },
655            format: img.format,
656            width: img.width,
657            height: img.height,
658            original_image: include_media
659                .then(|| original.map(|image| b64.encode(&image.data)))
660                .flatten(),
661            original_width: original.map(|image| image.width),
662            original_height: original.map(|image| image.height),
663            seed_used: response.seed_used,
664            generation_time_ms: response.generation_time_ms,
665            model: response.model.clone(),
666            video_frames: None,
667            video_fps: None,
668            video_thumbnail: None,
669            video_gif_preview: None,
670            video_has_audio: false,
671            video_duration_ms: None,
672            video_audio_sample_rate: None,
673            video_audio_channels: None,
674            gpu: response.gpu,
675            filename: saved.output.clone(),
676            original_filename: saved.original.clone(),
677            metadata: event_metadata,
678        }
679    }
680}
681
682pub(crate) fn build_sse_completion_message(
683    response: &mold_core::GenerateResponse,
684    img: &mold_core::ImageData,
685    original: Option<&mold_core::ImageData>,
686    metadata: Option<&OutputMetadata>,
687    saved: &SavedOutputNames,
688    payload: SseCompletionPayload,
689) -> SseMessage {
690    if payload == SseCompletionPayload::MetadataOnly && saved.output.is_none() {
691        return SseMessage::Error(SseErrorEvent {
692            message: "generation completed but the output could not be saved for streaming"
693                .to_string(),
694        });
695    }
696    SseMessage::Complete(Box::new(build_sse_complete_event(
697        response, img, original, metadata, saved, payload,
698    )))
699}
700
701/// Dispatch gate shared through `AppState`, toggled by `POST /api/queue/pause`
702/// and `POST /api/queue/resume`. When paused the dispatch loops stop pulling
703/// *new* jobs off the channel; the job already running on a worker finishes
704/// untouched. Cheap to poll (a single relaxed-ish atomic) so it can sit at the
705/// top of every loop iteration.
706pub struct QueuePause {
707    paused: AtomicBool,
708    /// Wakes every gated dispatch loop on resume. Resume calls
709    /// `notify_waiters()` so *all* loops (single- and multi-GPU) proceed, not
710    /// just one.
711    notify: Notify,
712}
713
714impl QueuePause {
715    pub fn new() -> Arc<Self> {
716        Arc::new(Self {
717            paused: AtomicBool::new(false),
718            notify: Notify::new(),
719        })
720    }
721
722    /// Pause new-job dispatch. Returns `true` iff this call flipped the state
723    /// (was running); idempotent repeat pauses return `false` so the route can
724    /// suppress a duplicate `queue_paused` event.
725    pub fn pause(&self) -> bool {
726        !self.paused.swap(true, Ordering::SeqCst)
727    }
728
729    /// Resume dispatch and wake every gated loop. Returns `true` iff this call
730    /// flipped the state (was paused).
731    pub fn resume(&self) -> bool {
732        let was_paused = self.paused.swap(false, Ordering::SeqCst);
733        if was_paused {
734            self.notify.notify_waiters();
735        }
736        was_paused
737    }
738
739    pub fn is_paused(&self) -> bool {
740        self.paused.load(Ordering::SeqCst)
741    }
742
743    /// Park the caller while paused, returning as soon as dispatch is resumed
744    /// (immediately when not paused). Registers the wakeup *before* the second
745    /// flag check so a concurrent `resume()`'s `notify_waiters()` can't slip
746    /// between the check and the await — the classic lost-wakeup race — and
747    /// re-loops in case of a spurious wake.
748    pub async fn wait_if_paused(&self) {
749        while self.paused.load(Ordering::SeqCst) {
750            let notified = self.notify.notified();
751            tokio::pin!(notified);
752            notified.as_mut().enable();
753            if !self.paused.load(Ordering::SeqCst) {
754                break;
755            }
756            notified.await;
757        }
758    }
759}
760
761/// Runs the generation queue worker loop. Processes one job at a time (FIFO),
762/// but uses a small bounded lookahead buffer to prefer jobs whose model is
763/// already loaded — minimizing model swaps when the queue interleaves models.
764/// Exits when the sender half of the channel is dropped (server shutdown).
765pub async fn run_queue_worker(
766    mut job_rx: tokio::sync::mpsc::Receiver<GenerationJob>,
767    state: AppState,
768) {
769    tracing::debug!("generation queue worker started");
770    let buffer_size = resolve_lookahead_buffer();
771    let max_deferrals = resolve_max_deferrals();
772    let mut buffer: VecDeque<BufferedJob> = VecDeque::with_capacity(buffer_size);
773
774    loop {
775        // Hold new-job dispatch while paused. A job already running finishes
776        // untouched — this only gates the pull of the *next* job.
777        state.queue_pause.wait_if_paused().await;
778        if buffer.is_empty() {
779            match job_rx.recv().await {
780                Some(j) => buffer.push_back(BufferedJob::new(j)),
781                None => break,
782            }
783        }
784        // Top up the buffer without blocking — drain the channel up to capacity.
785        top_up_buffer(&mut buffer, &mut job_rx, buffer_size);
786        // Re-check after the recv: a pause that landed while this loop was
787        // parked waiting for work must hold the job that woke it, not leak
788        // it into dispatch.
789        state.queue_pause.wait_if_paused().await;
790
791        // Honor user reorders (`PATCH /api/queue/:id {position}`) before the
792        // model-swap picker runs — the registry is the single source of truth
793        // for order, so aligning the buffer to it makes a reorder change real
794        // dispatch order rather than only the `GET /api/queue` snapshot.
795        align_buffer_to_registry_order(&mut buffer, &state.job_registry.queued_ids_in_order());
796
797        let loaded = single_gpu_loaded_models(&state).await;
798        let job = pick_next_job(&mut buffer, &loaded, max_deferrals);
799        let job_id = job.id.clone();
800
801        #[cfg(feature = "metrics")]
802        crate::metrics::record_queue_depth(state.queue.pending());
803        process_job(&state, job).await;
804        state.queue.decrement();
805        // Drop the registry entry on every terminal path — the worker
806        // here doesn't own a drop guard, so we do it inline alongside
807        // the queue counter decrement.
808        state.job_registry.remove(&job_id);
809        #[cfg(feature = "metrics")]
810        crate::metrics::record_queue_depth(state.queue.pending());
811    }
812    tracing::info!("generation queue worker shutting down");
813}
814
815async fn single_gpu_loaded_models(state: &AppState) -> std::collections::HashSet<String> {
816    let mut set = std::collections::HashSet::new();
817    let cache = state.model_cache.lock().await;
818    if let Some(name) = cache.active_model() {
819        set.insert(name.to_string());
820    }
821    set
822}
823
824/// Build the set of "currently loaded somewhere" model names across every
825/// worker in the multi-GPU pool. A worker counts the model as loaded if
826/// either it's in the worker's cache as Gpu-resident OR it's the worker's
827/// `active_generation` (covering the take-and-restore window where the
828/// cache entry briefly disappears).
829fn multi_gpu_loaded_models(state: &AppState) -> std::collections::HashSet<String> {
830    let mut set = std::collections::HashSet::new();
831    for worker in &state.gpu_pool.workers {
832        if let Ok(active_gen) = worker.active_generation.read() {
833            if let Some(g) = active_gen.as_ref() {
834                set.insert(g.model.clone());
835            }
836        }
837        if let Ok(cache) = worker.model_cache.lock() {
838            if let Some(name) = cache.active_model() {
839                set.insert(name.to_string());
840            }
841        }
842    }
843    set
844}
845
846/// In-flight wrapper that tracks how many times the picker has skipped this
847/// job. Once the count exceeds `max_deferrals`, the picker force-dispatches
848/// it to bound starvation.
849pub(crate) struct BufferedJob {
850    pub(crate) job: GenerationJob,
851    pub(crate) deferred: usize,
852}
853
854impl BufferedJob {
855    fn new(job: GenerationJob) -> Self {
856        Self { job, deferred: 0 }
857    }
858}
859
860/// Drain the receive channel into the lookahead buffer, capped at
861/// `buffer_size`. Returns when the buffer is full or the channel has no
862/// immediately-available jobs (the receiver is unchanged on `Empty`). Pure
863/// helper extracted so tests can lock in the cap as a load-bearing invariant
864/// without spinning up the full async dispatcher.
865pub(crate) fn top_up_buffer(
866    buffer: &mut VecDeque<BufferedJob>,
867    job_rx: &mut tokio::sync::mpsc::Receiver<GenerationJob>,
868    buffer_size: usize,
869) {
870    while buffer.len() < buffer_size {
871        match job_rx.try_recv() {
872            Ok(j) => buffer.push_back(BufferedJob::new(j)),
873            Err(_) => break,
874        }
875    }
876}
877
878/// Pure picker for the lookahead buffer. Selects the buffered job whose
879/// model is already loaded somewhere in `loaded`; ties broken by arrival
880/// order (front of the deque wins). The head's `deferred` count bounds
881/// starvation: if the head has been skipped `max_deferrals` times, it wins
882/// regardless of `loaded` membership.
883///
884/// The returned job is removed from the buffer; remaining buffered jobs that
885/// were skipped have their `deferred` count incremented. Increments
886/// `mold_queue_reorders_total` whenever a non-head job is picked.
887pub(crate) fn pick_next_job(
888    buffer: &mut VecDeque<BufferedJob>,
889    loaded: &std::collections::HashSet<String>,
890    max_deferrals: usize,
891) -> GenerationJob {
892    debug_assert!(
893        !buffer.is_empty(),
894        "pick_next_job requires non-empty buffer"
895    );
896
897    // Force-dispatch the head if it's hit the starvation budget.
898    if let Some(head) = buffer.pop_front_if(|head| head.deferred >= max_deferrals) {
899        return head.job;
900    }
901
902    // Find the front-most buffered job whose model is already loaded.
903    let pick_idx = buffer
904        .iter()
905        .position(|b| loaded.contains(&b.job.request.model))
906        .unwrap_or(0);
907
908    if pick_idx > 0 {
909        for (i, b) in buffer.iter_mut().enumerate() {
910            if i < pick_idx {
911                b.deferred += 1;
912            }
913        }
914        let model = buffer[pick_idx].job.request.model.clone();
915        tracing::debug!(
916            picked_model = %model,
917            head_model = %buffer.front().map(|b| b.job.request.model.as_str()).unwrap_or(""),
918            picked_index = pick_idx,
919            "queue reorder picked non-head job"
920        );
921        #[cfg(feature = "metrics")]
922        crate::metrics::record_queue_reorder();
923    }
924
925    buffer.remove(pick_idx).expect("pick_idx in range").job
926}
927
928/// Reorder the lookahead `buffer` to follow the registry's queued order — the
929/// single source of truth that `PATCH /api/queue/:id {position}` mutates.
930///
931/// The dispatch loops pull jobs off the channel in submission order into a
932/// bounded buffer and hand it to [`pick_next_job`], which treats the buffer
933/// front as highest priority. Without this step a `reorder_queued` would only
934/// change the `GET /api/queue` snapshot while the loop kept consuming its FIFO
935/// buffer, so a user's "run this next" would be silently ignored. Aligning the
936/// buffer to `order` here makes the reorder drive real dispatch order, while
937/// the model-swap picker still runs on top of the (now user-ordered) sequence,
938/// bounded by the starvation budget.
939///
940/// Jobs are stably reordered by their index in `order`; a job whose id isn't
941/// present — an empty-id test job, one cancelled out of the registry while it
942/// still holds a buffer slot, or one not yet registered — keeps its relative
943/// arrival order behind every registry-tracked job, so nothing untracked is
944/// promoted ahead of tracked work. When the buffer already matches the
945/// registry (the no-reorder steady state) this is a stable no-op that
946/// preserves each job's `deferred` starvation count.
947pub(crate) fn align_buffer_to_registry_order(buffer: &mut VecDeque<BufferedJob>, order: &[String]) {
948    if buffer.len() < 2 {
949        return;
950    }
951    let rank: std::collections::HashMap<&str, usize> = order
952        .iter()
953        .enumerate()
954        .map(|(i, id)| (id.as_str(), i))
955        .collect();
956    // `sort_by_key` is stable, so jobs sharing the fallback rank (unknown ids)
957    // keep their arrival order relative to one another.
958    let mut items: Vec<BufferedJob> = buffer.drain(..).collect();
959    items.sort_by_key(|b| {
960        if b.job.id.is_empty() {
961            usize::MAX
962        } else {
963            rank.get(b.job.id.as_str()).copied().unwrap_or(usize::MAX)
964        }
965    });
966    buffer.extend(items);
967}
968
969pub(crate) const DEFAULT_LOOKAHEAD_BUFFER: usize = 8;
970pub(crate) const DEFAULT_MAX_DEFERRALS: usize = 3;
971pub(crate) const LOOKAHEAD_BUFFER_ENV: &str = "MOLD_QUEUE_LOOKAHEAD_BUFFER";
972pub(crate) const MAX_DEFERRALS_ENV: &str = "MOLD_QUEUE_MAX_DEFERRALS";
973const LOOKAHEAD_BUFFER_LOWER: usize = 1;
974const LOOKAHEAD_BUFFER_UPPER: usize = 64;
975const MAX_DEFERRALS_UPPER: usize = 32;
976
977/// Resolve the lookahead buffer size from env, falling back to the default.
978/// Out-of-range or unparseable values log a warning and use the default —
979/// matching the warn-then-default pattern of `resolve_max_cached_models`.
980pub(crate) fn resolve_lookahead_buffer() -> usize {
981    match std::env::var(LOOKAHEAD_BUFFER_ENV) {
982        Ok(raw) => match raw.trim().parse::<usize>() {
983            Ok(n) if (LOOKAHEAD_BUFFER_LOWER..=LOOKAHEAD_BUFFER_UPPER).contains(&n) => n,
984            Ok(n) => {
985                tracing::warn!(
986                    env = LOOKAHEAD_BUFFER_ENV,
987                    value = n,
988                    lower = LOOKAHEAD_BUFFER_LOWER,
989                    upper = LOOKAHEAD_BUFFER_UPPER,
990                    "ignoring out-of-range queue lookahead buffer; using default"
991                );
992                DEFAULT_LOOKAHEAD_BUFFER
993            }
994            Err(e) => {
995                tracing::warn!(
996                    env = LOOKAHEAD_BUFFER_ENV,
997                    raw = %raw,
998                    error = %e,
999                    "ignoring unparseable queue lookahead buffer; using default"
1000                );
1001                DEFAULT_LOOKAHEAD_BUFFER
1002            }
1003        },
1004        Err(_) => DEFAULT_LOOKAHEAD_BUFFER,
1005    }
1006}
1007
1008/// Resolve the max-deferrals starvation budget from env. Out-of-range or
1009/// unparseable values log a warning and use the default.
1010pub(crate) fn resolve_max_deferrals() -> usize {
1011    match std::env::var(MAX_DEFERRALS_ENV) {
1012        Ok(raw) => match raw.trim().parse::<usize>() {
1013            Ok(n) if n <= MAX_DEFERRALS_UPPER => n,
1014            Ok(n) => {
1015                tracing::warn!(
1016                    env = MAX_DEFERRALS_ENV,
1017                    value = n,
1018                    upper = MAX_DEFERRALS_UPPER,
1019                    "ignoring out-of-range queue max-deferrals; using default"
1020                );
1021                DEFAULT_MAX_DEFERRALS
1022            }
1023            Err(e) => {
1024                tracing::warn!(
1025                    env = MAX_DEFERRALS_ENV,
1026                    raw = %raw,
1027                    error = %e,
1028                    "ignoring unparseable queue max-deferrals; using default"
1029                );
1030                DEFAULT_MAX_DEFERRALS
1031            }
1032        },
1033        Err(_) => DEFAULT_MAX_DEFERRALS,
1034    }
1035}
1036
1037async fn process_job(state: &AppState, job: GenerationJob) {
1038    // Check if client already disconnected before doing any work
1039    if job.result_tx.is_closed() {
1040        tracing::debug!("skipping queued job — client disconnected");
1041        return;
1042    }
1043
1044    // Single-GPU path: there's only one slot. `gpu=None` keeps the wire
1045    // shape consistent with multi-GPU even when we don't know the ordinal.
1046    state.job_registry.mark_running(&job.id, None);
1047
1048    // Send "now processing" event (position 0). `id` echoes the
1049    // server-assigned UUID so reconnecting clients can match progress
1050    // updates to their persisted card.
1051    if let Some(ref tx) = job.progress_tx {
1052        let _ = tx.send(SseMessage::Progress(SseProgressEvent::Queued {
1053            position: 0,
1054            id: job.id.clone(),
1055        }));
1056    }
1057
1058    // 1. Ensure model is ready (with progress forwarding)
1059    let progress_callback = job.progress_tx.as_ref().map(|tx| {
1060        let tx = tx.clone();
1061        Arc::new(move |event: mold_inference::ProgressEvent| {
1062            let _ = tx.send(SseMessage::Progress(progress_to_sse(event)));
1063        }) as model_manager::EngineProgressCallback
1064    });
1065
1066    let activation_hint = model_manager::activation_hint_for_request(state, &job.request).await;
1067    let request_has_lora = model_manager::request_has_effective_lora(&job.request);
1068    if let Err(api_err) = model_manager::ensure_model_ready(
1069        state,
1070        &job.request.model,
1071        progress_callback,
1072        activation_hint,
1073        request_has_lora,
1074    )
1075    .await
1076    {
1077        let err_msg = api_err.error.clone();
1078        if let Some(ref tx) = job.progress_tx {
1079            let _ = tx.send(SseMessage::Error(SseErrorEvent {
1080                message: err_msg.clone(),
1081            }));
1082        }
1083        let _ = job.result_tx.send(Err(err_msg));
1084        return;
1085    }
1086
1087    // 2. Low-memory warning (MPS/unified memory only — observability aid)
1088    #[cfg(target_os = "macos")]
1089    if let Some(available) = mold_inference::device::available_system_memory_bytes() {
1090        if available < 1_000_000_000 {
1091            tracing::warn!(
1092                available_mb = available / 1_000_000,
1093                "low memory before inference — system may become unstable"
1094            );
1095        }
1096    }
1097
1098    // 3. Take the engine out of the cache so the cache mutex stays free during
1099    //    generation. Mirrors the multi-GPU `gpu_worker::process_job` pattern —
1100    //    holding the cache lock through inference would block /api/models,
1101    //    /api/cache, and any concurrent gallery/admin reads.
1102    let taken = {
1103        let mut cache = state.model_cache.lock().await;
1104        cache.take(&job.request.model)
1105    };
1106    let Some(mut cached_engine) = taken else {
1107        let err_msg = "no engine available after model readiness check".to_string();
1108        if let Some(ref tx) = job.progress_tx {
1109            let _ = tx.send(SseMessage::Error(SseErrorEvent {
1110                message: err_msg.clone(),
1111            }));
1112        }
1113        let _ = job.result_tx.send(Err(err_msg));
1114        return;
1115    };
1116
1117    let active_gen = state.active_generation.clone();
1118    let gen_req = job.request.clone();
1119    let progress_tx = job.progress_tx.clone();
1120
1121    set_active_generation(state, &job.request.model, &job.request.prompt);
1122
1123    // Install progress callback before crossing into spawn_blocking — keeps
1124    // the callback installation off the blocking thread. Mirrors the pre-
1125    // refactor behavior: when streaming, set the callback; when not, clear
1126    // it (and only clear after generate when streaming).
1127    let was_streaming = progress_tx.is_some();
1128    if let Some(ref ptx) = progress_tx {
1129        let ptx = ptx.clone();
1130        cached_engine.engine.set_on_progress(Box::new(move |event| {
1131            let _ = ptx.send(SseMessage::Progress(progress_to_sse(event)));
1132        }));
1133    } else {
1134        cached_engine.engine.clear_on_progress();
1135    }
1136
1137    #[cfg(feature = "metrics")]
1138    let inference_start = Instant::now();
1139    // RSS sample taken just before inference; the post-inference sample below
1140    // logs the per-job delta so RAM growth can be attributed to a specific
1141    // generation rather than tracked at process granularity.
1142    let rss_before = crate::resources::ram_snapshot().used_by_mold;
1143    // Run generation on the blocking pool. Move the engine in, return it back
1144    // out (alongside the result + any panic payload) so we can restore it to
1145    // the cache in async context regardless of outcome.
1146    let join_result = tokio::task::spawn_blocking(move || {
1147        let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
1148            cached_engine.engine.generate(&gen_req)
1149        }));
1150        if was_streaming {
1151            cached_engine.engine.clear_on_progress();
1152        }
1153        (cached_engine, result)
1154    })
1155    .await;
1156
1157    let rss_after = crate::resources::ram_snapshot().used_by_mold;
1158    let rss_delta = rss_after as i64 - rss_before as i64;
1159    tracing::info!(
1160        model = %job.request.model,
1161        rss_before_mb = rss_before / 1_000_000,
1162        rss_after_mb = rss_after / 1_000_000,
1163        rss_delta_mb = rss_delta / 1_000_000,
1164        "generation memory delta"
1165    );
1166
1167    #[cfg(feature = "metrics")]
1168    let inference_duration = inference_start.elapsed().as_secs_f64();
1169
1170    // Restore the engine to the cache as soon as the blocking task joins —
1171    // even panics must restore so the cache isn't left with a hole. If the
1172    // tokio task itself failed (JoinError), the engine is gone — restoration
1173    // is impossible. Without `clear_in_flight` the model name would leak
1174    // forever in `in_flight`, so `ensure_model_ready` keeps fast-pathing
1175    // through `contains()` while every subsequent `take()` returns `None`,
1176    // permanently jamming this model. Clear the marker so the cache will
1177    // legitimately re-load the engine on the next request.
1178    let result = match join_result {
1179        Ok((cached_engine, panic_or_result)) => {
1180            {
1181                let mut cache = state.model_cache.lock().await;
1182                cache.restore(cached_engine);
1183            }
1184            clear_active_generation(state);
1185            Ok(panic_or_result)
1186        }
1187        Err(join_err) => {
1188            {
1189                let mut cache = state.model_cache.lock().await;
1190                cache.clear_in_flight(&job.request.model);
1191            }
1192            clear_active_generation(state);
1193            Err(join_err)
1194        }
1195    };
1196
1197    match result {
1198        Ok(Ok(Ok(mut response))) => {
1199            #[cfg(feature = "metrics")]
1200            crate::metrics::record_generation(&job.request.model, inference_duration);
1201
1202            if response.images.is_empty() && response.video.is_none() {
1203                let err_msg = "generation error: engine returned no images or video".to_string();
1204                if let Some(ref tx) = job.progress_tx {
1205                    let _ = tx.send(SseMessage::Error(SseErrorEvent {
1206                        message: err_msg.clone(),
1207                    }));
1208                }
1209                let _ = job.result_tx.send(Err(err_msg));
1210                return;
1211            }
1212            // For video-only responses, synthesize an ImageData from the thumbnail
1213            // so the existing queue/SSE pipeline can handle it.
1214            let mut img = if !response.images.is_empty() {
1215                response.images.remove(0)
1216            } else if let Some(ref video) = response.video {
1217                ImageData {
1218                    data: video.thumbnail.clone(),
1219                    format: OutputFormat::Png,
1220                    width: video.width,
1221                    height: video.height,
1222                    index: 0,
1223                }
1224            } else {
1225                unreachable!("checked above");
1226            };
1227            let mut original_img = None;
1228            if response.video.is_none() && requested_post_upscale_model(&job.request).is_some() {
1229                let upscale_result = upscale_generated_image_on_single_worker(
1230                    state,
1231                    &job.request,
1232                    response.seed_used,
1233                    img.clone(),
1234                    job.progress_tx.as_ref(),
1235                )
1236                .await;
1237                let (output, preserved_original, upscale_error) =
1238                    settle_post_generation_upscale(img, upscale_result);
1239                img = output;
1240                original_img = preserved_original;
1241                if let Some(error) = upscale_error {
1242                    tracing::warn!(%error, "post-generation upscale failed; keeping original image");
1243                }
1244            }
1245
1246            // Save to output directory if configured.
1247            // Builds OutputMetadata from the request + the engine's actual
1248            // seed_used so the DB and embedded chunks agree. Awaited (still
1249            // off the async loop via spawn_blocking) so the complete event
1250            // below can carry the saved gallery filenames.
1251            let metadata = OutputMetadata::from_generate_request(
1252                &job.request,
1253                response.seed_used,
1254                None,
1255                mold_core::build_info::version_string(),
1256            );
1257            let mut saved_names = SavedOutputNames::default();
1258            if let Some(ref dir) = job.output_dir {
1259                let dir = dir.clone();
1260                let model = job.request.model.clone();
1261                let batch_size = job.request.batch_size;
1262                let generation_time_ms = response.generation_time_ms as i64;
1263                let db = state.metadata_db.clone();
1264                let events = state.events.clone();
1265                let save_task = if let Some(ref video) = response.video {
1266                    let video_data = video.data.clone();
1267                    let video_gif_preview = video.gif_preview.clone();
1268                    let video_format = video.format;
1269                    let video_metadata = metadata.clone();
1270                    tokio::task::spawn_blocking(move || SavedOutputNames {
1271                        output: save_video_to_dir(
1272                            &dir,
1273                            &video_data,
1274                            &video_gif_preview,
1275                            video_format,
1276                            &model,
1277                            &video_metadata,
1278                            Some(generation_time_ms),
1279                            db.as_ref().as_ref(),
1280                            Some(&events),
1281                        ),
1282                        original: None,
1283                    })
1284                } else {
1285                    let img_clone = img.clone();
1286                    let original_clone = original_img.clone();
1287                    let metadata_clone = metadata.clone();
1288                    tokio::task::spawn_blocking(move || {
1289                        save_generated_image_outputs(
1290                            &dir,
1291                            original_clone.as_ref(),
1292                            &img_clone,
1293                            &model,
1294                            batch_size,
1295                            &metadata_clone,
1296                            Some(generation_time_ms),
1297                            db.as_ref().as_ref(),
1298                            Some(&events),
1299                        )
1300                    })
1301                };
1302                saved_names = save_task.await.unwrap_or_default();
1303            }
1304
1305            // Send SSE complete event
1306            if let Some(ref tx) = job.progress_tx {
1307                let message = build_sse_completion_message(
1308                    &response,
1309                    &img,
1310                    original_img.as_ref(),
1311                    Some(&metadata),
1312                    &saved_names,
1313                    job.completion_payload,
1314                );
1315                let _ = tx.send(message);
1316            }
1317
1318            // Send result through oneshot
1319            let _ = job.result_tx.send(Ok(GenerationJobResult {
1320                image: img,
1321                response,
1322            }));
1323        }
1324        Ok(Ok(Err(e))) => {
1325            #[cfg(feature = "metrics")]
1326            crate::metrics::record_generation_error(&job.request.model);
1327
1328            *active_gen.write().unwrap_or_else(|e| e.into_inner()) = None;
1329            tracing::error!("generation error: {e:#}");
1330            let err_msg = format!("generation error: {}", clean_error_message(&e));
1331            if let Some(ref tx) = job.progress_tx {
1332                let _ = tx.send(SseMessage::Error(SseErrorEvent {
1333                    message: err_msg.clone(),
1334                }));
1335            }
1336            let _ = job.result_tx.send(Err(err_msg));
1337        }
1338        Ok(Err(panic_payload)) => {
1339            #[cfg(feature = "metrics")]
1340            crate::metrics::record_generation_error(&job.request.model);
1341
1342            *active_gen.write().unwrap_or_else(|e| e.into_inner()) = None;
1343            let msg = panic_payload
1344                .downcast_ref::<String>()
1345                .map(|s| s.as_str())
1346                .or_else(|| panic_payload.downcast_ref::<&str>().copied())
1347                .unwrap_or("unknown panic");
1348            tracing::error!("inference panicked: {msg}");
1349            let err_msg = format!("inference panicked: {msg}");
1350            if let Some(ref tx) = job.progress_tx {
1351                let _ = tx.send(SseMessage::Error(SseErrorEvent {
1352                    message: err_msg.clone(),
1353                }));
1354            }
1355            let _ = job.result_tx.send(Err(err_msg));
1356        }
1357        Err(join_err) => {
1358            #[cfg(feature = "metrics")]
1359            crate::metrics::record_generation_error(&job.request.model);
1360
1361            *active_gen.write().unwrap_or_else(|e| e.into_inner()) = None;
1362            tracing::error!("inference task join error: {join_err:?}");
1363            let err_msg = "inference task failed".to_string();
1364            if let Some(ref tx) = job.progress_tx {
1365                let _ = tx.send(SseMessage::Error(SseErrorEvent {
1366                    message: err_msg.clone(),
1367                }));
1368            }
1369            let _ = job.result_tx.send(Err(err_msg));
1370        }
1371    }
1372}
1373
1374// ── Multi-GPU queue dispatcher ──────────────────────────────────────────────
1375
1376/// Runs the multi-GPU dispatch loop. Routes each generation job to the best
1377/// GPU worker per `select_worker`'s tier order: model loaded + idle > idle
1378/// empty GPU that fits > model loaded but busy > idle empty GPU that does
1379/// not fit > most-headroom fallback (evict LRU there).
1380/// Uses a small lookahead buffer so an interleaved queue (`[A, B, A, B]`)
1381/// doesn't force a sibling worker to swap models when one already has the
1382/// right one warm.
1383///
1384/// Exits when the sender half of the channel is dropped (server shutdown).
1385pub async fn run_queue_dispatcher(
1386    job_rx: tokio::sync::mpsc::Receiver<GenerationJob>,
1387    state: AppState,
1388) {
1389    tracing::debug!("multi-GPU queue dispatcher started");
1390    let buffer_size = resolve_lookahead_buffer();
1391    let max_deferrals = resolve_max_deferrals();
1392    run_queue_dispatcher_with_tuning(job_rx, state, buffer_size, max_deferrals).await;
1393}
1394
1395async fn run_queue_dispatcher_with_tuning(
1396    mut job_rx: tokio::sync::mpsc::Receiver<GenerationJob>,
1397    state: AppState,
1398    buffer_size: usize,
1399    max_deferrals: usize,
1400) {
1401    let mut buffer: VecDeque<BufferedJob> = VecDeque::with_capacity(buffer_size);
1402
1403    loop {
1404        // Hold new-job dispatch while paused; in-flight worker jobs continue.
1405        state.queue_pause.wait_if_paused().await;
1406        if buffer.is_empty() {
1407            match job_rx.recv().await {
1408                Some(j) => buffer.push_back(BufferedJob::new(j)),
1409                None => break,
1410            }
1411        }
1412        top_up_buffer(&mut buffer, &mut job_rx, buffer_size);
1413        // Re-check after the recv: a pause that landed while this loop was
1414        // parked waiting for work must hold the job that woke it, not leak
1415        // it into dispatch.
1416        state.queue_pause.wait_if_paused().await;
1417
1418        // Honor user reorders (`PATCH /api/queue/:id {position}`) before the
1419        // model-swap picker runs — the registry is the single source of truth
1420        // for order, so aligning the buffer to it makes a reorder change real
1421        // dispatch order rather than only the `GET /api/queue` snapshot.
1422        // Reorder is within-host; per-lane `target_gpu` semantics are applied
1423        // later when a worker is selected for the picked job.
1424        align_buffer_to_registry_order(&mut buffer, &state.job_registry.queued_ids_in_order());
1425
1426        let loaded = multi_gpu_loaded_models(&state);
1427        let job = pick_next_job(&mut buffer, &loaded, max_deferrals);
1428
1429        #[cfg(feature = "metrics")]
1430        crate::metrics::record_queue_depth(state.queue.pending());
1431
1432        let job_id = job.id.clone();
1433        let model_name = job.request.model.clone();
1434        let estimated_vram = estimate_model_vram(&model_name);
1435
1436        if let Some(err_msg) = crate::gpu_pool::model_unschedulable_message(&model_name) {
1437            tracing::warn!(model = %model_name, "{err_msg}");
1438            if let Some(tx) = job.progress_tx {
1439                let _ = tx.send(SseMessage::Error(SseErrorEvent {
1440                    message: err_msg.clone(),
1441                }));
1442            }
1443            let _ = job.result_tx.send(Err(err_msg));
1444            state.queue.decrement();
1445            state.job_registry.remove(&job_id);
1446            #[cfg(feature = "metrics")]
1447            crate::metrics::record_queue_depth(state.queue.pending());
1448            continue;
1449        }
1450
1451        let placement_gpu = match state
1452            .gpu_pool
1453            .resolve_explicit_placement_gpu(job.request.placement.as_ref())
1454        {
1455            Ok(ordinal) => ordinal,
1456            Err(err_msg) => {
1457                tracing::warn!(model = %model_name, "{err_msg}");
1458                if let Some(tx) = job.progress_tx {
1459                    let _ = tx.send(SseMessage::Error(SseErrorEvent {
1460                        message: err_msg.clone(),
1461                    }));
1462                }
1463                let _ = job.result_tx.send(Err(err_msg));
1464                state.queue.decrement();
1465                state.job_registry.remove(&job_id);
1466                #[cfg(feature = "metrics")]
1467                crate::metrics::record_queue_depth(state.queue.pending());
1468                continue;
1469            }
1470        };
1471        let preferred_gpu = state
1472            .job_registry
1473            .target_gpu(&job_id)
1474            .flatten()
1475            .or(placement_gpu);
1476
1477        if job.result_tx.is_closed() {
1478            tracing::debug!(model = %model_name, "skipping queued multi-GPU job — client disconnected");
1479            state.queue.decrement();
1480            state.job_registry.remove(&job_id);
1481            #[cfg(feature = "metrics")]
1482            crate::metrics::record_queue_depth(state.queue.pending());
1483            continue;
1484        }
1485
1486        // Multi-GPU workers are synchronous threads and cannot pull missing
1487        // assets themselves. Resolve a first-use post-generation upscaler at
1488        // this async boundary, on the server/host that accepted the job,
1489        // before handing it to the selected GPU.
1490        if let Err(err_msg) =
1491            ensure_post_upscale_model_downloaded(&state, &job.request, job.progress_tx.as_ref())
1492                .await
1493        {
1494            tracing::warn!(
1495                model = %model_name,
1496                upscaler = ?job.request.upscale_model,
1497                "{err_msg}"
1498            );
1499            if let Some(tx) = job.progress_tx {
1500                let _ = tx.send(SseMessage::Error(SseErrorEvent {
1501                    message: err_msg.clone(),
1502                }));
1503            }
1504            let _ = job.result_tx.send(Err(err_msg));
1505            state.queue.decrement();
1506            state.job_registry.remove(&job_id);
1507            #[cfg(feature = "metrics")]
1508            crate::metrics::record_queue_depth(state.queue.pending());
1509            continue;
1510        }
1511
1512        // Build the GpuJob once; the retry loop moves it between attempts.
1513        let mut gpu_job = Some(GpuJob {
1514            id: job.id.clone(),
1515            model: model_name.clone(),
1516            request: job.request,
1517            completion_payload: job.completion_payload,
1518            progress_tx: job.progress_tx,
1519            result_tx: job.result_tx,
1520            output_dir: job.output_dir,
1521            config: state.config.clone(),
1522            metadata_db: state.metadata_db.clone(),
1523            queue: state.queue.clone(),
1524            registry: state.job_registry.clone(),
1525            events: state.events.clone(),
1526        });
1527
1528        let mut skip: Vec<usize> = if preferred_gpu.is_none() {
1529            let failed = crate::gpu_pool::failed_ordinals_for_model(&model_name);
1530            if failed.len() < state.gpu_pool.worker_count() {
1531                failed
1532            } else {
1533                Vec::new()
1534            }
1535        } else {
1536            Vec::new()
1537        };
1538        let mut dispatched = false;
1539
1540        while !dispatched {
1541            if gpu_job
1542                .as_ref()
1543                .is_some_and(|pending| pending.result_tx.is_closed())
1544            {
1545                tracing::debug!(
1546                    model = %model_name,
1547                    "dropping queued multi-GPU job before dispatch — client disconnected"
1548                );
1549                state.queue.decrement();
1550                state.job_registry.remove(&job_id);
1551                break;
1552            }
1553
1554            let worker = if let Some(ordinal) = preferred_gpu {
1555                state.gpu_pool.worker_by_ordinal(ordinal)
1556            } else {
1557                state
1558                    .gpu_pool
1559                    .select_worker_excluding(&model_name, estimated_vram, &skip)
1560            };
1561
1562            let Some(worker) = worker else {
1563                if preferred_gpu.is_none() && state.gpu_pool.worker_count() > 0 {
1564                    tracing::warn!(
1565                        model = %model_name,
1566                        "all GPU workers are temporarily unavailable; keeping job queued"
1567                    );
1568                    tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1569                    continue;
1570                }
1571                let rejected = gpu_job
1572                    .take()
1573                    .expect("gpu_job retained after failed dispatch");
1574                let err_msg = if state.gpu_pool.worker_count() == 0 {
1575                    format!("no GPU available for model {model_name}")
1576                } else if let Some(ordinal) = preferred_gpu {
1577                    format!("gpu:{ordinal} is not available for model {model_name}")
1578                } else {
1579                    format!("no GPU worker available for model {model_name}")
1580                };
1581                tracing::error!(model = %model_name, "{err_msg}");
1582                if let Some(tx) = rejected.progress_tx {
1583                    let _ = tx.send(SseMessage::Error(SseErrorEvent {
1584                        message: err_msg.clone(),
1585                    }));
1586                }
1587                let _ = rejected.result_tx.send(Err(err_msg));
1588                state.queue.decrement();
1589                state.job_registry.remove(&job_id);
1590                break;
1591            };
1592
1593            // Increment in-flight BEFORE sending to reserve the slot.
1594            worker.in_flight.fetch_add(1, Ordering::SeqCst);
1595            let pending = gpu_job.take().expect("gpu_job present in retry loop");
1596            if preferred_gpu.is_none() {
1597                let _ = state
1598                    .job_registry
1599                    .set_target_gpu(&job_id, Some(worker.gpu.ordinal));
1600            }
1601            match worker.job_tx.try_send(pending) {
1602                Ok(()) => {
1603                    dispatched = true;
1604                }
1605                Err(std::sync::mpsc::TrySendError::Full(j)) => {
1606                    worker.in_flight.fetch_sub(1, Ordering::SeqCst);
1607                    if preferred_gpu.is_none() {
1608                        let _ = state.job_registry.set_target_gpu(&job_id, None);
1609                    }
1610                    gpu_job = Some(j);
1611                    if preferred_gpu.is_none() {
1612                        skip.push(worker.gpu.ordinal);
1613                        if skip.len() >= state.gpu_pool.worker_count().max(1) {
1614                            skip.clear();
1615                            tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1616                        }
1617                    } else {
1618                        tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1619                    }
1620                }
1621                Err(std::sync::mpsc::TrySendError::Disconnected(j)) => {
1622                    worker.in_flight.fetch_sub(1, Ordering::SeqCst);
1623                    if preferred_gpu.is_none() {
1624                        let _ = state.job_registry.set_target_gpu(&job_id, None);
1625                    }
1626                    tracing::warn!(
1627                        gpu = worker.gpu.ordinal,
1628                        "GPU worker disconnected — retrying dispatch"
1629                    );
1630                    gpu_job = Some(j);
1631                    if preferred_gpu.is_none() {
1632                        skip.push(worker.gpu.ordinal);
1633                    } else {
1634                        let rejected = gpu_job.take().expect("gpu_job retained after disconnect");
1635                        let err_msg = format!(
1636                            "gpu:{} disconnected while dispatching model {model_name}",
1637                            worker.gpu.ordinal
1638                        );
1639                        if let Some(tx) = rejected.progress_tx {
1640                            let _ = tx.send(SseMessage::Error(SseErrorEvent {
1641                                message: err_msg.clone(),
1642                            }));
1643                        }
1644                        let _ = rejected.result_tx.send(Err(err_msg));
1645                        state.queue.decrement();
1646                        state.job_registry.remove(&job_id);
1647                        break;
1648                    }
1649                }
1650            }
1651        }
1652        #[cfg(feature = "metrics")]
1653        crate::metrics::record_queue_depth(state.queue.pending());
1654    }
1655    tracing::info!("multi-GPU queue dispatcher shutting down");
1656}
1657
1658/// Rough VRAM estimate for a model (used for placement decisions).
1659pub fn estimate_model_vram(model_name: &str) -> u64 {
1660    // Use a simple heuristic based on model name patterns.
1661    // Quantized models are smaller; BF16/FP16 are larger.
1662    let lower = model_name.to_lowercase();
1663    if lower.contains("flux2")
1664        && lower.contains("9b")
1665        && (lower.contains(":bf16") || lower.contains(":fp16"))
1666    {
1667        32_000_000_000 // Klein-9B BF16 needs a 32GB-class card in practice.
1668    } else if lower.contains(":q4") {
1669        6_000_000_000 // ~6GB
1670    } else if lower.contains(":q8") || lower.contains(":fp8") {
1671        12_000_000_000 // ~12GB
1672    } else if lower.contains(":bf16") || lower.contains(":fp16") {
1673        24_000_000_000 // ~24GB
1674    } else if lower.contains("sd15") || lower.contains("sd1.5") {
1675        4_000_000_000 // ~4GB
1676    } else {
1677        // SDXL (~8GB) and other models default to 8GB.
1678        8_000_000_000
1679    }
1680}
1681
1682#[cfg(test)]
1683mod tests {
1684    use super::*;
1685    use crate::gpu_pool::{GpuPool, GpuWorker};
1686    use crate::model_cache::ModelCache;
1687    use crate::state::QueueHandle;
1688    use mold_core::{GenerateRequest, ImageData, ModelConfig, OutputFormat};
1689    use mold_db::MetadataDb;
1690    use mold_inference::device::DiscoveredGpu;
1691    use mold_inference::shared_pool::SharedPool;
1692    use std::sync::atomic::AtomicUsize;
1693    use std::sync::{Arc, Mutex, RwLock};
1694    use tempfile::TempDir;
1695
1696    /// A `GenerateRequest` with the bare minimum fields populated — enough to
1697    /// hand to `OutputMetadata::from_generate_request` in tests.
1698    fn fake_request(model: &str) -> GenerateRequest {
1699        GenerateRequest {
1700            prompt: "a cat".to_string(),
1701            negative_prompt: None,
1702            model: model.to_string(),
1703            width: 512,
1704            height: 512,
1705            steps: 4,
1706            guidance: 3.5,
1707            seed: Some(7),
1708            batch_size: 1,
1709            output_format: Some(OutputFormat::Png),
1710            embed_metadata: None,
1711            scheduler: None,
1712            cfg_plus: None,
1713            source_image: None,
1714            source_image_name: None,
1715            edit_images: None,
1716            strength: 0.75,
1717            mask_image: None,
1718            control_image: None,
1719            control_model: None,
1720            control_scale: 1.0,
1721            expand: None,
1722            original_prompt: None,
1723            batch_id: None,
1724            batch_index: None,
1725            batch_count: None,
1726            lora: None,
1727            frames: None,
1728            fps: None,
1729            upscale_model: None,
1730            gif_preview: false,
1731            enable_audio: None,
1732            audio_file: None,
1733            audio_file_path: None,
1734            source_video: None,
1735            source_video_path: None,
1736            keyframes: None,
1737            pipeline: None,
1738            loras: None,
1739            retake_range: None,
1740            spatial_upscale: None,
1741            temporal_upscale: None,
1742            placement: None,
1743        }
1744    }
1745
1746    fn fake_image() -> ImageData {
1747        ImageData {
1748            // PNG magic bytes — the helpers don't validate, but this keeps
1749            // the on-disk file from being trivially mistaken for empty.
1750            data: vec![0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A],
1751            format: OutputFormat::Png,
1752            width: 512,
1753            height: 512,
1754            index: 0,
1755        }
1756    }
1757
1758    #[test]
1759    fn multi_gpu_dispatch_identifies_missing_post_upscaler_for_auto_pull() {
1760        let mut req = fake_request("flux-dev:q4");
1761        req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
1762
1763        assert_eq!(
1764            post_upscale_model_to_pull(&mold_core::Config::default(), &req).unwrap(),
1765            Some("real-esrgan-x4plus:fp16".to_string())
1766        );
1767
1768        let tmp = TempDir::new().unwrap();
1769        let weights = tmp.path().join("realesrgan.safetensors");
1770        std::fs::write(&weights, b"test weights").unwrap();
1771        let mut config = mold_core::Config::default();
1772        config.models.insert(
1773            "real-esrgan-x4plus:fp16".to_string(),
1774            ModelConfig {
1775                transformer: Some(weights.display().to_string()),
1776                ..Default::default()
1777            },
1778        );
1779        assert_eq!(post_upscale_model_to_pull(&config, &req).unwrap(), None);
1780
1781        config
1782            .models
1783            .get_mut("real-esrgan-x4plus:fp16")
1784            .unwrap()
1785            .transformer = Some(tmp.path().join("missing.safetensors").display().to_string());
1786        assert_eq!(
1787            post_upscale_model_to_pull(&config, &req).unwrap(),
1788            Some("real-esrgan-x4plus:fp16".to_string()),
1789            "stale config paths should trigger a repair pull"
1790        );
1791    }
1792
1793    fn test_worker(
1794        ordinal: usize,
1795        channel_size: usize,
1796    ) -> (
1797        Arc<GpuWorker>,
1798        std::sync::mpsc::Receiver<crate::gpu_pool::GpuJob>,
1799    ) {
1800        let (job_tx, job_rx) = std::sync::mpsc::sync_channel(channel_size);
1801        let worker = Arc::new(GpuWorker {
1802            gpu: DiscoveredGpu {
1803                ordinal,
1804                name: format!("gpu{ordinal}"),
1805                total_vram_bytes: 24_000_000_000,
1806                free_vram_bytes: 24_000_000_000,
1807            },
1808            model_cache: Arc::new(Mutex::new(ModelCache::new(3))),
1809            active_generation: Arc::new(RwLock::new(None)),
1810            model_load_lock: Arc::new(Mutex::new(())),
1811            shared_pool: Arc::new(Mutex::new(SharedPool::new())),
1812            in_flight: AtomicUsize::new(0),
1813            consecutive_failures: AtomicUsize::new(0),
1814            poisoned: AtomicBool::new(false),
1815            fatal_cuda_error: Arc::new(AtomicBool::new(false)),
1816            fatal_cuda_shutdown: Arc::new(tokio::sync::Notify::new()),
1817            degraded_until: RwLock::new(None),
1818            job_tx,
1819        });
1820        (worker, job_rx)
1821    }
1822
1823    fn empty_test_state(config: mold_core::Config) -> crate::state::AppState {
1824        crate::state::AppState::empty(
1825            config,
1826            QueueHandle::new(tokio::sync::mpsc::channel(1).0),
1827            crate::state::AppState::empty_gpu_pool(),
1828            200,
1829        )
1830    }
1831
1832    #[test]
1833    fn save_image_to_dir_writes_file_and_creates_missing_dir() {
1834        let tmp = TempDir::new().unwrap();
1835        let nested = tmp.path().join("sub/output");
1836        assert!(!nested.exists());
1837
1838        save_image_to_dir(
1839            &nested,
1840            &fake_image(),
1841            "flux-dev:q4",
1842            1,
1843            None,
1844            None,
1845            None,
1846            None,
1847        );
1848
1849        assert!(nested.exists(), "save should mkdir -p");
1850        let entries: Vec<_> = std::fs::read_dir(&nested).unwrap().collect();
1851        assert_eq!(entries.len(), 1);
1852        let name = entries[0].as_ref().unwrap().file_name();
1853        let name_str = name.to_string_lossy();
1854        // Filename uses model-with-colon-replaced-by-dash + ms timestamp + .png.
1855        assert!(name_str.starts_with("mold-flux-dev-q4-"), "{name_str}");
1856        assert!(name_str.ends_with(".png"), "{name_str}");
1857    }
1858
1859    #[test]
1860    fn save_image_to_dir_includes_batch_index_when_batch_size_gt_1() {
1861        let tmp = TempDir::new().unwrap();
1862        let mut img = fake_image();
1863        img.index = 3;
1864        img.format = OutputFormat::Jpeg;
1865        img.data = vec![0xFF, 0xD8, 0xFF, 0xE0]; // JPEG magic
1866
1867        save_image_to_dir(tmp.path(), &img, "sdxl", 4, None, None, None, None);
1868
1869        let entries: Vec<_> = std::fs::read_dir(tmp.path()).unwrap().collect();
1870        let name = entries[0]
1871            .as_ref()
1872            .unwrap()
1873            .file_name()
1874            .to_string_lossy()
1875            .to_string();
1876        assert!(
1877            name.contains("-3.jpeg"),
1878            "expected batch index suffix: {name}"
1879        );
1880    }
1881
1882    #[test]
1883    fn save_image_to_dir_upserts_metadata_row_when_db_provided() {
1884        let tmp = TempDir::new().unwrap();
1885        let db = MetadataDb::open_in_memory().unwrap();
1886        let req = fake_request("flux-dev:q4");
1887        let meta = OutputMetadata::from_generate_request(&req, 42, None, "test-version");
1888
1889        save_image_to_dir(
1890            tmp.path(),
1891            &fake_image(),
1892            "flux-dev:q4",
1893            1,
1894            Some(&meta),
1895            Some(1234),
1896            Some(&db),
1897            None,
1898        );
1899
1900        let rows = db.list(Some(tmp.path())).unwrap();
1901        assert_eq!(rows.len(), 1, "exactly one DB row for the saved file");
1902        let rec = &rows[0];
1903        assert_eq!(rec.metadata.prompt, "a cat");
1904        assert_eq!(rec.metadata.seed, 42);
1905        assert_eq!(rec.metadata.version, "test-version");
1906        assert_eq!(rec.format, OutputFormat::Png);
1907        assert_eq!(rec.generation_time_ms, Some(1234));
1908        // stat_from_disk should have populated the size from the actual file.
1909        assert!(rec.file_size_bytes.unwrap_or(0) > 0);
1910    }
1911
1912    #[test]
1913    fn save_generated_image_outputs_persists_original_and_upscaled_dimensions() {
1914        let tmp = TempDir::new().unwrap();
1915        let db = MetadataDb::open_in_memory().unwrap();
1916        let mut req = fake_request("flux-dev:q4");
1917        req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
1918        let meta = OutputMetadata::from_generate_request(&req, 42, None, "test-version");
1919        let original = fake_image();
1920        let mut upscaled = fake_image();
1921        upscaled.width = 2048;
1922        upscaled.height = 2048;
1923        upscaled.data = vec![4, 5, 6];
1924
1925        save_generated_image_outputs(
1926            tmp.path(),
1927            Some(&original),
1928            &upscaled,
1929            "flux-dev:q4",
1930            1,
1931            &meta,
1932            Some(1234),
1933            Some(&db),
1934            None,
1935        );
1936
1937        let rows = db.list(Some(tmp.path())).unwrap();
1938        assert_eq!(rows.len(), 2);
1939        let original_row = rows
1940            .iter()
1941            .find(|row| row.filename.contains("-original."))
1942            .expect("original row");
1943        let upscaled_row = rows
1944            .iter()
1945            .find(|row| row.filename.contains("-upscaled."))
1946            .expect("upscaled row");
1947        assert_eq!(
1948            (original_row.metadata.width, original_row.metadata.height),
1949            (512, 512)
1950        );
1951        assert_eq!(
1952            (upscaled_row.metadata.width, upscaled_row.metadata.height),
1953            (2048, 2048)
1954        );
1955        assert_eq!(upscaled_row.metadata.generation_width, Some(512));
1956        assert_eq!(upscaled_row.metadata.generation_height, Some(512));
1957    }
1958
1959    #[test]
1960    fn save_image_to_dir_skips_db_when_metadata_is_none() {
1961        let tmp = TempDir::new().unwrap();
1962        let db = MetadataDb::open_in_memory().unwrap();
1963
1964        save_image_to_dir(
1965            tmp.path(),
1966            &fake_image(),
1967            "flux-dev:q4",
1968            1,
1969            None, // ← metadata absent
1970            Some(1234),
1971            Some(&db),
1972            None,
1973        );
1974
1975        // File still on disk, but no DB row recorded — both gates must hold
1976        // for the upsert to fire.
1977        assert_eq!(std::fs::read_dir(tmp.path()).unwrap().count(), 1);
1978        assert_eq!(db.list(None).unwrap().len(), 0);
1979    }
1980
1981    #[test]
1982    fn save_image_to_dir_invalid_path_does_not_panic() {
1983        // /dev/null is a file, not a directory — create_dir_all should fail
1984        // and the helper must log + return cleanly rather than panic.
1985        save_image_to_dir(
1986            std::path::Path::new("/dev/null/cant-mkdir-here"),
1987            &fake_image(),
1988            "test",
1989            1,
1990            None,
1991            None,
1992            None,
1993            None,
1994        );
1995    }
1996
1997    #[test]
1998    fn save_image_to_dir_emits_gallery_added_with_row_when_db_records() {
1999        let tmp = TempDir::new().unwrap();
2000        let db = MetadataDb::open_in_memory().unwrap();
2001        let req = fake_request("flux-dev:q4");
2002        let meta = OutputMetadata::from_generate_request(&req, 42, None, "test-version");
2003        let events = crate::events::EventBroadcaster::new();
2004        let mut rx = events.subscribe();
2005
2006        save_image_to_dir(
2007            tmp.path(),
2008            &fake_image(),
2009            "flux-dev:q4",
2010            1,
2011            Some(&meta),
2012            Some(1234),
2013            Some(&db),
2014            Some(&events),
2015        );
2016
2017        match rx.try_recv().unwrap() {
2018            mold_core::ServerEvent::GalleryAdded { filename, image } => {
2019                assert!(filename.ends_with(".png"), "{filename}");
2020                let img = image.expect("DB recorded — event must carry the gallery row");
2021                assert_eq!(img.filename, filename);
2022                assert_eq!(img.metadata.prompt, "a cat");
2023            }
2024            other => panic!("expected gallery_added, got {other:?}"),
2025        }
2026    }
2027
2028    #[test]
2029    fn save_image_to_dir_emits_gallery_added_without_row_when_db_absent() {
2030        let tmp = TempDir::new().unwrap();
2031        let events = crate::events::EventBroadcaster::new();
2032        let mut rx = events.subscribe();
2033
2034        save_image_to_dir(
2035            tmp.path(),
2036            &fake_image(),
2037            "flux-dev:q4",
2038            1,
2039            None,
2040            None,
2041            None, // no DB
2042            Some(&events),
2043        );
2044
2045        match rx.try_recv().unwrap() {
2046            mold_core::ServerEvent::GalleryAdded { image, .. } => {
2047                assert!(image.is_none(), "no DB → clients must refetch");
2048            }
2049            other => panic!("expected gallery_added, got {other:?}"),
2050        }
2051    }
2052
2053    #[test]
2054    fn save_image_to_dir_emits_nothing_on_write_failure() {
2055        let events = crate::events::EventBroadcaster::new();
2056        let mut rx = events.subscribe();
2057
2058        save_image_to_dir(
2059            std::path::Path::new("/dev/null/cant-mkdir-here"),
2060            &fake_image(),
2061            "test",
2062            1,
2063            None,
2064            None,
2065            None,
2066            Some(&events),
2067        );
2068
2069        assert!(
2070            rx.try_recv().is_err(),
2071            "failed save must not announce a gallery entry"
2072        );
2073    }
2074
2075    #[test]
2076    fn save_video_to_dir_emits_gallery_added() {
2077        let tmp = TempDir::new().unwrap();
2078        let db = MetadataDb::open_in_memory().unwrap();
2079        let req = fake_request("ltx-video:fp16");
2080        let meta = OutputMetadata::from_generate_request(&req, 1, None, "v");
2081        let events = crate::events::EventBroadcaster::new();
2082        let mut rx = events.subscribe();
2083
2084        save_video_to_dir(
2085            tmp.path(),
2086            b"fake mp4 bytes",
2087            b"",
2088            OutputFormat::Mp4,
2089            "ltx-video:fp16",
2090            &meta,
2091            Some(5000),
2092            Some(&db),
2093            Some(&events),
2094        );
2095
2096        match rx.try_recv().unwrap() {
2097            mold_core::ServerEvent::GalleryAdded { filename, image } => {
2098                assert!(filename.ends_with(".mp4"), "{filename}");
2099                assert!(image.is_some());
2100            }
2101            other => panic!("expected gallery_added, got {other:?}"),
2102        }
2103    }
2104
2105    #[test]
2106    fn save_video_to_dir_writes_mp4_and_records_metadata() {
2107        let tmp = TempDir::new().unwrap();
2108        let db = MetadataDb::open_in_memory().unwrap();
2109        let mut req = fake_request("ltx-video:fp16");
2110        req.frames = Some(25);
2111        req.fps = Some(24);
2112        let meta = OutputMetadata::from_generate_request(&req, 99, None, "test-version");
2113
2114        // Minimal MP4-ish bytes: an `ftyp` box header. The helper writes
2115        // bytes verbatim — content validation happens at gallery scan time.
2116        let bytes = b"\x00\x00\x00\x18ftypmp42\x00\x00\x00\x00mp42isom".to_vec();
2117
2118        save_video_to_dir(
2119            tmp.path(),
2120            &bytes,
2121            b"",
2122            OutputFormat::Mp4,
2123            "ltx-video:fp16",
2124            &meta,
2125            Some(5000),
2126            Some(&db),
2127            None,
2128        );
2129
2130        let entries: Vec<_> = std::fs::read_dir(tmp.path()).unwrap().collect();
2131        assert_eq!(entries.len(), 1);
2132        let name = entries[0]
2133            .as_ref()
2134            .unwrap()
2135            .file_name()
2136            .to_string_lossy()
2137            .to_string();
2138        assert!(name.starts_with("mold-ltx-video-fp16-"), "{name}");
2139        assert!(name.ends_with(".mp4"), "{name}");
2140
2141        let rows = db.list(Some(tmp.path())).unwrap();
2142        assert_eq!(rows.len(), 1);
2143        assert_eq!(rows[0].format, OutputFormat::Mp4);
2144        assert_eq!(rows[0].metadata.frames, Some(25));
2145        assert_eq!(rows[0].metadata.fps, Some(24));
2146        assert_eq!(rows[0].generation_time_ms, Some(5000));
2147    }
2148
2149    #[test]
2150    fn save_video_to_dir_without_db_still_writes_file() {
2151        let tmp = TempDir::new().unwrap();
2152        let req = fake_request("ltx-video:fp16");
2153        let meta = OutputMetadata::from_generate_request(&req, 1, None, "v");
2154
2155        save_video_to_dir(
2156            tmp.path(),
2157            b"fake gif bytes",
2158            b"",
2159            OutputFormat::Gif,
2160            "ltx-video:fp16",
2161            &meta,
2162            None,
2163            None,
2164            None,
2165        );
2166
2167        let entries: Vec<_> = std::fs::read_dir(tmp.path()).unwrap().collect();
2168        assert_eq!(entries.len(), 1);
2169        let name = entries[0]
2170            .as_ref()
2171            .unwrap()
2172            .file_name()
2173            .to_string_lossy()
2174            .to_string();
2175        assert!(name.ends_with(".gif"), "{name}");
2176    }
2177
2178    #[test]
2179    fn save_video_to_dir_invalid_path_does_not_panic() {
2180        let req = fake_request("ltx-video:fp16");
2181        let meta = OutputMetadata::from_generate_request(&req, 1, None, "v");
2182        save_video_to_dir(
2183            std::path::Path::new("/dev/null/nope"),
2184            b"x",
2185            b"",
2186            OutputFormat::Mp4,
2187            "test",
2188            &meta,
2189            None,
2190            None,
2191            None,
2192        );
2193    }
2194
2195    /// `save_video_preview_gif_to` must write to
2196    /// `<preview_dir>/<filename>.preview.gif` — the exact location
2197    /// `GET /api/gallery/preview/:filename` streams from. Without this
2198    /// sidecar the preview endpoint would 404 on every real generation
2199    /// and the TUI detail pane would only ever see the PNG thumbnail
2200    /// fallback.
2201    #[test]
2202    fn save_video_preview_gif_writes_to_preview_cache() {
2203        let td = tempfile::tempdir().unwrap();
2204        let preview_dir = td.path().join("cache").join("previews");
2205
2206        const GIF: &[u8] = b"GIF89a\x01\x00\x01\x00\x00\x00\x00\x3b";
2207        save_video_preview_gif_to(&preview_dir, "ltx2-42.mp4", GIF);
2208
2209        let expected = preview_dir.join("ltx2-42.mp4.preview.gif");
2210        assert!(
2211            expected.is_file(),
2212            "preview gif should land at {}",
2213            expected.display()
2214        );
2215        assert_eq!(std::fs::read(&expected).unwrap(), GIF);
2216    }
2217
2218    #[test]
2219    fn build_sse_complete_event_video_carries_mp4_payload_and_metadata() {
2220        // Regression guard for the multi-GPU bug: if `response.video` is set,
2221        // the SSE complete event must encode the actual video bytes and
2222        // populate every `video_*` field so the client can reconstruct a
2223        // `VideoData`. Before the shared helper, `gpu_worker.rs` encoded the
2224        // thumbnail PNG and hard-coded every `video_*` field to `None`,
2225        // silently degrading every LTX-Video / LTX-2 response to an image.
2226        let video = mold_core::VideoData {
2227            data: vec![0x00, 0x00, 0x00, 0x18, b'f', b't', b'y', b'p'],
2228            format: OutputFormat::Mp4,
2229            width: 768,
2230            height: 512,
2231            frames: 25,
2232            fps: 24,
2233            thumbnail: vec![0x89, 0x50, 0x4E, 0x47],
2234            gif_preview: vec![b'G', b'I', b'F', b'8'],
2235            has_audio: true,
2236            duration_ms: Some(1040),
2237            audio_sample_rate: Some(44100),
2238            audio_channels: Some(2),
2239        };
2240        let resp = mold_core::GenerateResponse {
2241            images: vec![],
2242            video: Some(video.clone()),
2243            generation_time_ms: 1234,
2244            model: "ltx-2-19b-distilled:fp8".to_string(),
2245            seed_used: 7,
2246            gpu: Some(0),
2247        };
2248        // The `img` the caller synthesizes from the video thumbnail — must be
2249        // ignored for the video branch.
2250        let thumb_img = ImageData {
2251            data: video.thumbnail.clone(),
2252            format: OutputFormat::Png,
2253            width: video.width,
2254            height: video.height,
2255            index: 0,
2256        };
2257
2258        let event = build_sse_complete_event(
2259            &resp,
2260            &thumb_img,
2261            None,
2262            None,
2263            &SavedOutputNames::default(),
2264            SseCompletionPayload::Full,
2265        );
2266
2267        let b64 = base64::engine::general_purpose::STANDARD;
2268        assert_eq!(event.image, b64.encode(&video.data));
2269        assert_eq!(event.format, OutputFormat::Mp4);
2270        assert_eq!(event.video_frames, Some(25));
2271        assert_eq!(event.video_fps, Some(24));
2272        assert_eq!(event.video_thumbnail, Some(b64.encode(&video.thumbnail)));
2273        assert_eq!(
2274            event.video_gif_preview,
2275            Some(b64.encode(&video.gif_preview))
2276        );
2277        assert!(event.video_has_audio);
2278        assert_eq!(event.video_duration_ms, Some(1040));
2279        assert_eq!(event.gpu, Some(0));
2280
2281        let saved = SavedOutputNames {
2282            output: Some("generated-video.mp4".to_string()),
2283            original: None,
2284        };
2285        let metadata_only = build_sse_complete_event(
2286            &resp,
2287            &thumb_img,
2288            None,
2289            None,
2290            &saved,
2291            SseCompletionPayload::MetadataOnly,
2292        );
2293        assert!(metadata_only.image.is_empty());
2294        assert!(metadata_only.video_thumbnail.is_none());
2295        assert!(metadata_only.video_gif_preview.is_none());
2296        assert_eq!(metadata_only.video_frames, Some(25));
2297        assert_eq!(
2298            metadata_only.filename.as_deref(),
2299            Some("generated-video.mp4")
2300        );
2301    }
2302
2303    #[test]
2304    fn build_sse_complete_event_video_empty_gif_preview_omits_field() {
2305        let video = mold_core::VideoData {
2306            data: vec![0x00, 0x00, 0x00, 0x18],
2307            format: OutputFormat::Mp4,
2308            width: 256,
2309            height: 256,
2310            frames: 17,
2311            fps: 12,
2312            thumbnail: vec![0x89, 0x50],
2313            gif_preview: Vec::new(),
2314            has_audio: false,
2315            duration_ms: None,
2316            audio_sample_rate: None,
2317            audio_channels: None,
2318        };
2319        let resp = mold_core::GenerateResponse {
2320            images: vec![],
2321            video: Some(video),
2322            generation_time_ms: 0,
2323            model: "m".to_string(),
2324            seed_used: 0,
2325            gpu: None,
2326        };
2327        let event = build_sse_complete_event(
2328            &resp,
2329            &fake_image(),
2330            None,
2331            None,
2332            &SavedOutputNames::default(),
2333            SseCompletionPayload::Full,
2334        );
2335        assert!(event.video_gif_preview.is_none());
2336        assert!(!event.video_has_audio);
2337    }
2338
2339    #[test]
2340    fn build_sse_complete_event_image_clears_all_video_fields() {
2341        let resp = mold_core::GenerateResponse {
2342            images: vec![fake_image()],
2343            video: None,
2344            generation_time_ms: 100,
2345            model: "flux-schnell:q8".to_string(),
2346            seed_used: 5,
2347            gpu: None,
2348        };
2349        let event = build_sse_complete_event(
2350            &resp,
2351            &fake_image(),
2352            None,
2353            None,
2354            &SavedOutputNames::default(),
2355            SseCompletionPayload::Full,
2356        );
2357        assert_eq!(event.format, OutputFormat::Png);
2358        assert!(event.video_frames.is_none());
2359        assert!(event.video_fps.is_none());
2360        assert!(event.video_thumbnail.is_none());
2361        assert!(event.video_gif_preview.is_none());
2362        assert!(!event.video_has_audio);
2363        assert!(event.video_duration_ms.is_none());
2364    }
2365
2366    #[test]
2367    fn build_sse_complete_event_carries_saved_names_and_recorded_metadata() {
2368        let mut req = fake_request("flux-dev:q4");
2369        req.batch_id = Some("prepared-batch-1".to_string());
2370        req.batch_index = Some(2);
2371        req.batch_count = Some(3);
2372        let resp = mold_core::GenerateResponse {
2373            images: vec![fake_image()],
2374            video: None,
2375            generation_time_ms: 100,
2376            model: "flux-dev:q4".to_string(),
2377            seed_used: 5,
2378            gpu: None,
2379        };
2380        let metadata =
2381            OutputMetadata::from_generate_request(&req, resp.seed_used, None, "test-version");
2382        let saved = SavedOutputNames {
2383            output: Some("flux-dev-q4-123.png".to_string()),
2384            original: Some("flux-dev-q4-123-original.png".to_string()),
2385        };
2386        let event = build_sse_complete_event(
2387            &resp,
2388            &fake_image(),
2389            None,
2390            Some(&metadata),
2391            &saved,
2392            SseCompletionPayload::Full,
2393        );
2394        assert_eq!(event.filename.as_deref(), Some("flux-dev-q4-123.png"));
2395        assert_eq!(
2396            event.original_filename.as_deref(),
2397            Some("flux-dev-q4-123-original.png")
2398        );
2399        // The event metadata mirrors what the save path records: the
2400        // payload's actual dimensions, not the request's.
2401        let meta = event.metadata.expect("metadata rides the complete event");
2402        assert_eq!(meta.seed, 5);
2403        assert_eq!(meta.width, fake_image().width);
2404        assert_eq!(meta.height, fake_image().height);
2405        assert_eq!(meta.batch_id.as_deref(), Some("prepared-batch-1"));
2406        assert_eq!(meta.batch_index, Some(2));
2407        assert_eq!(meta.batch_count, Some(3));
2408
2409        let metadata_only = build_sse_complete_event(
2410            &resp,
2411            &fake_image(),
2412            Some(&fake_image()),
2413            Some(&metadata),
2414            &saved,
2415            SseCompletionPayload::MetadataOnly,
2416        );
2417        assert!(metadata_only.image.is_empty());
2418        assert!(metadata_only.original_image.is_none());
2419        assert_eq!(
2420            metadata_only.filename.as_deref(),
2421            Some("flux-dev-q4-123.png")
2422        );
2423        assert!(metadata_only.metadata.is_some());
2424    }
2425
2426    #[test]
2427    fn metadata_only_completion_fails_when_the_output_was_not_saved() {
2428        let response = mold_core::GenerateResponse {
2429            images: vec![fake_image()],
2430            video: None,
2431            generation_time_ms: 100,
2432            model: "flux-dev:q4".to_string(),
2433            seed_used: 5,
2434            gpu: None,
2435        };
2436        let message = build_sse_completion_message(
2437            &response,
2438            &fake_image(),
2439            None,
2440            None,
2441            &SavedOutputNames::default(),
2442            SseCompletionPayload::MetadataOnly,
2443        );
2444        match message {
2445            SseMessage::Error(error) => assert!(error.message.contains("could not be saved")),
2446            _ => panic!("metadata-only completion without a file must be an SSE error"),
2447        }
2448    }
2449
2450    #[test]
2451    fn post_generation_upscale_replaces_image_response_dimensions() {
2452        let mut req = fake_request("flux-dev:q4");
2453        req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
2454        let mut response = mold_core::GenerateResponse {
2455            images: vec![],
2456            video: None,
2457            generation_time_ms: 100,
2458            model: "flux-dev:q4".to_string(),
2459            seed_used: 5,
2460            gpu: None,
2461        };
2462        let img = fake_image();
2463        let upscaled = mold_core::UpscaleResponse {
2464            image: ImageData {
2465                data: vec![1, 2, 3],
2466                format: OutputFormat::Png,
2467                width: 2048,
2468                height: 2048,
2469                index: 0,
2470            },
2471            upscale_time_ms: 42,
2472            model: "real-esrgan-x4plus:fp16".to_string(),
2473            scale_factor: 4,
2474            original_width: 512,
2475            original_height: 512,
2476        };
2477
2478        let next = apply_upscale_response_to_image_generation(&req, &mut response, img, upscaled)
2479            .expect("image upscale should apply");
2480        let event = build_sse_complete_event(
2481            &response,
2482            &next,
2483            Some(&fake_image()),
2484            None,
2485            &SavedOutputNames::default(),
2486            SseCompletionPayload::Full,
2487        );
2488        assert!(event.original_image.is_some());
2489        assert_eq!(event.original_width, Some(512));
2490        assert_eq!(event.original_height, Some(512));
2491        let mut metadata =
2492            OutputMetadata::from_generate_request(&req, response.seed_used, None, "test-version");
2493        apply_output_dimensions_to_metadata(&mut metadata, &next);
2494
2495        assert_eq!(next.width, 2048);
2496        assert_eq!(next.height, 2048);
2497        assert_eq!(event.width, 2048);
2498        assert_eq!(event.height, 2048);
2499        assert_eq!(metadata.width, 2048);
2500        assert_eq!(metadata.height, 2048);
2501        assert_eq!(metadata.generation_width, Some(512));
2502        assert_eq!(metadata.generation_height, Some(512));
2503        assert_eq!(
2504            metadata.upscale_model.as_deref(),
2505            Some("real-esrgan-x4plus:fp16")
2506        );
2507    }
2508
2509    #[test]
2510    fn failed_post_generation_upscale_keeps_only_the_original_output() {
2511        let original = fake_image();
2512        let (output, preserved_original, error) = settle_post_generation_upscale(
2513            original.clone(),
2514            Err("upscaler unavailable".to_string()),
2515        );
2516
2517        assert_eq!(output.data, original.data);
2518        assert!(preserved_original.is_none());
2519        assert_eq!(error.as_deref(), Some("upscaler unavailable"));
2520    }
2521
2522    #[test]
2523    fn post_generation_upscale_skips_video_responses() {
2524        let mut req = fake_request("ltx-video:fp16");
2525        req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
2526        let video = mold_core::VideoData {
2527            data: vec![0, 0, 0, 24],
2528            format: OutputFormat::Mp4,
2529            width: 512,
2530            height: 512,
2531            frames: 25,
2532            fps: 24,
2533            thumbnail: vec![9, 9],
2534            gif_preview: vec![],
2535            has_audio: false,
2536            duration_ms: None,
2537            audio_sample_rate: None,
2538            audio_channels: None,
2539        };
2540        let mut response = mold_core::GenerateResponse {
2541            images: vec![],
2542            video: Some(video),
2543            generation_time_ms: 100,
2544            model: "ltx-video:fp16".to_string(),
2545            seed_used: 5,
2546            gpu: None,
2547        };
2548        let img = fake_image();
2549        let upscaled = mold_core::UpscaleResponse {
2550            image: ImageData {
2551                data: vec![1, 2, 3],
2552                format: OutputFormat::Png,
2553                width: 2048,
2554                height: 2048,
2555                index: 0,
2556            },
2557            upscale_time_ms: 42,
2558            model: "real-esrgan-x4plus:fp16".to_string(),
2559            scale_factor: 4,
2560            original_width: 512,
2561            original_height: 512,
2562        };
2563
2564        let next = apply_upscale_response_to_image_generation(&req, &mut response, img, upscaled)
2565            .expect("video upscale should be skipped");
2566
2567        assert_eq!(next.width, 512);
2568        assert_eq!(next.height, 512);
2569        assert!(response.video.is_some());
2570    }
2571
2572    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2573    async fn single_worker_post_upscale_noops_without_model() {
2574        let state = empty_test_state(mold_core::Config::default());
2575        let req = fake_request("flux-dev:q4");
2576
2577        let next = upscale_generated_image_on_single_worker(&state, &req, 5, fake_image(), None)
2578            .await
2579            .expect("missing upscale model should leave the image unchanged");
2580
2581        assert_eq!(next.width, 512);
2582        assert_eq!(next.height, 512);
2583        assert_eq!(next.index, 0);
2584    }
2585
2586    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2587    async fn single_worker_post_upscale_rejects_unknown_upscaler_manifest() {
2588        let state = empty_test_state(mold_core::Config::default());
2589        let mut req = fake_request("flux-dev:q4");
2590        req.upscale_model = Some("definitely-not-a-real-upscaler:fp16".to_string());
2591        let (progress_tx, mut progress_rx) = tokio::sync::mpsc::unbounded_channel();
2592
2593        let err = upscale_generated_image_on_single_worker(
2594            &state,
2595            &req,
2596            5,
2597            fake_image(),
2598            Some(&progress_tx),
2599        )
2600        .await
2601        .expect_err("unknown upscalers should fail before generation completes");
2602
2603        assert!(err.contains("unknown upscaler model"), "got: {err}");
2604        let first_progress = progress_rx
2605            .try_recv()
2606            .expect("loading stage should be emitted before validation fails");
2607        assert!(matches!(
2608            first_progress,
2609            SseMessage::Progress(SseProgressEvent::StageStart { .. })
2610        ));
2611    }
2612
2613    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2614    async fn single_worker_post_upscale_surfaces_missing_weights_path() {
2615        let tmp = TempDir::new().unwrap();
2616        let missing_weights = tmp.path().join("missing-upscaler.safetensors");
2617        let mut config = mold_core::Config::default();
2618        config.models.insert(
2619            "real-esrgan-x4plus:fp16".to_string(),
2620            ModelConfig {
2621                transformer: Some(missing_weights.display().to_string()),
2622                ..Default::default()
2623            },
2624        );
2625        let state = empty_test_state(config);
2626        let mut req = fake_request("flux-dev:q4");
2627        req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
2628        let (progress_tx, mut progress_rx) = tokio::sync::mpsc::unbounded_channel();
2629
2630        let err = upscale_generated_image_on_single_worker(
2631            &state,
2632            &req,
2633            5,
2634            fake_image(),
2635            Some(&progress_tx),
2636        )
2637        .await
2638        .expect_err("missing weight files should be surfaced");
2639
2640        assert!(err.contains("upscale failed"), "got: {err}");
2641        assert!(err.contains("upscaler weights not found"), "got: {err}");
2642        let first_progress = progress_rx
2643            .try_recv()
2644            .expect("loading stage should be emitted before loading fails");
2645        assert!(matches!(
2646            first_progress,
2647            SseMessage::Progress(SseProgressEvent::StageStart { .. })
2648        ));
2649    }
2650
2651    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2652    async fn queue_dispatcher_waits_for_worker_capacity_instead_of_rejecting() {
2653        let (worker, worker_rx) = test_worker(0, 1);
2654        let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
2655        let queue = QueueHandle::new(job_tx.clone());
2656        let state = crate::state::AppState::empty(
2657            mold_core::Config::default(),
2658            queue.clone(),
2659            Arc::new(GpuPool {
2660                workers: vec![worker.clone()],
2661            }),
2662            8,
2663        );
2664
2665        let (filler_result_tx, _filler_result_rx) = tokio::sync::oneshot::channel();
2666        let filler_job = crate::gpu_pool::GpuJob {
2667            id: String::new(),
2668            model: "busy-model".to_string(),
2669            request: fake_request("busy-model"),
2670            completion_payload: SseCompletionPayload::Full,
2671            progress_tx: None,
2672            result_tx: filler_result_tx,
2673            output_dir: None,
2674            config: state.config.clone(),
2675            metadata_db: state.metadata_db.clone(),
2676            queue: state.queue.clone(),
2677            registry: state.job_registry.clone(),
2678            events: state.events.clone(),
2679        };
2680        worker.job_tx.send(filler_job).unwrap();
2681
2682        let dispatcher = tokio::spawn(run_queue_dispatcher_with_tuning(
2683            job_rx,
2684            state.clone(),
2685            8,
2686            DEFAULT_MAX_DEFERRALS,
2687        ));
2688
2689        let (result_tx, mut result_rx) = tokio::sync::oneshot::channel();
2690        let job = crate::state::GenerationJob {
2691            id: String::new(),
2692            request: fake_request("flux-dev:q4"),
2693            completion_payload: SseCompletionPayload::Full,
2694            progress_tx: None,
2695            result_tx,
2696            output_dir: None,
2697        };
2698        let _position = queue.submit(job, 8).await.unwrap();
2699
2700        tokio::time::sleep(std::time::Duration::from_millis(25)).await;
2701        assert!(
2702            result_rx.try_recv().is_err(),
2703            "dispatcher should keep the job pending while all worker channels are full"
2704        );
2705
2706        let _filler = worker_rx
2707            .recv()
2708            .expect("filler job should occupy the local channel");
2709        let dispatched = worker_rx
2710            .recv_timeout(std::time::Duration::from_secs(1))
2711            .expect("queued job should dispatch once capacity is available");
2712        assert_eq!(dispatched.model, "flux-dev:q4");
2713
2714        drop(job_tx);
2715        dispatcher.abort();
2716    }
2717
2718    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2719    async fn queue_dispatcher_waits_for_degraded_worker_recovery_instead_of_rejecting() {
2720        let (worker, worker_rx) = test_worker(0, 1);
2721        worker.consecutive_failures.store(3, Ordering::SeqCst);
2722        *worker.degraded_until.write().unwrap() =
2723            Some(Instant::now() + std::time::Duration::from_secs(60));
2724
2725        let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
2726        let queue = QueueHandle::new(job_tx.clone());
2727        let state = crate::state::AppState::empty(
2728            mold_core::Config::default(),
2729            queue.clone(),
2730            Arc::new(GpuPool {
2731                workers: vec![worker.clone()],
2732            }),
2733            8,
2734        );
2735        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
2736
2737        let (result_tx, mut result_rx) = tokio::sync::oneshot::channel();
2738        let job = crate::state::GenerationJob {
2739            id: String::new(),
2740            request: fake_request("flux-dev:q4"),
2741            completion_payload: SseCompletionPayload::Full,
2742            progress_tx: None,
2743            result_tx,
2744            output_dir: None,
2745        };
2746        queue.submit(job, 8).await.unwrap();
2747
2748        tokio::time::sleep(std::time::Duration::from_millis(25)).await;
2749        assert!(
2750            result_rx.try_recv().is_err(),
2751            "dispatcher should keep the job pending while all workers are degraded"
2752        );
2753        assert!(
2754            worker_rx.try_recv().is_err(),
2755            "degraded worker must not receive work before recovery"
2756        );
2757
2758        worker.consecutive_failures.store(0, Ordering::SeqCst);
2759        *worker.degraded_until.write().unwrap() = None;
2760
2761        let dispatched = worker_rx
2762            .recv_timeout(std::time::Duration::from_secs(1))
2763            .expect("queued job should dispatch once a worker recovers");
2764        assert_eq!(dispatched.model, "flux-dev:q4");
2765
2766        drop(job_tx);
2767        dispatcher.abort();
2768    }
2769
2770    /// Regression for the take-and-restore refactor in `process_job`: when
2771    /// the engine vanishes from the cache between `ensure_model_ready` and
2772    /// `cache.take()`, the take path must produce `None` (handled with a
2773    /// clean error in `process_job`) rather than panicking. The pure cache
2774    /// invariant — `take()` on an absent model returns `None` — is what
2775    /// keeps the take-and-restore safe.
2776    #[tokio::test]
2777    async fn cache_take_on_vanished_engine_returns_none_not_panic() {
2778        use crate::model_cache::ModelCache;
2779        use mold_core::GenerateResponse;
2780        use mold_inference::InferenceEngine;
2781
2782        struct StubEngine(&'static str);
2783        impl InferenceEngine for StubEngine {
2784            fn generate(&mut self, _r: &GenerateRequest) -> anyhow::Result<GenerateResponse> {
2785                unimplemented!()
2786            }
2787            fn model_name(&self) -> &str {
2788                self.0
2789            }
2790            fn is_loaded(&self) -> bool {
2791                true
2792            }
2793            fn load(&mut self) -> anyhow::Result<()> {
2794                Ok(())
2795            }
2796        }
2797
2798        let mut cache = ModelCache::new(3);
2799        // Cache empty (engine never inserted, or evicted/removed by a
2800        // concurrent admin call between `ensure_model_ready` and `take`).
2801        assert!(cache.take("vanished-model").is_none());
2802
2803        // After a take of a present engine, a subsequent take of the same
2804        // name must also return None — guards against double-take in the
2805        // restore path.
2806        cache.insert(Box::new(StubEngine("present-model")), 0);
2807        let first = cache.take("present-model");
2808        assert!(first.is_some());
2809        assert!(
2810            cache.take("present-model").is_none(),
2811            "double-take must return None"
2812        );
2813    }
2814
2815    fn buf_job(model: &str) -> BufferedJob {
2816        let (tx, _rx) = tokio::sync::oneshot::channel();
2817        BufferedJob::new(crate::state::GenerationJob {
2818            id: String::new(),
2819            request: fake_request(model),
2820            completion_payload: SseCompletionPayload::Full,
2821            progress_tx: None,
2822            result_tx: tx,
2823            output_dir: None,
2824        })
2825    }
2826
2827    fn buf_job_with_id(id: &str, model: &str) -> BufferedJob {
2828        let (tx, _rx) = tokio::sync::oneshot::channel();
2829        BufferedJob::new(crate::state::GenerationJob {
2830            id: id.to_string(),
2831            request: fake_request(model),
2832            completion_payload: SseCompletionPayload::Full,
2833            progress_tx: None,
2834            result_tx: tx,
2835            output_dir: None,
2836        })
2837    }
2838
2839    #[test]
2840    fn align_buffer_reorders_to_match_registry_queued_order() {
2841        use std::collections::VecDeque;
2842        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2843        for id in ["a", "b", "c"] {
2844            buffer.push_back(buf_job_with_id(id, &format!("model-{id}")));
2845        }
2846        // Registry moved c to the front, then a, then b.
2847        let order = vec!["c".to_string(), "a".to_string(), "b".to_string()];
2848        align_buffer_to_registry_order(&mut buffer, &order);
2849        let ids: Vec<&str> = buffer.iter().map(|b| b.job.id.as_str()).collect();
2850        assert_eq!(ids, vec!["c", "a", "b"]);
2851    }
2852
2853    #[test]
2854    fn align_buffer_is_a_noop_when_already_in_registry_order() {
2855        use std::collections::VecDeque;
2856        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2857        buffer.push_back(buf_job_with_id("a", "model-a"));
2858        // Give the middle job a non-zero deferral count so we can prove the
2859        // no-op path preserves per-job starvation accounting.
2860        let mut b = buf_job_with_id("b", "model-b");
2861        b.deferred = 2;
2862        buffer.push_back(b);
2863        buffer.push_back(buf_job_with_id("c", "model-c"));
2864        let order = vec!["a".to_string(), "b".to_string(), "c".to_string()];
2865        align_buffer_to_registry_order(&mut buffer, &order);
2866        let ids: Vec<&str> = buffer.iter().map(|b| b.job.id.as_str()).collect();
2867        assert_eq!(ids, vec!["a", "b", "c"]);
2868        assert_eq!(
2869            buffer[1].deferred, 2,
2870            "a no-op align must preserve the deferred starvation count"
2871        );
2872    }
2873
2874    #[test]
2875    fn align_buffer_keeps_unregistered_jobs_in_arrival_order_at_the_back() {
2876        use std::collections::VecDeque;
2877        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2878        // "x" and "y" aren't in the registry order (e.g. cancelled out but
2879        // still holding a buffer slot); "b" and "a" are.
2880        for id in ["x", "b", "y", "a"] {
2881            buffer.push_back(buf_job_with_id(id, "m"));
2882        }
2883        let order = vec!["a".to_string(), "b".to_string()];
2884        align_buffer_to_registry_order(&mut buffer, &order);
2885        let ids: Vec<&str> = buffer.iter().map(|b| b.job.id.as_str()).collect();
2886        // Registry-tracked jobs first in registry order (a, b), then the
2887        // untracked ones in their original arrival order (x, y).
2888        assert_eq!(ids, vec!["a", "b", "x", "y"]);
2889    }
2890
2891    #[test]
2892    fn align_buffer_leaves_empty_id_jobs_untouched() {
2893        // Tests that submit `GenerationJob`s directly (empty ids) never
2894        // register in the registry, so the align pass must be a stable no-op
2895        // that preserves the model-swap picker's interleaving assumptions.
2896        use std::collections::VecDeque;
2897        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2898        for model in ["a", "b", "a", "b"] {
2899            buffer.push_back(buf_job(model));
2900        }
2901        align_buffer_to_registry_order(&mut buffer, &[]);
2902        let models: Vec<&str> = buffer
2903            .iter()
2904            .map(|b| b.job.request.model.as_str())
2905            .collect();
2906        assert_eq!(models, vec!["a", "b", "a", "b"]);
2907    }
2908
2909    #[test]
2910    fn pick_next_job_picks_head_when_head_model_loaded() {
2911        use std::collections::{HashSet, VecDeque};
2912        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2913        buffer.push_back(buf_job("a"));
2914        buffer.push_back(buf_job("b"));
2915        buffer.push_back(buf_job("a"));
2916        let loaded: HashSet<String> = ["a".to_string()].into_iter().collect();
2917        let picked = pick_next_job(&mut buffer, &loaded, 3);
2918        assert_eq!(picked.request.model, "a");
2919        assert_eq!(buffer.len(), 2);
2920        assert_eq!(buffer.front().unwrap().job.request.model, "b");
2921        assert_eq!(
2922            buffer.front().unwrap().deferred,
2923            0,
2924            "head shouldn't be deferred when picker chose the head itself"
2925        );
2926    }
2927
2928    #[test]
2929    fn pick_next_job_picks_non_head_when_only_non_head_model_loaded() {
2930        use std::collections::{HashSet, VecDeque};
2931        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2932        buffer.push_back(buf_job("a"));
2933        buffer.push_back(buf_job("b"));
2934        buffer.push_back(buf_job("a"));
2935        let loaded: HashSet<String> = ["b".to_string()].into_iter().collect();
2936        let picked = pick_next_job(&mut buffer, &loaded, 3);
2937        assert_eq!(picked.request.model, "b");
2938        assert_eq!(buffer.len(), 2);
2939        // The head ("a") was skipped once and now sits at deferral=1.
2940        assert_eq!(buffer.front().unwrap().job.request.model, "a");
2941        assert_eq!(buffer.front().unwrap().deferred, 1);
2942    }
2943
2944    #[test]
2945    fn pick_next_job_force_dispatches_head_after_max_deferrals() {
2946        use std::collections::{HashSet, VecDeque};
2947        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2948        let mut head = buf_job("a");
2949        head.deferred = 3;
2950        buffer.push_back(head);
2951        buffer.push_back(buf_job("b"));
2952        // Even though only `b` is loaded, head ("a") has hit the budget and wins.
2953        let loaded: HashSet<String> = ["b".to_string()].into_iter().collect();
2954        let picked = pick_next_job(&mut buffer, &loaded, 3);
2955        assert_eq!(picked.request.model, "a");
2956        assert_eq!(buffer.len(), 1);
2957        assert_eq!(buffer.front().unwrap().job.request.model, "b");
2958    }
2959
2960    #[test]
2961    fn pick_next_job_falls_back_to_head_when_nothing_loaded() {
2962        use std::collections::{HashSet, VecDeque};
2963        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2964        buffer.push_back(buf_job("a"));
2965        buffer.push_back(buf_job("b"));
2966        let loaded: HashSet<String> = HashSet::new();
2967        let picked = pick_next_job(&mut buffer, &loaded, 3);
2968        assert_eq!(picked.request.model, "a");
2969    }
2970
2971    /// Fix D: with `max_deferrals = 0`, every reorder would exceed the
2972    /// budget on the very first skip, so the picker degenerates to FIFO —
2973    /// the head wins regardless of which model is loaded.
2974    #[test]
2975    fn pick_next_job_max_deferrals_zero_picks_head_even_when_non_head_loaded() {
2976        use std::collections::{HashSet, VecDeque};
2977        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2978        buffer.push_back(buf_job("b")); // head
2979        buffer.push_back(buf_job("a")); // non-head
2980        let loaded: HashSet<String> = ["a".to_string()].into_iter().collect();
2981        let picked = pick_next_job(&mut buffer, &loaded, 0);
2982        assert_eq!(
2983            picked.request.model, "b",
2984            "max_deferrals=0 must force FIFO — head must win even when only the non-head model is loaded"
2985        );
2986        assert_eq!(buffer.len(), 1);
2987        assert_eq!(buffer.front().unwrap().job.request.model, "a");
2988    }
2989
2990    /// Fix D: with `max_deferrals = 0` and an empty `loaded` set, the head
2991    /// is the only candidate anyway. Locks in the FIFO behaviour when
2992    /// nothing is warm.
2993    #[test]
2994    fn pick_next_job_max_deferrals_zero_with_empty_loaded_picks_head() {
2995        use std::collections::{HashSet, VecDeque};
2996        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2997        buffer.push_back(buf_job("a")); // head
2998        buffer.push_back(buf_job("b"));
2999        let loaded: HashSet<String> = HashSet::new();
3000        let picked = pick_next_job(&mut buffer, &loaded, 0);
3001        assert_eq!(picked.request.model, "a");
3002        assert_eq!(buffer.len(), 1);
3003        assert_eq!(buffer.front().unwrap().job.request.model, "b");
3004    }
3005
3006    /// Fix E: when both head and a non-head match `loaded`, the picker must
3007    /// pick the front-most match — i.e. the first `A` in `[A, B, A, B]`
3008    /// when both `A` and `B` are loaded. Locks in arrival-order stability
3009    /// across multiple matching jobs.
3010    #[test]
3011    fn pick_next_job_picks_front_most_match_when_multiple_loaded() {
3012        use std::collections::{HashSet, VecDeque};
3013        let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
3014        buffer.push_back(buf_job("a"));
3015        buffer.push_back(buf_job("b"));
3016        buffer.push_back(buf_job("a"));
3017        buffer.push_back(buf_job("b"));
3018        let loaded: HashSet<String> = ["a".to_string(), "b".to_string()].into_iter().collect();
3019        let picked = pick_next_job(&mut buffer, &loaded, 3);
3020        assert_eq!(
3021            picked.request.model, "a",
3022            "front-most match wins (the first `a`), not the loaded model with the most copies later in the buffer"
3023        );
3024        // Three jobs remain: [b, a, b]; head was the picked first `a` so the
3025        // new head is the original-index-1 `b`. Nothing was deferred because
3026        // the picker chose the head itself.
3027        assert_eq!(buffer.len(), 3);
3028        let remaining: Vec<&str> = buffer
3029            .iter()
3030            .map(|b| b.job.request.model.as_str())
3031            .collect();
3032        assert_eq!(remaining, vec!["b", "a", "b"]);
3033        assert_eq!(buffer.front().unwrap().deferred, 0);
3034    }
3035
3036    /// Integration: an interleaved `[A, B, A, B]` queue dispatched against a
3037    /// single worker that has model `A` warm should reorder so both `A` jobs
3038    /// run first, then both `B` jobs — minimizing model swaps from 4 → 1.
3039    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3040    async fn queue_dispatcher_reorders_interleaved_jobs_to_minimize_swaps() {
3041        let (worker, worker_rx) = test_worker(0, 8);
3042        // Pre-mark the worker as having model "a" loaded so the picker
3043        // recognises it as warm.
3044        {
3045            let mut cache = worker.model_cache.lock().unwrap();
3046            struct Engine(&'static str);
3047            impl mold_inference::InferenceEngine for Engine {
3048                fn generate(
3049                    &mut self,
3050                    _r: &GenerateRequest,
3051                ) -> anyhow::Result<mold_core::GenerateResponse> {
3052                    unimplemented!()
3053                }
3054                fn model_name(&self) -> &str {
3055                    self.0
3056                }
3057                fn is_loaded(&self) -> bool {
3058                    true
3059                }
3060                fn load(&mut self) -> anyhow::Result<()> {
3061                    Ok(())
3062                }
3063            }
3064            cache.insert(Box::new(Engine("a")), 0);
3065        }
3066
3067        let (job_tx, job_rx) = tokio::sync::mpsc::channel(8);
3068        let queue = QueueHandle::new(job_tx.clone());
3069        let state = crate::state::AppState::empty(
3070            mold_core::Config::default(),
3071            queue.clone(),
3072            Arc::new(GpuPool {
3073                workers: vec![worker.clone()],
3074            }),
3075            8,
3076        );
3077
3078        // Submit [a, b, a, b] BEFORE the dispatcher spins up so the buffer
3079        // top-up sees all four at once.
3080        let mut result_rxs = Vec::new();
3081        for model in ["a", "b", "a", "b"] {
3082            let (tx, rx) = tokio::sync::oneshot::channel();
3083            let job = crate::state::GenerationJob {
3084                id: String::new(),
3085                request: fake_request(model),
3086                completion_payload: SseCompletionPayload::Full,
3087                progress_tx: None,
3088                result_tx: tx,
3089                output_dir: None,
3090            };
3091            queue.submit(job, 8).await.unwrap();
3092            result_rxs.push(rx);
3093        }
3094
3095        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3096
3097        let mut order = Vec::new();
3098        for _ in 0..4 {
3099            let dispatched = worker_rx
3100                .recv_timeout(std::time::Duration::from_secs(2))
3101                .expect("worker should receive the dispatched job");
3102            order.push(dispatched.model);
3103        }
3104        drop(job_tx);
3105        dispatcher.abort();
3106
3107        assert_eq!(
3108            order,
3109            vec![
3110                "a".to_string(),
3111                "a".to_string(),
3112                "b".to_string(),
3113                "b".to_string(),
3114            ],
3115            "lookahead reorder should batch all `a` jobs together before swapping to `b`"
3116        );
3117    }
3118
3119    /// A `PATCH /api/queue/:id {position}` reorder must change *real* dispatch
3120    /// order, not just the `GET /api/queue` snapshot. Pause the queue so the
3121    /// dispatcher parks before pulling anything, submit A, B, C (all buffered
3122    /// together on resume), move C to the front of the registry, then resume —
3123    /// the worker must receive C first, then A, B. Distinct models with nothing
3124    /// warm keep the model-swap picker on the buffer head, so the registry
3125    /// reorder is the only thing that can change the order.
3126    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3127    async fn queue_dispatcher_honors_registry_reorder_in_real_dispatch() {
3128        let (worker, worker_rx) = test_worker(0, 8);
3129        let (job_tx, job_rx) = tokio::sync::mpsc::channel(8);
3130        let queue = QueueHandle::new(job_tx.clone());
3131        let state = crate::state::AppState::empty(
3132            mold_core::Config::default(),
3133            queue.clone(),
3134            Arc::new(GpuPool {
3135                workers: vec![worker.clone()],
3136            }),
3137            8,
3138        );
3139
3140        // Pause *before* the dispatcher exists so its first `wait_if_paused`
3141        // parks it ahead of the pre-recv gate — nothing is pulled off the
3142        // channel until we resume, so all three jobs buffer together.
3143        state.queue_pause.pause();
3144        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3145
3146        // Register + submit A, B, C with matching ids so the dispatcher's
3147        // registry lookups line up with the queue payloads.
3148        let mut result_rxs = Vec::new();
3149        for id in ["a", "b", "c"] {
3150            state
3151                .job_registry
3152                .register_job(id, format!("model-{id}"), None, None, None);
3153            let (tx, rx) = tokio::sync::oneshot::channel();
3154            let job = crate::state::GenerationJob {
3155                id: id.to_string(),
3156                request: fake_request(&format!("model-{id}")),
3157                completion_payload: SseCompletionPayload::Full,
3158                progress_tx: None,
3159                result_tx: tx,
3160                output_dir: None,
3161            };
3162            queue.submit(job, 8).await.unwrap();
3163            result_rxs.push(rx);
3164        }
3165
3166        // Move C to the front of the registry — the single source of truth for
3167        // dispatch order.
3168        state.job_registry.reorder_queued("c", 0).unwrap();
3169
3170        // Resume: the dispatcher drains A, B, C (still submission-ordered in the
3171        // channel), aligns its buffer to the registry, and dispatches.
3172        state.queue_pause.resume();
3173
3174        let mut order = Vec::new();
3175        for _ in 0..3 {
3176            let dispatched = worker_rx
3177                .recv_timeout(std::time::Duration::from_secs(2))
3178                .expect("worker should receive the dispatched job");
3179            order.push(dispatched.model);
3180        }
3181        drop(job_tx);
3182        dispatcher.abort();
3183
3184        assert_eq!(
3185            order,
3186            vec![
3187                "model-c".to_string(),
3188                "model-a".to_string(),
3189                "model-b".to_string(),
3190            ],
3191            "registry reorder must drive real dispatch order, not just the snapshot"
3192        );
3193    }
3194
3195    /// Fix F: the `top_up_buffer` helper must never grow the buffer past
3196    /// `buffer_size`, no matter how many jobs are sitting in the channel.
3197    /// This is the load-bearing invariant that bounds the working set the
3198    /// picker considers — without it a burst submission could let the
3199    /// dispatcher reorder across the entire pending queue, defeating the
3200    /// fairness guarantees the `deferred` counter is built around.
3201    #[tokio::test]
3202    async fn top_up_buffer_never_exceeds_capacity() {
3203        use std::collections::VecDeque;
3204        let (job_tx, mut job_rx) = tokio::sync::mpsc::channel::<GenerationJob>(32);
3205
3206        // Submit 10 jobs into the channel synchronously so the buffer's top-up
3207        // call sees them all immediately available via try_recv.
3208        for i in 0..10 {
3209            let (tx, _rx) = tokio::sync::oneshot::channel();
3210            let job = GenerationJob {
3211                id: String::new(),
3212                request: fake_request(&format!("model-{i}")),
3213                completion_payload: SseCompletionPayload::Full,
3214                progress_tx: None,
3215                result_tx: tx,
3216                output_dir: None,
3217            };
3218            job_tx.send(job).await.unwrap();
3219        }
3220
3221        // buffer_size = 4 — top_up must stop at 4 even with 10 in the channel.
3222        let mut buffer: VecDeque<BufferedJob> = VecDeque::with_capacity(4);
3223        top_up_buffer(&mut buffer, &mut job_rx, 4);
3224        assert_eq!(
3225            buffer.len(),
3226            4,
3227            "top_up_buffer must cap at buffer_size, leaving the rest in the channel"
3228        );
3229
3230        // Drain the four buffered jobs, then top up again; the next call must
3231        // pull only the next four from the channel (FIFO order preserved).
3232        while buffer.pop_front().is_some() {}
3233        top_up_buffer(&mut buffer, &mut job_rx, 4);
3234        assert_eq!(buffer.len(), 4);
3235        let names: Vec<&str> = buffer
3236            .iter()
3237            .map(|b| b.job.request.model.as_str())
3238            .collect();
3239        assert_eq!(
3240            names,
3241            vec!["model-4", "model-5", "model-6", "model-7"],
3242            "second top-up must drain the next FIFO window from the channel"
3243        );
3244
3245        // Drop sender so the channel reports closed; remaining 2 jobs still
3246        // arrive via try_recv before the channel goes dry.
3247        drop(job_tx);
3248        while buffer.pop_front().is_some() {}
3249        top_up_buffer(&mut buffer, &mut job_rx, 4);
3250        assert_eq!(
3251            buffer.len(),
3252            2,
3253            "top_up_buffer drains the channel tail when fewer jobs than capacity remain"
3254        );
3255        let names: Vec<&str> = buffer
3256            .iter()
3257            .map(|b| b.job.request.model.as_str())
3258            .collect();
3259        assert_eq!(names, vec!["model-8", "model-9"]);
3260    }
3261
3262    /// Same invariant, but reached via the dispatcher loop (integration). A
3263    /// burst of N > buffer_size jobs must still dispatch in FIFO order with
3264    /// no jobs lost — the buffer cap can't drop traffic, only delay it. We
3265    /// drain the worker channel as fast as the dispatcher fills it, so the
3266    /// test exercises buffer rotation rather than worker-channel back-pressure.
3267    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3268    async fn queue_dispatcher_dispatches_all_jobs_when_submission_exceeds_buffer() {
3269        let (worker, worker_rx) = test_worker(0, 4);
3270        let (job_tx, job_rx) = tokio::sync::mpsc::channel(32);
3271        let queue = QueueHandle::new(job_tx.clone());
3272        let state = crate::state::AppState::empty(
3273            mold_core::Config::default(),
3274            queue.clone(),
3275            Arc::new(GpuPool {
3276                workers: vec![worker.clone()],
3277            }),
3278            32,
3279        );
3280
3281        // Drain the worker channel concurrently and decrement in_flight as
3282        // a real worker would, so the dispatcher's worker-selection sees the
3283        // worker as idle for each subsequent send (otherwise `in_flight`
3284        // grows unbounded and the worker never re-classifies as eligible
3285        // when the sync-channel fills).
3286        let drain_worker = worker.clone();
3287        let drainer = std::thread::spawn(move || {
3288            let mut order = Vec::new();
3289            while order.len() < 10 {
3290                match worker_rx.recv_timeout(std::time::Duration::from_secs(5)) {
3291                    Ok(j) => {
3292                        drain_worker.in_flight.fetch_sub(1, Ordering::SeqCst);
3293                        order.push(j.model);
3294                    }
3295                    Err(e) => panic!("drain stalled at {:?}: {e:?}", order),
3296                }
3297            }
3298            order
3299        });
3300
3301        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3302
3303        // Submit AFTER the dispatcher and drainer are running so we exercise
3304        // the live top-up loop rather than a one-shot drain of a pre-filled
3305        // channel. Hold result_rx values past the dispatch — the dispatcher
3306        // skips jobs whose result_tx is closed, which would otherwise drop
3307        // every job before it reaches the worker channel.
3308        let mut held_rxs = Vec::new();
3309        for i in 0..10 {
3310            let (tx, rx) = tokio::sync::oneshot::channel();
3311            held_rxs.push(rx);
3312            let job = crate::state::GenerationJob {
3313                id: String::new(),
3314                request: fake_request(&format!("model-{i}")),
3315                completion_payload: SseCompletionPayload::Full,
3316                progress_tx: None,
3317                result_tx: tx,
3318                output_dir: None,
3319            };
3320            queue.submit(job, 32).await.unwrap();
3321        }
3322
3323        let order = drainer.join().expect("drainer thread panic");
3324        drop(job_tx);
3325        dispatcher.abort();
3326
3327        let expected: Vec<String> = (0..10).map(|i| format!("model-{i}")).collect();
3328        assert_eq!(
3329            order, expected,
3330            "10 distinct jobs must come out in FIFO across buffer rotations"
3331        );
3332    }
3333
3334    /// Serializes every test that mutates queue env vars (process-global).
3335    static QUEUE_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
3336
3337    fn with_queue_env<R>(name: &str, value: Option<&str>, f: impl FnOnce() -> R) -> R {
3338        let _g = QUEUE_ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
3339        let prev = std::env::var(name).ok();
3340        match value {
3341            Some(v) => std::env::set_var(name, v),
3342            None => std::env::remove_var(name),
3343        }
3344        let out = f();
3345        match prev {
3346            Some(v) => std::env::set_var(name, v),
3347            None => std::env::remove_var(name),
3348        }
3349        out
3350    }
3351
3352    #[test]
3353    fn resolve_lookahead_buffer_uses_default_when_env_missing() {
3354        let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, None, resolve_lookahead_buffer);
3355        assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3356    }
3357
3358    #[test]
3359    fn resolve_lookahead_buffer_honors_env_within_range() {
3360        let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, Some("4"), resolve_lookahead_buffer);
3361        assert_eq!(n, 4);
3362    }
3363
3364    #[test]
3365    fn resolve_lookahead_buffer_falls_back_when_out_of_range() {
3366        // 0 is below the 1 lower bound; 999 is above the 64 upper bound.
3367        let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, Some("0"), resolve_lookahead_buffer);
3368        assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3369        let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, Some("999"), resolve_lookahead_buffer);
3370        assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3371    }
3372
3373    #[test]
3374    fn resolve_lookahead_buffer_falls_back_when_unparseable() {
3375        let n = with_queue_env(
3376            LOOKAHEAD_BUFFER_ENV,
3377            Some("not-a-number"),
3378            resolve_lookahead_buffer,
3379        );
3380        assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3381    }
3382
3383    #[test]
3384    fn resolve_max_deferrals_uses_default_when_env_missing() {
3385        let n = with_queue_env(MAX_DEFERRALS_ENV, None, resolve_max_deferrals);
3386        assert_eq!(n, DEFAULT_MAX_DEFERRALS);
3387    }
3388
3389    #[test]
3390    fn resolve_max_deferrals_honors_env_within_range() {
3391        // 0 is the in-range "FIFO" sentinel, 32 is the upper edge.
3392        let n = with_queue_env(MAX_DEFERRALS_ENV, Some("0"), resolve_max_deferrals);
3393        assert_eq!(n, 0);
3394        let n = with_queue_env(MAX_DEFERRALS_ENV, Some("32"), resolve_max_deferrals);
3395        assert_eq!(n, 32);
3396        let n = with_queue_env(MAX_DEFERRALS_ENV, Some("5"), resolve_max_deferrals);
3397        assert_eq!(n, 5);
3398    }
3399
3400    #[test]
3401    fn resolve_max_deferrals_falls_back_when_out_of_range() {
3402        let n = with_queue_env(MAX_DEFERRALS_ENV, Some("999"), resolve_max_deferrals);
3403        assert_eq!(n, DEFAULT_MAX_DEFERRALS);
3404    }
3405
3406    #[test]
3407    fn resolve_max_deferrals_falls_back_when_unparseable() {
3408        let n = with_queue_env(
3409            MAX_DEFERRALS_ENV,
3410            Some("not-a-number"),
3411            resolve_max_deferrals,
3412        );
3413        assert_eq!(n, DEFAULT_MAX_DEFERRALS);
3414    }
3415
3416    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3417    async fn queue_dispatcher_honors_explicit_placement_gpu() {
3418        let (worker0, rx0) = test_worker(0, 1);
3419        let (worker1, rx1) = test_worker(1, 1);
3420        let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3421        let queue = QueueHandle::new(job_tx.clone());
3422        let state = crate::state::AppState::empty(
3423            mold_core::Config::default(),
3424            queue.clone(),
3425            Arc::new(GpuPool {
3426                workers: vec![worker0, worker1],
3427            }),
3428            8,
3429        );
3430
3431        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state));
3432
3433        let mut request = fake_request("flux-dev:q4");
3434        request.placement = Some(mold_core::types::DevicePlacement {
3435            text_encoders: mold_core::types::DeviceRef::Auto,
3436            advanced: Some(mold_core::types::AdvancedPlacement {
3437                transformer: mold_core::types::DeviceRef::gpu(1),
3438                ..mold_core::types::AdvancedPlacement::default()
3439            }),
3440        });
3441
3442        let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3443        let job = crate::state::GenerationJob {
3444            id: String::new(),
3445            request,
3446            completion_payload: SseCompletionPayload::Full,
3447            progress_tx: None,
3448            result_tx,
3449            output_dir: None,
3450        };
3451        let _position = queue.submit(job, 8).await.unwrap();
3452
3453        let dispatched = rx1
3454            .recv_timeout(std::time::Duration::from_secs(1))
3455            .expect("explicit placement should route to gpu 1");
3456        assert_eq!(dispatched.model, "flux-dev:q4");
3457        assert!(rx0.try_recv().is_err(), "gpu 0 should not receive the job");
3458
3459        drop(job_tx);
3460        dispatcher.abort();
3461    }
3462
3463    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3464    async fn queue_dispatcher_records_auto_selected_gpu_before_worker_starts() {
3465        let (worker0, rx0) = test_worker(0, 1);
3466        let (worker1, rx1) = test_worker(1, 1);
3467        let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3468        let queue = QueueHandle::new(job_tx.clone());
3469        let state = crate::state::AppState::empty(
3470            mold_core::Config::default(),
3471            queue.clone(),
3472            Arc::new(GpuPool {
3473                workers: vec![worker0, worker1],
3474            }),
3475            8,
3476        );
3477        state.job_registry.register("auto-job", "flux-dev:q4");
3478
3479        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3480
3481        let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3482        let job = crate::state::GenerationJob {
3483            id: "auto-job".to_string(),
3484            request: fake_request("flux-dev:q4"),
3485            completion_payload: SseCompletionPayload::Full,
3486            progress_tx: None,
3487            result_tx,
3488            output_dir: None,
3489        };
3490        let _position = queue.submit(job, 8).await.unwrap();
3491
3492        let (dispatched, ordinal) = match rx0.recv_timeout(std::time::Duration::from_secs(1)) {
3493            Ok(job) => (job, 0),
3494            Err(_) => (
3495                rx1.recv_timeout(std::time::Duration::from_secs(1))
3496                    .expect("auto job should dispatch to one GPU"),
3497                1,
3498            ),
3499        };
3500        assert_eq!(dispatched.model, "flux-dev:q4");
3501        assert_eq!(
3502            state.job_registry.target_gpu("auto-job"),
3503            Some(Some(ordinal))
3504        );
3505
3506        drop(job_tx);
3507        dispatcher.abort();
3508    }
3509
3510    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3511    async fn paused_dispatcher_holds_new_jobs_until_resumed() {
3512        let (worker0, rx0) = test_worker(0, 1);
3513        let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3514        let queue = QueueHandle::new(job_tx.clone());
3515        let state = crate::state::AppState::empty(
3516            mold_core::Config::default(),
3517            queue.clone(),
3518            Arc::new(GpuPool {
3519                workers: vec![worker0],
3520            }),
3521            8,
3522        );
3523
3524        // Pause before the dispatcher runs — a submitted job must stay queued.
3525        assert!(state.queue_pause.pause());
3526        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3527
3528        let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3529        let job = crate::state::GenerationJob {
3530            id: "paused-job".to_string(),
3531            request: fake_request("flux-dev:q4"),
3532            completion_payload: SseCompletionPayload::Full,
3533            progress_tx: None,
3534            result_tx,
3535            output_dir: None,
3536        };
3537        let _position = queue.submit(job, 8).await.unwrap();
3538
3539        // While paused the worker never receives the job.
3540        assert!(
3541            rx0.recv_timeout(std::time::Duration::from_millis(200))
3542                .is_err(),
3543            "paused dispatcher must not hand a job to a worker"
3544        );
3545
3546        // Resume → the queued job dispatches.
3547        assert!(state.queue_pause.resume());
3548        let dispatched = rx0
3549            .recv_timeout(std::time::Duration::from_secs(1))
3550            .expect("resumed dispatcher should dispatch the queued job");
3551        assert_eq!(dispatched.model, "flux-dev:q4");
3552
3553        drop(job_tx);
3554        dispatcher.abort();
3555    }
3556
3557    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3558    async fn pause_while_dispatcher_is_parked_on_an_empty_queue_still_holds_the_next_job() {
3559        // The subtle ordering: the dispatcher passes the top-of-loop gate,
3560        // then parks in job_rx.recv() on an EMPTY queue. A pause that lands
3561        // while it is parked must hold the very job whose arrival wakes it —
3562        // without the post-recv re-check, that job leaks into dispatch.
3563        let (worker0, rx0) = test_worker(0, 1);
3564        let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3565        let queue = QueueHandle::new(job_tx.clone());
3566        let state = crate::state::AppState::empty(
3567            mold_core::Config::default(),
3568            queue.clone(),
3569            Arc::new(GpuPool {
3570                workers: vec![worker0],
3571            }),
3572            8,
3573        );
3574
3575        // Dispatcher starts UNPAUSED and parks waiting for work.
3576        let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3577        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
3578
3579        // Pause lands while it is parked, then a job arrives.
3580        assert!(state.queue_pause.pause());
3581        let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3582        let job = crate::state::GenerationJob {
3583            id: "parked-job".to_string(),
3584            request: fake_request("flux-dev:q4"),
3585            completion_payload: SseCompletionPayload::Full,
3586            progress_tx: None,
3587            result_tx,
3588            output_dir: None,
3589        };
3590        let _position = queue.submit(job, 8).await.unwrap();
3591
3592        assert!(
3593            rx0.recv_timeout(std::time::Duration::from_millis(200))
3594                .is_err(),
3595            "a job arriving while paused must not wake straight into dispatch"
3596        );
3597
3598        assert!(state.queue_pause.resume());
3599        let dispatched = rx0
3600            .recv_timeout(std::time::Duration::from_secs(1))
3601            .expect("resume should release the held job");
3602        assert_eq!(dispatched.model, "flux-dev:q4");
3603
3604        drop(job_tx);
3605        dispatcher.abort();
3606    }
3607}
3608
3609#[cfg(test)]
3610mod queue_pause_tests {
3611    use super::QueuePause;
3612    use std::time::Duration;
3613
3614    #[test]
3615    fn pause_and_resume_report_state_transitions() {
3616        let gate = QueuePause::new();
3617        assert!(!gate.is_paused());
3618        assert!(gate.pause(), "first pause flips state");
3619        assert!(gate.is_paused());
3620        assert!(!gate.pause(), "second pause is a no-op transition");
3621        assert!(gate.resume(), "first resume flips state");
3622        assert!(!gate.is_paused());
3623        assert!(!gate.resume(), "second resume is a no-op transition");
3624    }
3625
3626    #[tokio::test]
3627    async fn wait_if_paused_returns_immediately_when_not_paused() {
3628        let gate = QueuePause::new();
3629        // Not paused → the await resolves without needing a resume.
3630        tokio::time::timeout(Duration::from_secs(1), gate.wait_if_paused())
3631            .await
3632            .expect("wait_if_paused must not block when the gate is open");
3633    }
3634
3635    #[tokio::test]
3636    async fn wait_if_paused_blocks_until_resumed() {
3637        let gate = QueuePause::new();
3638        assert!(gate.pause());
3639
3640        let waiter = {
3641            let gate = gate.clone();
3642            tokio::spawn(async move { gate.wait_if_paused().await })
3643        };
3644
3645        // While paused the waiter must stay parked.
3646        tokio::time::sleep(Duration::from_millis(50)).await;
3647        assert!(!waiter.is_finished(), "waiter must block while paused");
3648
3649        // Resume wakes it via notify_waiters().
3650        assert!(gate.resume());
3651        tokio::time::timeout(Duration::from_secs(1), waiter)
3652            .await
3653            .expect("waiter must unblock within the timeout after resume")
3654            .expect("waiter task must not panic");
3655    }
3656
3657    #[tokio::test]
3658    async fn resume_wakes_every_gated_waiter() {
3659        // notify_waiters (not notify_one) so all dispatch loops proceed.
3660        let gate = QueuePause::new();
3661        assert!(gate.pause());
3662
3663        let waiters: Vec<_> = (0..3)
3664            .map(|_| {
3665                let gate = gate.clone();
3666                tokio::spawn(async move { gate.wait_if_paused().await })
3667            })
3668            .collect();
3669
3670        tokio::time::sleep(Duration::from_millis(50)).await;
3671        assert!(gate.resume());
3672
3673        for waiter in waiters {
3674            tokio::time::timeout(Duration::from_secs(1), waiter)
3675                .await
3676                .expect("every gated waiter must wake on a single resume")
3677                .expect("waiter task must not panic");
3678        }
3679    }
3680}