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
21fn progress_to_sse(event: mold_inference::ProgressEvent) -> SseProgressEvent {
23 event.into()
24}
25
26pub(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#[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 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#[derive(Debug, Default, Clone)]
185pub(crate) struct SavedOutputNames {
186 pub output: Option<String>,
188 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#[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
544pub(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
560fn 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
580pub(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 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
701pub struct QueuePause {
707 paused: AtomicBool,
708 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 pub fn pause(&self) -> bool {
726 !self.paused.swap(true, Ordering::SeqCst)
727 }
728
729 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 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
761pub 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 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_buffer(&mut buffer, &mut job_rx, buffer_size);
786 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 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
818fn 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
840pub(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
854pub(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
872pub(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 if let Some(head) = buffer.pop_front_if(|head| head.deferred >= max_deferrals) {
893 return head.job;
894 }
895
896 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
930pub(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
961pub(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 if job.result_tx.is_closed() {
993 tracing::debug!("skipping queued job — client disconnected");
994 return;
995 }
996
997 state.job_registry.mark_running(&job.id, None);
1000
1001 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 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 #[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 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 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 let rss_before = crate::resources::ram_snapshot().used_by_mold;
1096 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 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 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 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 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 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
1327pub 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 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 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 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 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 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
1603pub fn estimate_model_vram(model_name: &str) -> u64 {
1605 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 } else if lower.contains(":q4") {
1614 6_000_000_000 } else if lower.contains(":q8") || lower.contains(":fp8") {
1616 12_000_000_000 } else if lower.contains(":bf16") || lower.contains(":fp16") {
1618 24_000_000_000 } else if lower.contains("sd15") || lower.contains("sd1.5") {
1620 4_000_000_000 } else {
1622 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 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 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 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]; 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 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, Some(1234),
1916 Some(&db),
1917 None,
1918 );
1919
1920 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 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, 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 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 #[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 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 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 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 #[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 assert!(cache.take("vanished-model").is_none());
2747
2748 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 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 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 #[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")); buffer.push_back(buf_job("a")); 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 #[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")); 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 #[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 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 #[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 {
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 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 #[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 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 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 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(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 #[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 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 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 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 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 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 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 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 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 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 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3364 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
3365
3366 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 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 tokio::time::sleep(Duration::from_millis(50)).await;
3434 assert!(!waiter.is_finished(), "waiter must block while paused");
3435
3436 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 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}