Skip to main content

mold_server/
queue.rs

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