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