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 align_buffer_to_registry_order(&mut buffer, &state.job_registry.queued_ids_in_order());
796
797 let loaded = single_gpu_loaded_models(&state).await;
798 let job = pick_next_job(&mut buffer, &loaded, max_deferrals);
799 let job_id = job.id.clone();
800
801 #[cfg(feature = "metrics")]
802 crate::metrics::record_queue_depth(state.queue.pending());
803 process_job(&state, job).await;
804 state.queue.decrement();
805 state.job_registry.remove(&job_id);
809 #[cfg(feature = "metrics")]
810 crate::metrics::record_queue_depth(state.queue.pending());
811 }
812 tracing::info!("generation queue worker shutting down");
813}
814
815async fn single_gpu_loaded_models(state: &AppState) -> std::collections::HashSet<String> {
816 let mut set = std::collections::HashSet::new();
817 let cache = state.model_cache.lock().await;
818 if let Some(name) = cache.active_model() {
819 set.insert(name.to_string());
820 }
821 set
822}
823
824fn multi_gpu_loaded_models(state: &AppState) -> std::collections::HashSet<String> {
830 let mut set = std::collections::HashSet::new();
831 for worker in &state.gpu_pool.workers {
832 if let Ok(active_gen) = worker.active_generation.read() {
833 if let Some(g) = active_gen.as_ref() {
834 set.insert(g.model.clone());
835 }
836 }
837 if let Ok(cache) = worker.model_cache.lock() {
838 if let Some(name) = cache.active_model() {
839 set.insert(name.to_string());
840 }
841 }
842 }
843 set
844}
845
846pub(crate) struct BufferedJob {
850 pub(crate) job: GenerationJob,
851 pub(crate) deferred: usize,
852}
853
854impl BufferedJob {
855 fn new(job: GenerationJob) -> Self {
856 Self { job, deferred: 0 }
857 }
858}
859
860pub(crate) fn top_up_buffer(
866 buffer: &mut VecDeque<BufferedJob>,
867 job_rx: &mut tokio::sync::mpsc::Receiver<GenerationJob>,
868 buffer_size: usize,
869) {
870 while buffer.len() < buffer_size {
871 match job_rx.try_recv() {
872 Ok(j) => buffer.push_back(BufferedJob::new(j)),
873 Err(_) => break,
874 }
875 }
876}
877
878pub(crate) fn pick_next_job(
888 buffer: &mut VecDeque<BufferedJob>,
889 loaded: &std::collections::HashSet<String>,
890 max_deferrals: usize,
891) -> GenerationJob {
892 debug_assert!(
893 !buffer.is_empty(),
894 "pick_next_job requires non-empty buffer"
895 );
896
897 if let Some(head) = buffer.pop_front_if(|head| head.deferred >= max_deferrals) {
899 return head.job;
900 }
901
902 let pick_idx = buffer
904 .iter()
905 .position(|b| loaded.contains(&b.job.request.model))
906 .unwrap_or(0);
907
908 if pick_idx > 0 {
909 for (i, b) in buffer.iter_mut().enumerate() {
910 if i < pick_idx {
911 b.deferred += 1;
912 }
913 }
914 let model = buffer[pick_idx].job.request.model.clone();
915 tracing::debug!(
916 picked_model = %model,
917 head_model = %buffer.front().map(|b| b.job.request.model.as_str()).unwrap_or(""),
918 picked_index = pick_idx,
919 "queue reorder picked non-head job"
920 );
921 #[cfg(feature = "metrics")]
922 crate::metrics::record_queue_reorder();
923 }
924
925 buffer.remove(pick_idx).expect("pick_idx in range").job
926}
927
928pub(crate) fn align_buffer_to_registry_order(buffer: &mut VecDeque<BufferedJob>, order: &[String]) {
948 if buffer.len() < 2 {
949 return;
950 }
951 let rank: std::collections::HashMap<&str, usize> = order
952 .iter()
953 .enumerate()
954 .map(|(i, id)| (id.as_str(), i))
955 .collect();
956 let mut items: Vec<BufferedJob> = buffer.drain(..).collect();
959 items.sort_by_key(|b| {
960 if b.job.id.is_empty() {
961 usize::MAX
962 } else {
963 rank.get(b.job.id.as_str()).copied().unwrap_or(usize::MAX)
964 }
965 });
966 buffer.extend(items);
967}
968
969pub(crate) const DEFAULT_LOOKAHEAD_BUFFER: usize = 8;
970pub(crate) const DEFAULT_MAX_DEFERRALS: usize = 3;
971pub(crate) const LOOKAHEAD_BUFFER_ENV: &str = "MOLD_QUEUE_LOOKAHEAD_BUFFER";
972pub(crate) const MAX_DEFERRALS_ENV: &str = "MOLD_QUEUE_MAX_DEFERRALS";
973const LOOKAHEAD_BUFFER_LOWER: usize = 1;
974const LOOKAHEAD_BUFFER_UPPER: usize = 64;
975const MAX_DEFERRALS_UPPER: usize = 32;
976
977pub(crate) fn resolve_lookahead_buffer() -> usize {
981 match std::env::var(LOOKAHEAD_BUFFER_ENV) {
982 Ok(raw) => match raw.trim().parse::<usize>() {
983 Ok(n) if (LOOKAHEAD_BUFFER_LOWER..=LOOKAHEAD_BUFFER_UPPER).contains(&n) => n,
984 Ok(n) => {
985 tracing::warn!(
986 env = LOOKAHEAD_BUFFER_ENV,
987 value = n,
988 lower = LOOKAHEAD_BUFFER_LOWER,
989 upper = LOOKAHEAD_BUFFER_UPPER,
990 "ignoring out-of-range queue lookahead buffer; using default"
991 );
992 DEFAULT_LOOKAHEAD_BUFFER
993 }
994 Err(e) => {
995 tracing::warn!(
996 env = LOOKAHEAD_BUFFER_ENV,
997 raw = %raw,
998 error = %e,
999 "ignoring unparseable queue lookahead buffer; using default"
1000 );
1001 DEFAULT_LOOKAHEAD_BUFFER
1002 }
1003 },
1004 Err(_) => DEFAULT_LOOKAHEAD_BUFFER,
1005 }
1006}
1007
1008pub(crate) fn resolve_max_deferrals() -> usize {
1011 match std::env::var(MAX_DEFERRALS_ENV) {
1012 Ok(raw) => match raw.trim().parse::<usize>() {
1013 Ok(n) if n <= MAX_DEFERRALS_UPPER => n,
1014 Ok(n) => {
1015 tracing::warn!(
1016 env = MAX_DEFERRALS_ENV,
1017 value = n,
1018 upper = MAX_DEFERRALS_UPPER,
1019 "ignoring out-of-range queue max-deferrals; using default"
1020 );
1021 DEFAULT_MAX_DEFERRALS
1022 }
1023 Err(e) => {
1024 tracing::warn!(
1025 env = MAX_DEFERRALS_ENV,
1026 raw = %raw,
1027 error = %e,
1028 "ignoring unparseable queue max-deferrals; using default"
1029 );
1030 DEFAULT_MAX_DEFERRALS
1031 }
1032 },
1033 Err(_) => DEFAULT_MAX_DEFERRALS,
1034 }
1035}
1036
1037async fn process_job(state: &AppState, job: GenerationJob) {
1038 if job.result_tx.is_closed() {
1040 tracing::debug!("skipping queued job — client disconnected");
1041 return;
1042 }
1043
1044 state.job_registry.mark_running(&job.id, None);
1047
1048 if let Some(ref tx) = job.progress_tx {
1052 let _ = tx.send(SseMessage::Progress(SseProgressEvent::Queued {
1053 position: 0,
1054 id: job.id.clone(),
1055 }));
1056 }
1057
1058 let progress_callback = job.progress_tx.as_ref().map(|tx| {
1060 let tx = tx.clone();
1061 Arc::new(move |event: mold_inference::ProgressEvent| {
1062 let _ = tx.send(SseMessage::Progress(progress_to_sse(event)));
1063 }) as model_manager::EngineProgressCallback
1064 });
1065
1066 let activation_hint = model_manager::activation_hint_for_request(state, &job.request).await;
1067 let request_has_lora = model_manager::request_has_effective_lora(&job.request);
1068 if let Err(api_err) = model_manager::ensure_model_ready(
1069 state,
1070 &job.request.model,
1071 progress_callback,
1072 activation_hint,
1073 request_has_lora,
1074 )
1075 .await
1076 {
1077 let err_msg = api_err.error.clone();
1078 if let Some(ref tx) = job.progress_tx {
1079 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1080 message: err_msg.clone(),
1081 }));
1082 }
1083 let _ = job.result_tx.send(Err(err_msg));
1084 return;
1085 }
1086
1087 #[cfg(target_os = "macos")]
1089 if let Some(available) = mold_inference::device::available_system_memory_bytes() {
1090 if available < 1_000_000_000 {
1091 tracing::warn!(
1092 available_mb = available / 1_000_000,
1093 "low memory before inference — system may become unstable"
1094 );
1095 }
1096 }
1097
1098 let taken = {
1103 let mut cache = state.model_cache.lock().await;
1104 cache.take(&job.request.model)
1105 };
1106 let Some(mut cached_engine) = taken else {
1107 let err_msg = "no engine available after model readiness check".to_string();
1108 if let Some(ref tx) = job.progress_tx {
1109 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1110 message: err_msg.clone(),
1111 }));
1112 }
1113 let _ = job.result_tx.send(Err(err_msg));
1114 return;
1115 };
1116
1117 let active_gen = state.active_generation.clone();
1118 let gen_req = job.request.clone();
1119 let progress_tx = job.progress_tx.clone();
1120
1121 set_active_generation(state, &job.request.model, &job.request.prompt);
1122
1123 let was_streaming = progress_tx.is_some();
1128 if let Some(ref ptx) = progress_tx {
1129 let ptx = ptx.clone();
1130 cached_engine.engine.set_on_progress(Box::new(move |event| {
1131 let _ = ptx.send(SseMessage::Progress(progress_to_sse(event)));
1132 }));
1133 } else {
1134 cached_engine.engine.clear_on_progress();
1135 }
1136
1137 #[cfg(feature = "metrics")]
1138 let inference_start = Instant::now();
1139 let rss_before = crate::resources::ram_snapshot().used_by_mold;
1143 let join_result = tokio::task::spawn_blocking(move || {
1147 let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
1148 cached_engine.engine.generate(&gen_req)
1149 }));
1150 if was_streaming {
1151 cached_engine.engine.clear_on_progress();
1152 }
1153 (cached_engine, result)
1154 })
1155 .await;
1156
1157 let rss_after = crate::resources::ram_snapshot().used_by_mold;
1158 let rss_delta = rss_after as i64 - rss_before as i64;
1159 tracing::info!(
1160 model = %job.request.model,
1161 rss_before_mb = rss_before / 1_000_000,
1162 rss_after_mb = rss_after / 1_000_000,
1163 rss_delta_mb = rss_delta / 1_000_000,
1164 "generation memory delta"
1165 );
1166
1167 #[cfg(feature = "metrics")]
1168 let inference_duration = inference_start.elapsed().as_secs_f64();
1169
1170 let result = match join_result {
1179 Ok((cached_engine, panic_or_result)) => {
1180 {
1181 let mut cache = state.model_cache.lock().await;
1182 cache.restore(cached_engine);
1183 }
1184 clear_active_generation(state);
1185 Ok(panic_or_result)
1186 }
1187 Err(join_err) => {
1188 {
1189 let mut cache = state.model_cache.lock().await;
1190 cache.clear_in_flight(&job.request.model);
1191 }
1192 clear_active_generation(state);
1193 Err(join_err)
1194 }
1195 };
1196
1197 match result {
1198 Ok(Ok(Ok(mut response))) => {
1199 #[cfg(feature = "metrics")]
1200 crate::metrics::record_generation(&job.request.model, inference_duration);
1201
1202 if response.images.is_empty() && response.video.is_none() {
1203 let err_msg = "generation error: engine returned no images or video".to_string();
1204 if let Some(ref tx) = job.progress_tx {
1205 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1206 message: err_msg.clone(),
1207 }));
1208 }
1209 let _ = job.result_tx.send(Err(err_msg));
1210 return;
1211 }
1212 let mut img = if !response.images.is_empty() {
1215 response.images.remove(0)
1216 } else if let Some(ref video) = response.video {
1217 ImageData {
1218 data: video.thumbnail.clone(),
1219 format: OutputFormat::Png,
1220 width: video.width,
1221 height: video.height,
1222 index: 0,
1223 }
1224 } else {
1225 unreachable!("checked above");
1226 };
1227 let mut original_img = None;
1228 if response.video.is_none() && requested_post_upscale_model(&job.request).is_some() {
1229 let upscale_result = upscale_generated_image_on_single_worker(
1230 state,
1231 &job.request,
1232 response.seed_used,
1233 img.clone(),
1234 job.progress_tx.as_ref(),
1235 )
1236 .await;
1237 let (output, preserved_original, upscale_error) =
1238 settle_post_generation_upscale(img, upscale_result);
1239 img = output;
1240 original_img = preserved_original;
1241 if let Some(error) = upscale_error {
1242 tracing::warn!(%error, "post-generation upscale failed; keeping original image");
1243 }
1244 }
1245
1246 let metadata = OutputMetadata::from_generate_request(
1252 &job.request,
1253 response.seed_used,
1254 None,
1255 mold_core::build_info::version_string(),
1256 );
1257 let mut saved_names = SavedOutputNames::default();
1258 if let Some(ref dir) = job.output_dir {
1259 let dir = dir.clone();
1260 let model = job.request.model.clone();
1261 let batch_size = job.request.batch_size;
1262 let generation_time_ms = response.generation_time_ms as i64;
1263 let db = state.metadata_db.clone();
1264 let events = state.events.clone();
1265 let save_task = if let Some(ref video) = response.video {
1266 let video_data = video.data.clone();
1267 let video_gif_preview = video.gif_preview.clone();
1268 let video_format = video.format;
1269 let video_metadata = metadata.clone();
1270 tokio::task::spawn_blocking(move || SavedOutputNames {
1271 output: save_video_to_dir(
1272 &dir,
1273 &video_data,
1274 &video_gif_preview,
1275 video_format,
1276 &model,
1277 &video_metadata,
1278 Some(generation_time_ms),
1279 db.as_ref().as_ref(),
1280 Some(&events),
1281 ),
1282 original: None,
1283 })
1284 } else {
1285 let img_clone = img.clone();
1286 let original_clone = original_img.clone();
1287 let metadata_clone = metadata.clone();
1288 tokio::task::spawn_blocking(move || {
1289 save_generated_image_outputs(
1290 &dir,
1291 original_clone.as_ref(),
1292 &img_clone,
1293 &model,
1294 batch_size,
1295 &metadata_clone,
1296 Some(generation_time_ms),
1297 db.as_ref().as_ref(),
1298 Some(&events),
1299 )
1300 })
1301 };
1302 saved_names = save_task.await.unwrap_or_default();
1303 }
1304
1305 if let Some(ref tx) = job.progress_tx {
1307 let message = build_sse_completion_message(
1308 &response,
1309 &img,
1310 original_img.as_ref(),
1311 Some(&metadata),
1312 &saved_names,
1313 job.completion_payload,
1314 );
1315 let _ = tx.send(message);
1316 }
1317
1318 let _ = job.result_tx.send(Ok(GenerationJobResult {
1320 image: img,
1321 response,
1322 }));
1323 }
1324 Ok(Ok(Err(e))) => {
1325 #[cfg(feature = "metrics")]
1326 crate::metrics::record_generation_error(&job.request.model);
1327
1328 *active_gen.write().unwrap_or_else(|e| e.into_inner()) = None;
1329 tracing::error!("generation error: {e:#}");
1330 let err_msg = format!("generation error: {}", clean_error_message(&e));
1331 if let Some(ref tx) = job.progress_tx {
1332 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1333 message: err_msg.clone(),
1334 }));
1335 }
1336 let _ = job.result_tx.send(Err(err_msg));
1337 }
1338 Ok(Err(panic_payload)) => {
1339 #[cfg(feature = "metrics")]
1340 crate::metrics::record_generation_error(&job.request.model);
1341
1342 *active_gen.write().unwrap_or_else(|e| e.into_inner()) = None;
1343 let msg = panic_payload
1344 .downcast_ref::<String>()
1345 .map(|s| s.as_str())
1346 .or_else(|| panic_payload.downcast_ref::<&str>().copied())
1347 .unwrap_or("unknown panic");
1348 tracing::error!("inference panicked: {msg}");
1349 let err_msg = format!("inference panicked: {msg}");
1350 if let Some(ref tx) = job.progress_tx {
1351 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1352 message: err_msg.clone(),
1353 }));
1354 }
1355 let _ = job.result_tx.send(Err(err_msg));
1356 }
1357 Err(join_err) => {
1358 #[cfg(feature = "metrics")]
1359 crate::metrics::record_generation_error(&job.request.model);
1360
1361 *active_gen.write().unwrap_or_else(|e| e.into_inner()) = None;
1362 tracing::error!("inference task join error: {join_err:?}");
1363 let err_msg = "inference task failed".to_string();
1364 if let Some(ref tx) = job.progress_tx {
1365 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1366 message: err_msg.clone(),
1367 }));
1368 }
1369 let _ = job.result_tx.send(Err(err_msg));
1370 }
1371 }
1372}
1373
1374pub async fn run_queue_dispatcher(
1386 job_rx: tokio::sync::mpsc::Receiver<GenerationJob>,
1387 state: AppState,
1388) {
1389 tracing::debug!("multi-GPU queue dispatcher started");
1390 let buffer_size = resolve_lookahead_buffer();
1391 let max_deferrals = resolve_max_deferrals();
1392 run_queue_dispatcher_with_tuning(job_rx, state, buffer_size, max_deferrals).await;
1393}
1394
1395async fn run_queue_dispatcher_with_tuning(
1396 mut job_rx: tokio::sync::mpsc::Receiver<GenerationJob>,
1397 state: AppState,
1398 buffer_size: usize,
1399 max_deferrals: usize,
1400) {
1401 let mut buffer: VecDeque<BufferedJob> = VecDeque::with_capacity(buffer_size);
1402
1403 loop {
1404 state.queue_pause.wait_if_paused().await;
1406 if buffer.is_empty() {
1407 match job_rx.recv().await {
1408 Some(j) => buffer.push_back(BufferedJob::new(j)),
1409 None => break,
1410 }
1411 }
1412 top_up_buffer(&mut buffer, &mut job_rx, buffer_size);
1413 state.queue_pause.wait_if_paused().await;
1417
1418 align_buffer_to_registry_order(&mut buffer, &state.job_registry.queued_ids_in_order());
1425
1426 let loaded = multi_gpu_loaded_models(&state);
1427 let job = pick_next_job(&mut buffer, &loaded, max_deferrals);
1428
1429 #[cfg(feature = "metrics")]
1430 crate::metrics::record_queue_depth(state.queue.pending());
1431
1432 let job_id = job.id.clone();
1433 let model_name = job.request.model.clone();
1434 let estimated_vram = estimate_model_vram(&model_name);
1435
1436 if let Some(err_msg) = crate::gpu_pool::model_unschedulable_message(&model_name) {
1437 tracing::warn!(model = %model_name, "{err_msg}");
1438 if let Some(tx) = job.progress_tx {
1439 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1440 message: err_msg.clone(),
1441 }));
1442 }
1443 let _ = job.result_tx.send(Err(err_msg));
1444 state.queue.decrement();
1445 state.job_registry.remove(&job_id);
1446 #[cfg(feature = "metrics")]
1447 crate::metrics::record_queue_depth(state.queue.pending());
1448 continue;
1449 }
1450
1451 let placement_gpu = match state
1452 .gpu_pool
1453 .resolve_explicit_placement_gpu(job.request.placement.as_ref())
1454 {
1455 Ok(ordinal) => ordinal,
1456 Err(err_msg) => {
1457 tracing::warn!(model = %model_name, "{err_msg}");
1458 if let Some(tx) = job.progress_tx {
1459 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1460 message: err_msg.clone(),
1461 }));
1462 }
1463 let _ = job.result_tx.send(Err(err_msg));
1464 state.queue.decrement();
1465 state.job_registry.remove(&job_id);
1466 #[cfg(feature = "metrics")]
1467 crate::metrics::record_queue_depth(state.queue.pending());
1468 continue;
1469 }
1470 };
1471 let preferred_gpu = state
1472 .job_registry
1473 .target_gpu(&job_id)
1474 .flatten()
1475 .or(placement_gpu);
1476
1477 if job.result_tx.is_closed() {
1478 tracing::debug!(model = %model_name, "skipping queued multi-GPU job — client disconnected");
1479 state.queue.decrement();
1480 state.job_registry.remove(&job_id);
1481 #[cfg(feature = "metrics")]
1482 crate::metrics::record_queue_depth(state.queue.pending());
1483 continue;
1484 }
1485
1486 if let Err(err_msg) =
1491 ensure_post_upscale_model_downloaded(&state, &job.request, job.progress_tx.as_ref())
1492 .await
1493 {
1494 tracing::warn!(
1495 model = %model_name,
1496 upscaler = ?job.request.upscale_model,
1497 "{err_msg}"
1498 );
1499 if let Some(tx) = job.progress_tx {
1500 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1501 message: err_msg.clone(),
1502 }));
1503 }
1504 let _ = job.result_tx.send(Err(err_msg));
1505 state.queue.decrement();
1506 state.job_registry.remove(&job_id);
1507 #[cfg(feature = "metrics")]
1508 crate::metrics::record_queue_depth(state.queue.pending());
1509 continue;
1510 }
1511
1512 let mut gpu_job = Some(GpuJob {
1514 id: job.id.clone(),
1515 model: model_name.clone(),
1516 request: job.request,
1517 completion_payload: job.completion_payload,
1518 progress_tx: job.progress_tx,
1519 result_tx: job.result_tx,
1520 output_dir: job.output_dir,
1521 config: state.config.clone(),
1522 metadata_db: state.metadata_db.clone(),
1523 queue: state.queue.clone(),
1524 registry: state.job_registry.clone(),
1525 events: state.events.clone(),
1526 });
1527
1528 let mut skip: Vec<usize> = if preferred_gpu.is_none() {
1529 let failed = crate::gpu_pool::failed_ordinals_for_model(&model_name);
1530 if failed.len() < state.gpu_pool.worker_count() {
1531 failed
1532 } else {
1533 Vec::new()
1534 }
1535 } else {
1536 Vec::new()
1537 };
1538 let mut dispatched = false;
1539
1540 while !dispatched {
1541 if gpu_job
1542 .as_ref()
1543 .is_some_and(|pending| pending.result_tx.is_closed())
1544 {
1545 tracing::debug!(
1546 model = %model_name,
1547 "dropping queued multi-GPU job before dispatch — client disconnected"
1548 );
1549 state.queue.decrement();
1550 state.job_registry.remove(&job_id);
1551 break;
1552 }
1553
1554 let worker = if let Some(ordinal) = preferred_gpu {
1555 state.gpu_pool.worker_by_ordinal(ordinal)
1556 } else {
1557 state
1558 .gpu_pool
1559 .select_worker_excluding(&model_name, estimated_vram, &skip)
1560 };
1561
1562 let Some(worker) = worker else {
1563 if preferred_gpu.is_none() && state.gpu_pool.worker_count() > 0 {
1564 tracing::warn!(
1565 model = %model_name,
1566 "all GPU workers are temporarily unavailable; keeping job queued"
1567 );
1568 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1569 continue;
1570 }
1571 let rejected = gpu_job
1572 .take()
1573 .expect("gpu_job retained after failed dispatch");
1574 let err_msg = if state.gpu_pool.worker_count() == 0 {
1575 format!("no GPU available for model {model_name}")
1576 } else if let Some(ordinal) = preferred_gpu {
1577 format!("gpu:{ordinal} is not available for model {model_name}")
1578 } else {
1579 format!("no GPU worker available for model {model_name}")
1580 };
1581 tracing::error!(model = %model_name, "{err_msg}");
1582 if let Some(tx) = rejected.progress_tx {
1583 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1584 message: err_msg.clone(),
1585 }));
1586 }
1587 let _ = rejected.result_tx.send(Err(err_msg));
1588 state.queue.decrement();
1589 state.job_registry.remove(&job_id);
1590 break;
1591 };
1592
1593 worker.in_flight.fetch_add(1, Ordering::SeqCst);
1595 let pending = gpu_job.take().expect("gpu_job present in retry loop");
1596 if preferred_gpu.is_none() {
1597 let _ = state
1598 .job_registry
1599 .set_target_gpu(&job_id, Some(worker.gpu.ordinal));
1600 }
1601 match worker.job_tx.try_send(pending) {
1602 Ok(()) => {
1603 dispatched = true;
1604 }
1605 Err(std::sync::mpsc::TrySendError::Full(j)) => {
1606 worker.in_flight.fetch_sub(1, Ordering::SeqCst);
1607 if preferred_gpu.is_none() {
1608 let _ = state.job_registry.set_target_gpu(&job_id, None);
1609 }
1610 gpu_job = Some(j);
1611 if preferred_gpu.is_none() {
1612 skip.push(worker.gpu.ordinal);
1613 if skip.len() >= state.gpu_pool.worker_count().max(1) {
1614 skip.clear();
1615 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1616 }
1617 } else {
1618 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1619 }
1620 }
1621 Err(std::sync::mpsc::TrySendError::Disconnected(j)) => {
1622 worker.in_flight.fetch_sub(1, Ordering::SeqCst);
1623 if preferred_gpu.is_none() {
1624 let _ = state.job_registry.set_target_gpu(&job_id, None);
1625 }
1626 tracing::warn!(
1627 gpu = worker.gpu.ordinal,
1628 "GPU worker disconnected — retrying dispatch"
1629 );
1630 gpu_job = Some(j);
1631 if preferred_gpu.is_none() {
1632 skip.push(worker.gpu.ordinal);
1633 } else {
1634 let rejected = gpu_job.take().expect("gpu_job retained after disconnect");
1635 let err_msg = format!(
1636 "gpu:{} disconnected while dispatching model {model_name}",
1637 worker.gpu.ordinal
1638 );
1639 if let Some(tx) = rejected.progress_tx {
1640 let _ = tx.send(SseMessage::Error(SseErrorEvent {
1641 message: err_msg.clone(),
1642 }));
1643 }
1644 let _ = rejected.result_tx.send(Err(err_msg));
1645 state.queue.decrement();
1646 state.job_registry.remove(&job_id);
1647 break;
1648 }
1649 }
1650 }
1651 }
1652 #[cfg(feature = "metrics")]
1653 crate::metrics::record_queue_depth(state.queue.pending());
1654 }
1655 tracing::info!("multi-GPU queue dispatcher shutting down");
1656}
1657
1658pub fn estimate_model_vram(model_name: &str) -> u64 {
1660 let lower = model_name.to_lowercase();
1663 if lower.contains("flux2")
1664 && lower.contains("9b")
1665 && (lower.contains(":bf16") || lower.contains(":fp16"))
1666 {
1667 32_000_000_000 } else if lower.contains(":q4") {
1669 6_000_000_000 } else if lower.contains(":q8") || lower.contains(":fp8") {
1671 12_000_000_000 } else if lower.contains(":bf16") || lower.contains(":fp16") {
1673 24_000_000_000 } else if lower.contains("sd15") || lower.contains("sd1.5") {
1675 4_000_000_000 } else {
1677 8_000_000_000
1679 }
1680}
1681
1682#[cfg(test)]
1683mod tests {
1684 use super::*;
1685 use crate::gpu_pool::{GpuPool, GpuWorker};
1686 use crate::model_cache::ModelCache;
1687 use crate::state::QueueHandle;
1688 use mold_core::{GenerateRequest, ImageData, ModelConfig, OutputFormat};
1689 use mold_db::MetadataDb;
1690 use mold_inference::device::DiscoveredGpu;
1691 use mold_inference::shared_pool::SharedPool;
1692 use std::sync::atomic::AtomicUsize;
1693 use std::sync::{Arc, Mutex, RwLock};
1694 use tempfile::TempDir;
1695
1696 fn fake_request(model: &str) -> GenerateRequest {
1699 GenerateRequest {
1700 prompt: "a cat".to_string(),
1701 negative_prompt: None,
1702 model: model.to_string(),
1703 width: 512,
1704 height: 512,
1705 steps: 4,
1706 guidance: 3.5,
1707 seed: Some(7),
1708 batch_size: 1,
1709 output_format: Some(OutputFormat::Png),
1710 embed_metadata: None,
1711 scheduler: None,
1712 cfg_plus: None,
1713 source_image: None,
1714 source_image_name: None,
1715 edit_images: None,
1716 strength: 0.75,
1717 mask_image: None,
1718 control_image: None,
1719 control_model: None,
1720 control_scale: 1.0,
1721 expand: None,
1722 original_prompt: None,
1723 batch_id: None,
1724 batch_index: None,
1725 batch_count: None,
1726 lora: None,
1727 frames: None,
1728 fps: None,
1729 upscale_model: None,
1730 gif_preview: false,
1731 enable_audio: None,
1732 audio_file: None,
1733 audio_file_path: None,
1734 source_video: None,
1735 source_video_path: None,
1736 keyframes: None,
1737 pipeline: None,
1738 loras: None,
1739 retake_range: None,
1740 spatial_upscale: None,
1741 temporal_upscale: None,
1742 placement: None,
1743 }
1744 }
1745
1746 fn fake_image() -> ImageData {
1747 ImageData {
1748 data: vec![0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A],
1751 format: OutputFormat::Png,
1752 width: 512,
1753 height: 512,
1754 index: 0,
1755 }
1756 }
1757
1758 #[test]
1759 fn multi_gpu_dispatch_identifies_missing_post_upscaler_for_auto_pull() {
1760 let mut req = fake_request("flux-dev:q4");
1761 req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
1762
1763 assert_eq!(
1764 post_upscale_model_to_pull(&mold_core::Config::default(), &req).unwrap(),
1765 Some("real-esrgan-x4plus:fp16".to_string())
1766 );
1767
1768 let tmp = TempDir::new().unwrap();
1769 let weights = tmp.path().join("realesrgan.safetensors");
1770 std::fs::write(&weights, b"test weights").unwrap();
1771 let mut config = mold_core::Config::default();
1772 config.models.insert(
1773 "real-esrgan-x4plus:fp16".to_string(),
1774 ModelConfig {
1775 transformer: Some(weights.display().to_string()),
1776 ..Default::default()
1777 },
1778 );
1779 assert_eq!(post_upscale_model_to_pull(&config, &req).unwrap(), None);
1780
1781 config
1782 .models
1783 .get_mut("real-esrgan-x4plus:fp16")
1784 .unwrap()
1785 .transformer = Some(tmp.path().join("missing.safetensors").display().to_string());
1786 assert_eq!(
1787 post_upscale_model_to_pull(&config, &req).unwrap(),
1788 Some("real-esrgan-x4plus:fp16".to_string()),
1789 "stale config paths should trigger a repair pull"
1790 );
1791 }
1792
1793 fn test_worker(
1794 ordinal: usize,
1795 channel_size: usize,
1796 ) -> (
1797 Arc<GpuWorker>,
1798 std::sync::mpsc::Receiver<crate::gpu_pool::GpuJob>,
1799 ) {
1800 let (job_tx, job_rx) = std::sync::mpsc::sync_channel(channel_size);
1801 let worker = Arc::new(GpuWorker {
1802 gpu: DiscoveredGpu {
1803 ordinal,
1804 name: format!("gpu{ordinal}"),
1805 total_vram_bytes: 24_000_000_000,
1806 free_vram_bytes: 24_000_000_000,
1807 },
1808 model_cache: Arc::new(Mutex::new(ModelCache::new(3))),
1809 active_generation: Arc::new(RwLock::new(None)),
1810 model_load_lock: Arc::new(Mutex::new(())),
1811 shared_pool: Arc::new(Mutex::new(SharedPool::new())),
1812 in_flight: AtomicUsize::new(0),
1813 consecutive_failures: AtomicUsize::new(0),
1814 poisoned: AtomicBool::new(false),
1815 fatal_cuda_error: Arc::new(AtomicBool::new(false)),
1816 fatal_cuda_shutdown: Arc::new(tokio::sync::Notify::new()),
1817 degraded_until: RwLock::new(None),
1818 job_tx,
1819 });
1820 (worker, job_rx)
1821 }
1822
1823 fn empty_test_state(config: mold_core::Config) -> crate::state::AppState {
1824 crate::state::AppState::empty(
1825 config,
1826 QueueHandle::new(tokio::sync::mpsc::channel(1).0),
1827 crate::state::AppState::empty_gpu_pool(),
1828 200,
1829 )
1830 }
1831
1832 #[test]
1833 fn save_image_to_dir_writes_file_and_creates_missing_dir() {
1834 let tmp = TempDir::new().unwrap();
1835 let nested = tmp.path().join("sub/output");
1836 assert!(!nested.exists());
1837
1838 save_image_to_dir(
1839 &nested,
1840 &fake_image(),
1841 "flux-dev:q4",
1842 1,
1843 None,
1844 None,
1845 None,
1846 None,
1847 );
1848
1849 assert!(nested.exists(), "save should mkdir -p");
1850 let entries: Vec<_> = std::fs::read_dir(&nested).unwrap().collect();
1851 assert_eq!(entries.len(), 1);
1852 let name = entries[0].as_ref().unwrap().file_name();
1853 let name_str = name.to_string_lossy();
1854 assert!(name_str.starts_with("mold-flux-dev-q4-"), "{name_str}");
1856 assert!(name_str.ends_with(".png"), "{name_str}");
1857 }
1858
1859 #[test]
1860 fn save_image_to_dir_includes_batch_index_when_batch_size_gt_1() {
1861 let tmp = TempDir::new().unwrap();
1862 let mut img = fake_image();
1863 img.index = 3;
1864 img.format = OutputFormat::Jpeg;
1865 img.data = vec![0xFF, 0xD8, 0xFF, 0xE0]; save_image_to_dir(tmp.path(), &img, "sdxl", 4, None, None, None, None);
1868
1869 let entries: Vec<_> = std::fs::read_dir(tmp.path()).unwrap().collect();
1870 let name = entries[0]
1871 .as_ref()
1872 .unwrap()
1873 .file_name()
1874 .to_string_lossy()
1875 .to_string();
1876 assert!(
1877 name.contains("-3.jpeg"),
1878 "expected batch index suffix: {name}"
1879 );
1880 }
1881
1882 #[test]
1883 fn save_image_to_dir_upserts_metadata_row_when_db_provided() {
1884 let tmp = TempDir::new().unwrap();
1885 let db = MetadataDb::open_in_memory().unwrap();
1886 let req = fake_request("flux-dev:q4");
1887 let meta = OutputMetadata::from_generate_request(&req, 42, None, "test-version");
1888
1889 save_image_to_dir(
1890 tmp.path(),
1891 &fake_image(),
1892 "flux-dev:q4",
1893 1,
1894 Some(&meta),
1895 Some(1234),
1896 Some(&db),
1897 None,
1898 );
1899
1900 let rows = db.list(Some(tmp.path())).unwrap();
1901 assert_eq!(rows.len(), 1, "exactly one DB row for the saved file");
1902 let rec = &rows[0];
1903 assert_eq!(rec.metadata.prompt, "a cat");
1904 assert_eq!(rec.metadata.seed, 42);
1905 assert_eq!(rec.metadata.version, "test-version");
1906 assert_eq!(rec.format, OutputFormat::Png);
1907 assert_eq!(rec.generation_time_ms, Some(1234));
1908 assert!(rec.file_size_bytes.unwrap_or(0) > 0);
1910 }
1911
1912 #[test]
1913 fn save_generated_image_outputs_persists_original_and_upscaled_dimensions() {
1914 let tmp = TempDir::new().unwrap();
1915 let db = MetadataDb::open_in_memory().unwrap();
1916 let mut req = fake_request("flux-dev:q4");
1917 req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
1918 let meta = OutputMetadata::from_generate_request(&req, 42, None, "test-version");
1919 let original = fake_image();
1920 let mut upscaled = fake_image();
1921 upscaled.width = 2048;
1922 upscaled.height = 2048;
1923 upscaled.data = vec![4, 5, 6];
1924
1925 save_generated_image_outputs(
1926 tmp.path(),
1927 Some(&original),
1928 &upscaled,
1929 "flux-dev:q4",
1930 1,
1931 &meta,
1932 Some(1234),
1933 Some(&db),
1934 None,
1935 );
1936
1937 let rows = db.list(Some(tmp.path())).unwrap();
1938 assert_eq!(rows.len(), 2);
1939 let original_row = rows
1940 .iter()
1941 .find(|row| row.filename.contains("-original."))
1942 .expect("original row");
1943 let upscaled_row = rows
1944 .iter()
1945 .find(|row| row.filename.contains("-upscaled."))
1946 .expect("upscaled row");
1947 assert_eq!(
1948 (original_row.metadata.width, original_row.metadata.height),
1949 (512, 512)
1950 );
1951 assert_eq!(
1952 (upscaled_row.metadata.width, upscaled_row.metadata.height),
1953 (2048, 2048)
1954 );
1955 assert_eq!(upscaled_row.metadata.generation_width, Some(512));
1956 assert_eq!(upscaled_row.metadata.generation_height, Some(512));
1957 }
1958
1959 #[test]
1960 fn save_image_to_dir_skips_db_when_metadata_is_none() {
1961 let tmp = TempDir::new().unwrap();
1962 let db = MetadataDb::open_in_memory().unwrap();
1963
1964 save_image_to_dir(
1965 tmp.path(),
1966 &fake_image(),
1967 "flux-dev:q4",
1968 1,
1969 None, Some(1234),
1971 Some(&db),
1972 None,
1973 );
1974
1975 assert_eq!(std::fs::read_dir(tmp.path()).unwrap().count(), 1);
1978 assert_eq!(db.list(None).unwrap().len(), 0);
1979 }
1980
1981 #[test]
1982 fn save_image_to_dir_invalid_path_does_not_panic() {
1983 save_image_to_dir(
1986 std::path::Path::new("/dev/null/cant-mkdir-here"),
1987 &fake_image(),
1988 "test",
1989 1,
1990 None,
1991 None,
1992 None,
1993 None,
1994 );
1995 }
1996
1997 #[test]
1998 fn save_image_to_dir_emits_gallery_added_with_row_when_db_records() {
1999 let tmp = TempDir::new().unwrap();
2000 let db = MetadataDb::open_in_memory().unwrap();
2001 let req = fake_request("flux-dev:q4");
2002 let meta = OutputMetadata::from_generate_request(&req, 42, None, "test-version");
2003 let events = crate::events::EventBroadcaster::new();
2004 let mut rx = events.subscribe();
2005
2006 save_image_to_dir(
2007 tmp.path(),
2008 &fake_image(),
2009 "flux-dev:q4",
2010 1,
2011 Some(&meta),
2012 Some(1234),
2013 Some(&db),
2014 Some(&events),
2015 );
2016
2017 match rx.try_recv().unwrap() {
2018 mold_core::ServerEvent::GalleryAdded { filename, image } => {
2019 assert!(filename.ends_with(".png"), "{filename}");
2020 let img = image.expect("DB recorded — event must carry the gallery row");
2021 assert_eq!(img.filename, filename);
2022 assert_eq!(img.metadata.prompt, "a cat");
2023 }
2024 other => panic!("expected gallery_added, got {other:?}"),
2025 }
2026 }
2027
2028 #[test]
2029 fn save_image_to_dir_emits_gallery_added_without_row_when_db_absent() {
2030 let tmp = TempDir::new().unwrap();
2031 let events = crate::events::EventBroadcaster::new();
2032 let mut rx = events.subscribe();
2033
2034 save_image_to_dir(
2035 tmp.path(),
2036 &fake_image(),
2037 "flux-dev:q4",
2038 1,
2039 None,
2040 None,
2041 None, Some(&events),
2043 );
2044
2045 match rx.try_recv().unwrap() {
2046 mold_core::ServerEvent::GalleryAdded { image, .. } => {
2047 assert!(image.is_none(), "no DB → clients must refetch");
2048 }
2049 other => panic!("expected gallery_added, got {other:?}"),
2050 }
2051 }
2052
2053 #[test]
2054 fn save_image_to_dir_emits_nothing_on_write_failure() {
2055 let events = crate::events::EventBroadcaster::new();
2056 let mut rx = events.subscribe();
2057
2058 save_image_to_dir(
2059 std::path::Path::new("/dev/null/cant-mkdir-here"),
2060 &fake_image(),
2061 "test",
2062 1,
2063 None,
2064 None,
2065 None,
2066 Some(&events),
2067 );
2068
2069 assert!(
2070 rx.try_recv().is_err(),
2071 "failed save must not announce a gallery entry"
2072 );
2073 }
2074
2075 #[test]
2076 fn save_video_to_dir_emits_gallery_added() {
2077 let tmp = TempDir::new().unwrap();
2078 let db = MetadataDb::open_in_memory().unwrap();
2079 let req = fake_request("ltx-video:fp16");
2080 let meta = OutputMetadata::from_generate_request(&req, 1, None, "v");
2081 let events = crate::events::EventBroadcaster::new();
2082 let mut rx = events.subscribe();
2083
2084 save_video_to_dir(
2085 tmp.path(),
2086 b"fake mp4 bytes",
2087 b"",
2088 OutputFormat::Mp4,
2089 "ltx-video:fp16",
2090 &meta,
2091 Some(5000),
2092 Some(&db),
2093 Some(&events),
2094 );
2095
2096 match rx.try_recv().unwrap() {
2097 mold_core::ServerEvent::GalleryAdded { filename, image } => {
2098 assert!(filename.ends_with(".mp4"), "{filename}");
2099 assert!(image.is_some());
2100 }
2101 other => panic!("expected gallery_added, got {other:?}"),
2102 }
2103 }
2104
2105 #[test]
2106 fn save_video_to_dir_writes_mp4_and_records_metadata() {
2107 let tmp = TempDir::new().unwrap();
2108 let db = MetadataDb::open_in_memory().unwrap();
2109 let mut req = fake_request("ltx-video:fp16");
2110 req.frames = Some(25);
2111 req.fps = Some(24);
2112 let meta = OutputMetadata::from_generate_request(&req, 99, None, "test-version");
2113
2114 let bytes = b"\x00\x00\x00\x18ftypmp42\x00\x00\x00\x00mp42isom".to_vec();
2117
2118 save_video_to_dir(
2119 tmp.path(),
2120 &bytes,
2121 b"",
2122 OutputFormat::Mp4,
2123 "ltx-video:fp16",
2124 &meta,
2125 Some(5000),
2126 Some(&db),
2127 None,
2128 );
2129
2130 let entries: Vec<_> = std::fs::read_dir(tmp.path()).unwrap().collect();
2131 assert_eq!(entries.len(), 1);
2132 let name = entries[0]
2133 .as_ref()
2134 .unwrap()
2135 .file_name()
2136 .to_string_lossy()
2137 .to_string();
2138 assert!(name.starts_with("mold-ltx-video-fp16-"), "{name}");
2139 assert!(name.ends_with(".mp4"), "{name}");
2140
2141 let rows = db.list(Some(tmp.path())).unwrap();
2142 assert_eq!(rows.len(), 1);
2143 assert_eq!(rows[0].format, OutputFormat::Mp4);
2144 assert_eq!(rows[0].metadata.frames, Some(25));
2145 assert_eq!(rows[0].metadata.fps, Some(24));
2146 assert_eq!(rows[0].generation_time_ms, Some(5000));
2147 }
2148
2149 #[test]
2150 fn save_video_to_dir_without_db_still_writes_file() {
2151 let tmp = TempDir::new().unwrap();
2152 let req = fake_request("ltx-video:fp16");
2153 let meta = OutputMetadata::from_generate_request(&req, 1, None, "v");
2154
2155 save_video_to_dir(
2156 tmp.path(),
2157 b"fake gif bytes",
2158 b"",
2159 OutputFormat::Gif,
2160 "ltx-video:fp16",
2161 &meta,
2162 None,
2163 None,
2164 None,
2165 );
2166
2167 let entries: Vec<_> = std::fs::read_dir(tmp.path()).unwrap().collect();
2168 assert_eq!(entries.len(), 1);
2169 let name = entries[0]
2170 .as_ref()
2171 .unwrap()
2172 .file_name()
2173 .to_string_lossy()
2174 .to_string();
2175 assert!(name.ends_with(".gif"), "{name}");
2176 }
2177
2178 #[test]
2179 fn save_video_to_dir_invalid_path_does_not_panic() {
2180 let req = fake_request("ltx-video:fp16");
2181 let meta = OutputMetadata::from_generate_request(&req, 1, None, "v");
2182 save_video_to_dir(
2183 std::path::Path::new("/dev/null/nope"),
2184 b"x",
2185 b"",
2186 OutputFormat::Mp4,
2187 "test",
2188 &meta,
2189 None,
2190 None,
2191 None,
2192 );
2193 }
2194
2195 #[test]
2202 fn save_video_preview_gif_writes_to_preview_cache() {
2203 let td = tempfile::tempdir().unwrap();
2204 let preview_dir = td.path().join("cache").join("previews");
2205
2206 const GIF: &[u8] = b"GIF89a\x01\x00\x01\x00\x00\x00\x00\x3b";
2207 save_video_preview_gif_to(&preview_dir, "ltx2-42.mp4", GIF);
2208
2209 let expected = preview_dir.join("ltx2-42.mp4.preview.gif");
2210 assert!(
2211 expected.is_file(),
2212 "preview gif should land at {}",
2213 expected.display()
2214 );
2215 assert_eq!(std::fs::read(&expected).unwrap(), GIF);
2216 }
2217
2218 #[test]
2219 fn build_sse_complete_event_video_carries_mp4_payload_and_metadata() {
2220 let video = mold_core::VideoData {
2227 data: vec![0x00, 0x00, 0x00, 0x18, b'f', b't', b'y', b'p'],
2228 format: OutputFormat::Mp4,
2229 width: 768,
2230 height: 512,
2231 frames: 25,
2232 fps: 24,
2233 thumbnail: vec![0x89, 0x50, 0x4E, 0x47],
2234 gif_preview: vec![b'G', b'I', b'F', b'8'],
2235 has_audio: true,
2236 duration_ms: Some(1040),
2237 audio_sample_rate: Some(44100),
2238 audio_channels: Some(2),
2239 };
2240 let resp = mold_core::GenerateResponse {
2241 images: vec![],
2242 video: Some(video.clone()),
2243 generation_time_ms: 1234,
2244 model: "ltx-2-19b-distilled:fp8".to_string(),
2245 seed_used: 7,
2246 gpu: Some(0),
2247 };
2248 let thumb_img = ImageData {
2251 data: video.thumbnail.clone(),
2252 format: OutputFormat::Png,
2253 width: video.width,
2254 height: video.height,
2255 index: 0,
2256 };
2257
2258 let event = build_sse_complete_event(
2259 &resp,
2260 &thumb_img,
2261 None,
2262 None,
2263 &SavedOutputNames::default(),
2264 SseCompletionPayload::Full,
2265 );
2266
2267 let b64 = base64::engine::general_purpose::STANDARD;
2268 assert_eq!(event.image, b64.encode(&video.data));
2269 assert_eq!(event.format, OutputFormat::Mp4);
2270 assert_eq!(event.video_frames, Some(25));
2271 assert_eq!(event.video_fps, Some(24));
2272 assert_eq!(event.video_thumbnail, Some(b64.encode(&video.thumbnail)));
2273 assert_eq!(
2274 event.video_gif_preview,
2275 Some(b64.encode(&video.gif_preview))
2276 );
2277 assert!(event.video_has_audio);
2278 assert_eq!(event.video_duration_ms, Some(1040));
2279 assert_eq!(event.gpu, Some(0));
2280
2281 let saved = SavedOutputNames {
2282 output: Some("generated-video.mp4".to_string()),
2283 original: None,
2284 };
2285 let metadata_only = build_sse_complete_event(
2286 &resp,
2287 &thumb_img,
2288 None,
2289 None,
2290 &saved,
2291 SseCompletionPayload::MetadataOnly,
2292 );
2293 assert!(metadata_only.image.is_empty());
2294 assert!(metadata_only.video_thumbnail.is_none());
2295 assert!(metadata_only.video_gif_preview.is_none());
2296 assert_eq!(metadata_only.video_frames, Some(25));
2297 assert_eq!(
2298 metadata_only.filename.as_deref(),
2299 Some("generated-video.mp4")
2300 );
2301 }
2302
2303 #[test]
2304 fn build_sse_complete_event_video_empty_gif_preview_omits_field() {
2305 let video = mold_core::VideoData {
2306 data: vec![0x00, 0x00, 0x00, 0x18],
2307 format: OutputFormat::Mp4,
2308 width: 256,
2309 height: 256,
2310 frames: 17,
2311 fps: 12,
2312 thumbnail: vec![0x89, 0x50],
2313 gif_preview: Vec::new(),
2314 has_audio: false,
2315 duration_ms: None,
2316 audio_sample_rate: None,
2317 audio_channels: None,
2318 };
2319 let resp = mold_core::GenerateResponse {
2320 images: vec![],
2321 video: Some(video),
2322 generation_time_ms: 0,
2323 model: "m".to_string(),
2324 seed_used: 0,
2325 gpu: None,
2326 };
2327 let event = build_sse_complete_event(
2328 &resp,
2329 &fake_image(),
2330 None,
2331 None,
2332 &SavedOutputNames::default(),
2333 SseCompletionPayload::Full,
2334 );
2335 assert!(event.video_gif_preview.is_none());
2336 assert!(!event.video_has_audio);
2337 }
2338
2339 #[test]
2340 fn build_sse_complete_event_image_clears_all_video_fields() {
2341 let resp = mold_core::GenerateResponse {
2342 images: vec![fake_image()],
2343 video: None,
2344 generation_time_ms: 100,
2345 model: "flux-schnell:q8".to_string(),
2346 seed_used: 5,
2347 gpu: None,
2348 };
2349 let event = build_sse_complete_event(
2350 &resp,
2351 &fake_image(),
2352 None,
2353 None,
2354 &SavedOutputNames::default(),
2355 SseCompletionPayload::Full,
2356 );
2357 assert_eq!(event.format, OutputFormat::Png);
2358 assert!(event.video_frames.is_none());
2359 assert!(event.video_fps.is_none());
2360 assert!(event.video_thumbnail.is_none());
2361 assert!(event.video_gif_preview.is_none());
2362 assert!(!event.video_has_audio);
2363 assert!(event.video_duration_ms.is_none());
2364 }
2365
2366 #[test]
2367 fn build_sse_complete_event_carries_saved_names_and_recorded_metadata() {
2368 let mut req = fake_request("flux-dev:q4");
2369 req.batch_id = Some("prepared-batch-1".to_string());
2370 req.batch_index = Some(2);
2371 req.batch_count = Some(3);
2372 let resp = mold_core::GenerateResponse {
2373 images: vec![fake_image()],
2374 video: None,
2375 generation_time_ms: 100,
2376 model: "flux-dev:q4".to_string(),
2377 seed_used: 5,
2378 gpu: None,
2379 };
2380 let metadata =
2381 OutputMetadata::from_generate_request(&req, resp.seed_used, None, "test-version");
2382 let saved = SavedOutputNames {
2383 output: Some("flux-dev-q4-123.png".to_string()),
2384 original: Some("flux-dev-q4-123-original.png".to_string()),
2385 };
2386 let event = build_sse_complete_event(
2387 &resp,
2388 &fake_image(),
2389 None,
2390 Some(&metadata),
2391 &saved,
2392 SseCompletionPayload::Full,
2393 );
2394 assert_eq!(event.filename.as_deref(), Some("flux-dev-q4-123.png"));
2395 assert_eq!(
2396 event.original_filename.as_deref(),
2397 Some("flux-dev-q4-123-original.png")
2398 );
2399 let meta = event.metadata.expect("metadata rides the complete event");
2402 assert_eq!(meta.seed, 5);
2403 assert_eq!(meta.width, fake_image().width);
2404 assert_eq!(meta.height, fake_image().height);
2405 assert_eq!(meta.batch_id.as_deref(), Some("prepared-batch-1"));
2406 assert_eq!(meta.batch_index, Some(2));
2407 assert_eq!(meta.batch_count, Some(3));
2408
2409 let metadata_only = build_sse_complete_event(
2410 &resp,
2411 &fake_image(),
2412 Some(&fake_image()),
2413 Some(&metadata),
2414 &saved,
2415 SseCompletionPayload::MetadataOnly,
2416 );
2417 assert!(metadata_only.image.is_empty());
2418 assert!(metadata_only.original_image.is_none());
2419 assert_eq!(
2420 metadata_only.filename.as_deref(),
2421 Some("flux-dev-q4-123.png")
2422 );
2423 assert!(metadata_only.metadata.is_some());
2424 }
2425
2426 #[test]
2427 fn metadata_only_completion_fails_when_the_output_was_not_saved() {
2428 let response = mold_core::GenerateResponse {
2429 images: vec![fake_image()],
2430 video: None,
2431 generation_time_ms: 100,
2432 model: "flux-dev:q4".to_string(),
2433 seed_used: 5,
2434 gpu: None,
2435 };
2436 let message = build_sse_completion_message(
2437 &response,
2438 &fake_image(),
2439 None,
2440 None,
2441 &SavedOutputNames::default(),
2442 SseCompletionPayload::MetadataOnly,
2443 );
2444 match message {
2445 SseMessage::Error(error) => assert!(error.message.contains("could not be saved")),
2446 _ => panic!("metadata-only completion without a file must be an SSE error"),
2447 }
2448 }
2449
2450 #[test]
2451 fn post_generation_upscale_replaces_image_response_dimensions() {
2452 let mut req = fake_request("flux-dev:q4");
2453 req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
2454 let mut response = mold_core::GenerateResponse {
2455 images: vec![],
2456 video: None,
2457 generation_time_ms: 100,
2458 model: "flux-dev:q4".to_string(),
2459 seed_used: 5,
2460 gpu: None,
2461 };
2462 let img = fake_image();
2463 let upscaled = mold_core::UpscaleResponse {
2464 image: ImageData {
2465 data: vec![1, 2, 3],
2466 format: OutputFormat::Png,
2467 width: 2048,
2468 height: 2048,
2469 index: 0,
2470 },
2471 upscale_time_ms: 42,
2472 model: "real-esrgan-x4plus:fp16".to_string(),
2473 scale_factor: 4,
2474 original_width: 512,
2475 original_height: 512,
2476 };
2477
2478 let next = apply_upscale_response_to_image_generation(&req, &mut response, img, upscaled)
2479 .expect("image upscale should apply");
2480 let event = build_sse_complete_event(
2481 &response,
2482 &next,
2483 Some(&fake_image()),
2484 None,
2485 &SavedOutputNames::default(),
2486 SseCompletionPayload::Full,
2487 );
2488 assert!(event.original_image.is_some());
2489 assert_eq!(event.original_width, Some(512));
2490 assert_eq!(event.original_height, Some(512));
2491 let mut metadata =
2492 OutputMetadata::from_generate_request(&req, response.seed_used, None, "test-version");
2493 apply_output_dimensions_to_metadata(&mut metadata, &next);
2494
2495 assert_eq!(next.width, 2048);
2496 assert_eq!(next.height, 2048);
2497 assert_eq!(event.width, 2048);
2498 assert_eq!(event.height, 2048);
2499 assert_eq!(metadata.width, 2048);
2500 assert_eq!(metadata.height, 2048);
2501 assert_eq!(metadata.generation_width, Some(512));
2502 assert_eq!(metadata.generation_height, Some(512));
2503 assert_eq!(
2504 metadata.upscale_model.as_deref(),
2505 Some("real-esrgan-x4plus:fp16")
2506 );
2507 }
2508
2509 #[test]
2510 fn failed_post_generation_upscale_keeps_only_the_original_output() {
2511 let original = fake_image();
2512 let (output, preserved_original, error) = settle_post_generation_upscale(
2513 original.clone(),
2514 Err("upscaler unavailable".to_string()),
2515 );
2516
2517 assert_eq!(output.data, original.data);
2518 assert!(preserved_original.is_none());
2519 assert_eq!(error.as_deref(), Some("upscaler unavailable"));
2520 }
2521
2522 #[test]
2523 fn post_generation_upscale_skips_video_responses() {
2524 let mut req = fake_request("ltx-video:fp16");
2525 req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
2526 let video = mold_core::VideoData {
2527 data: vec![0, 0, 0, 24],
2528 format: OutputFormat::Mp4,
2529 width: 512,
2530 height: 512,
2531 frames: 25,
2532 fps: 24,
2533 thumbnail: vec![9, 9],
2534 gif_preview: vec![],
2535 has_audio: false,
2536 duration_ms: None,
2537 audio_sample_rate: None,
2538 audio_channels: None,
2539 };
2540 let mut response = mold_core::GenerateResponse {
2541 images: vec![],
2542 video: Some(video),
2543 generation_time_ms: 100,
2544 model: "ltx-video:fp16".to_string(),
2545 seed_used: 5,
2546 gpu: None,
2547 };
2548 let img = fake_image();
2549 let upscaled = mold_core::UpscaleResponse {
2550 image: ImageData {
2551 data: vec![1, 2, 3],
2552 format: OutputFormat::Png,
2553 width: 2048,
2554 height: 2048,
2555 index: 0,
2556 },
2557 upscale_time_ms: 42,
2558 model: "real-esrgan-x4plus:fp16".to_string(),
2559 scale_factor: 4,
2560 original_width: 512,
2561 original_height: 512,
2562 };
2563
2564 let next = apply_upscale_response_to_image_generation(&req, &mut response, img, upscaled)
2565 .expect("video upscale should be skipped");
2566
2567 assert_eq!(next.width, 512);
2568 assert_eq!(next.height, 512);
2569 assert!(response.video.is_some());
2570 }
2571
2572 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2573 async fn single_worker_post_upscale_noops_without_model() {
2574 let state = empty_test_state(mold_core::Config::default());
2575 let req = fake_request("flux-dev:q4");
2576
2577 let next = upscale_generated_image_on_single_worker(&state, &req, 5, fake_image(), None)
2578 .await
2579 .expect("missing upscale model should leave the image unchanged");
2580
2581 assert_eq!(next.width, 512);
2582 assert_eq!(next.height, 512);
2583 assert_eq!(next.index, 0);
2584 }
2585
2586 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2587 async fn single_worker_post_upscale_rejects_unknown_upscaler_manifest() {
2588 let state = empty_test_state(mold_core::Config::default());
2589 let mut req = fake_request("flux-dev:q4");
2590 req.upscale_model = Some("definitely-not-a-real-upscaler:fp16".to_string());
2591 let (progress_tx, mut progress_rx) = tokio::sync::mpsc::unbounded_channel();
2592
2593 let err = upscale_generated_image_on_single_worker(
2594 &state,
2595 &req,
2596 5,
2597 fake_image(),
2598 Some(&progress_tx),
2599 )
2600 .await
2601 .expect_err("unknown upscalers should fail before generation completes");
2602
2603 assert!(err.contains("unknown upscaler model"), "got: {err}");
2604 let first_progress = progress_rx
2605 .try_recv()
2606 .expect("loading stage should be emitted before validation fails");
2607 assert!(matches!(
2608 first_progress,
2609 SseMessage::Progress(SseProgressEvent::StageStart { .. })
2610 ));
2611 }
2612
2613 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2614 async fn single_worker_post_upscale_surfaces_missing_weights_path() {
2615 let tmp = TempDir::new().unwrap();
2616 let missing_weights = tmp.path().join("missing-upscaler.safetensors");
2617 let mut config = mold_core::Config::default();
2618 config.models.insert(
2619 "real-esrgan-x4plus:fp16".to_string(),
2620 ModelConfig {
2621 transformer: Some(missing_weights.display().to_string()),
2622 ..Default::default()
2623 },
2624 );
2625 let state = empty_test_state(config);
2626 let mut req = fake_request("flux-dev:q4");
2627 req.upscale_model = Some("real-esrgan-x4plus:fp16".to_string());
2628 let (progress_tx, mut progress_rx) = tokio::sync::mpsc::unbounded_channel();
2629
2630 let err = upscale_generated_image_on_single_worker(
2631 &state,
2632 &req,
2633 5,
2634 fake_image(),
2635 Some(&progress_tx),
2636 )
2637 .await
2638 .expect_err("missing weight files should be surfaced");
2639
2640 assert!(err.contains("upscale failed"), "got: {err}");
2641 assert!(err.contains("upscaler weights not found"), "got: {err}");
2642 let first_progress = progress_rx
2643 .try_recv()
2644 .expect("loading stage should be emitted before loading fails");
2645 assert!(matches!(
2646 first_progress,
2647 SseMessage::Progress(SseProgressEvent::StageStart { .. })
2648 ));
2649 }
2650
2651 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2652 async fn queue_dispatcher_waits_for_worker_capacity_instead_of_rejecting() {
2653 let (worker, worker_rx) = test_worker(0, 1);
2654 let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
2655 let queue = QueueHandle::new(job_tx.clone());
2656 let state = crate::state::AppState::empty(
2657 mold_core::Config::default(),
2658 queue.clone(),
2659 Arc::new(GpuPool {
2660 workers: vec![worker.clone()],
2661 }),
2662 8,
2663 );
2664
2665 let (filler_result_tx, _filler_result_rx) = tokio::sync::oneshot::channel();
2666 let filler_job = crate::gpu_pool::GpuJob {
2667 id: String::new(),
2668 model: "busy-model".to_string(),
2669 request: fake_request("busy-model"),
2670 completion_payload: SseCompletionPayload::Full,
2671 progress_tx: None,
2672 result_tx: filler_result_tx,
2673 output_dir: None,
2674 config: state.config.clone(),
2675 metadata_db: state.metadata_db.clone(),
2676 queue: state.queue.clone(),
2677 registry: state.job_registry.clone(),
2678 events: state.events.clone(),
2679 };
2680 worker.job_tx.send(filler_job).unwrap();
2681
2682 let dispatcher = tokio::spawn(run_queue_dispatcher_with_tuning(
2683 job_rx,
2684 state.clone(),
2685 8,
2686 DEFAULT_MAX_DEFERRALS,
2687 ));
2688
2689 let (result_tx, mut result_rx) = tokio::sync::oneshot::channel();
2690 let job = crate::state::GenerationJob {
2691 id: String::new(),
2692 request: fake_request("flux-dev:q4"),
2693 completion_payload: SseCompletionPayload::Full,
2694 progress_tx: None,
2695 result_tx,
2696 output_dir: None,
2697 };
2698 let _position = queue.submit(job, 8).await.unwrap();
2699
2700 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
2701 assert!(
2702 result_rx.try_recv().is_err(),
2703 "dispatcher should keep the job pending while all worker channels are full"
2704 );
2705
2706 let _filler = worker_rx
2707 .recv()
2708 .expect("filler job should occupy the local channel");
2709 let dispatched = worker_rx
2710 .recv_timeout(std::time::Duration::from_secs(1))
2711 .expect("queued job should dispatch once capacity is available");
2712 assert_eq!(dispatched.model, "flux-dev:q4");
2713
2714 drop(job_tx);
2715 dispatcher.abort();
2716 }
2717
2718 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2719 async fn queue_dispatcher_waits_for_degraded_worker_recovery_instead_of_rejecting() {
2720 let (worker, worker_rx) = test_worker(0, 1);
2721 worker.consecutive_failures.store(3, Ordering::SeqCst);
2722 *worker.degraded_until.write().unwrap() =
2723 Some(Instant::now() + std::time::Duration::from_secs(60));
2724
2725 let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
2726 let queue = QueueHandle::new(job_tx.clone());
2727 let state = crate::state::AppState::empty(
2728 mold_core::Config::default(),
2729 queue.clone(),
2730 Arc::new(GpuPool {
2731 workers: vec![worker.clone()],
2732 }),
2733 8,
2734 );
2735 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
2736
2737 let (result_tx, mut result_rx) = tokio::sync::oneshot::channel();
2738 let job = crate::state::GenerationJob {
2739 id: String::new(),
2740 request: fake_request("flux-dev:q4"),
2741 completion_payload: SseCompletionPayload::Full,
2742 progress_tx: None,
2743 result_tx,
2744 output_dir: None,
2745 };
2746 queue.submit(job, 8).await.unwrap();
2747
2748 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
2749 assert!(
2750 result_rx.try_recv().is_err(),
2751 "dispatcher should keep the job pending while all workers are degraded"
2752 );
2753 assert!(
2754 worker_rx.try_recv().is_err(),
2755 "degraded worker must not receive work before recovery"
2756 );
2757
2758 worker.consecutive_failures.store(0, Ordering::SeqCst);
2759 *worker.degraded_until.write().unwrap() = None;
2760
2761 let dispatched = worker_rx
2762 .recv_timeout(std::time::Duration::from_secs(1))
2763 .expect("queued job should dispatch once a worker recovers");
2764 assert_eq!(dispatched.model, "flux-dev:q4");
2765
2766 drop(job_tx);
2767 dispatcher.abort();
2768 }
2769
2770 #[tokio::test]
2777 async fn cache_take_on_vanished_engine_returns_none_not_panic() {
2778 use crate::model_cache::ModelCache;
2779 use mold_core::GenerateResponse;
2780 use mold_inference::InferenceEngine;
2781
2782 struct StubEngine(&'static str);
2783 impl InferenceEngine for StubEngine {
2784 fn generate(&mut self, _r: &GenerateRequest) -> anyhow::Result<GenerateResponse> {
2785 unimplemented!()
2786 }
2787 fn model_name(&self) -> &str {
2788 self.0
2789 }
2790 fn is_loaded(&self) -> bool {
2791 true
2792 }
2793 fn load(&mut self) -> anyhow::Result<()> {
2794 Ok(())
2795 }
2796 }
2797
2798 let mut cache = ModelCache::new(3);
2799 assert!(cache.take("vanished-model").is_none());
2802
2803 cache.insert(Box::new(StubEngine("present-model")), 0);
2807 let first = cache.take("present-model");
2808 assert!(first.is_some());
2809 assert!(
2810 cache.take("present-model").is_none(),
2811 "double-take must return None"
2812 );
2813 }
2814
2815 fn buf_job(model: &str) -> BufferedJob {
2816 let (tx, _rx) = tokio::sync::oneshot::channel();
2817 BufferedJob::new(crate::state::GenerationJob {
2818 id: String::new(),
2819 request: fake_request(model),
2820 completion_payload: SseCompletionPayload::Full,
2821 progress_tx: None,
2822 result_tx: tx,
2823 output_dir: None,
2824 })
2825 }
2826
2827 fn buf_job_with_id(id: &str, model: &str) -> BufferedJob {
2828 let (tx, _rx) = tokio::sync::oneshot::channel();
2829 BufferedJob::new(crate::state::GenerationJob {
2830 id: id.to_string(),
2831 request: fake_request(model),
2832 completion_payload: SseCompletionPayload::Full,
2833 progress_tx: None,
2834 result_tx: tx,
2835 output_dir: None,
2836 })
2837 }
2838
2839 #[test]
2840 fn align_buffer_reorders_to_match_registry_queued_order() {
2841 use std::collections::VecDeque;
2842 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2843 for id in ["a", "b", "c"] {
2844 buffer.push_back(buf_job_with_id(id, &format!("model-{id}")));
2845 }
2846 let order = vec!["c".to_string(), "a".to_string(), "b".to_string()];
2848 align_buffer_to_registry_order(&mut buffer, &order);
2849 let ids: Vec<&str> = buffer.iter().map(|b| b.job.id.as_str()).collect();
2850 assert_eq!(ids, vec!["c", "a", "b"]);
2851 }
2852
2853 #[test]
2854 fn align_buffer_is_a_noop_when_already_in_registry_order() {
2855 use std::collections::VecDeque;
2856 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2857 buffer.push_back(buf_job_with_id("a", "model-a"));
2858 let mut b = buf_job_with_id("b", "model-b");
2861 b.deferred = 2;
2862 buffer.push_back(b);
2863 buffer.push_back(buf_job_with_id("c", "model-c"));
2864 let order = vec!["a".to_string(), "b".to_string(), "c".to_string()];
2865 align_buffer_to_registry_order(&mut buffer, &order);
2866 let ids: Vec<&str> = buffer.iter().map(|b| b.job.id.as_str()).collect();
2867 assert_eq!(ids, vec!["a", "b", "c"]);
2868 assert_eq!(
2869 buffer[1].deferred, 2,
2870 "a no-op align must preserve the deferred starvation count"
2871 );
2872 }
2873
2874 #[test]
2875 fn align_buffer_keeps_unregistered_jobs_in_arrival_order_at_the_back() {
2876 use std::collections::VecDeque;
2877 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2878 for id in ["x", "b", "y", "a"] {
2881 buffer.push_back(buf_job_with_id(id, "m"));
2882 }
2883 let order = vec!["a".to_string(), "b".to_string()];
2884 align_buffer_to_registry_order(&mut buffer, &order);
2885 let ids: Vec<&str> = buffer.iter().map(|b| b.job.id.as_str()).collect();
2886 assert_eq!(ids, vec!["a", "b", "x", "y"]);
2889 }
2890
2891 #[test]
2892 fn align_buffer_leaves_empty_id_jobs_untouched() {
2893 use std::collections::VecDeque;
2897 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2898 for model in ["a", "b", "a", "b"] {
2899 buffer.push_back(buf_job(model));
2900 }
2901 align_buffer_to_registry_order(&mut buffer, &[]);
2902 let models: Vec<&str> = buffer
2903 .iter()
2904 .map(|b| b.job.request.model.as_str())
2905 .collect();
2906 assert_eq!(models, vec!["a", "b", "a", "b"]);
2907 }
2908
2909 #[test]
2910 fn pick_next_job_picks_head_when_head_model_loaded() {
2911 use std::collections::{HashSet, VecDeque};
2912 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2913 buffer.push_back(buf_job("a"));
2914 buffer.push_back(buf_job("b"));
2915 buffer.push_back(buf_job("a"));
2916 let loaded: HashSet<String> = ["a".to_string()].into_iter().collect();
2917 let picked = pick_next_job(&mut buffer, &loaded, 3);
2918 assert_eq!(picked.request.model, "a");
2919 assert_eq!(buffer.len(), 2);
2920 assert_eq!(buffer.front().unwrap().job.request.model, "b");
2921 assert_eq!(
2922 buffer.front().unwrap().deferred,
2923 0,
2924 "head shouldn't be deferred when picker chose the head itself"
2925 );
2926 }
2927
2928 #[test]
2929 fn pick_next_job_picks_non_head_when_only_non_head_model_loaded() {
2930 use std::collections::{HashSet, VecDeque};
2931 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2932 buffer.push_back(buf_job("a"));
2933 buffer.push_back(buf_job("b"));
2934 buffer.push_back(buf_job("a"));
2935 let loaded: HashSet<String> = ["b".to_string()].into_iter().collect();
2936 let picked = pick_next_job(&mut buffer, &loaded, 3);
2937 assert_eq!(picked.request.model, "b");
2938 assert_eq!(buffer.len(), 2);
2939 assert_eq!(buffer.front().unwrap().job.request.model, "a");
2941 assert_eq!(buffer.front().unwrap().deferred, 1);
2942 }
2943
2944 #[test]
2945 fn pick_next_job_force_dispatches_head_after_max_deferrals() {
2946 use std::collections::{HashSet, VecDeque};
2947 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2948 let mut head = buf_job("a");
2949 head.deferred = 3;
2950 buffer.push_back(head);
2951 buffer.push_back(buf_job("b"));
2952 let loaded: HashSet<String> = ["b".to_string()].into_iter().collect();
2954 let picked = pick_next_job(&mut buffer, &loaded, 3);
2955 assert_eq!(picked.request.model, "a");
2956 assert_eq!(buffer.len(), 1);
2957 assert_eq!(buffer.front().unwrap().job.request.model, "b");
2958 }
2959
2960 #[test]
2961 fn pick_next_job_falls_back_to_head_when_nothing_loaded() {
2962 use std::collections::{HashSet, VecDeque};
2963 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2964 buffer.push_back(buf_job("a"));
2965 buffer.push_back(buf_job("b"));
2966 let loaded: HashSet<String> = HashSet::new();
2967 let picked = pick_next_job(&mut buffer, &loaded, 3);
2968 assert_eq!(picked.request.model, "a");
2969 }
2970
2971 #[test]
2975 fn pick_next_job_max_deferrals_zero_picks_head_even_when_non_head_loaded() {
2976 use std::collections::{HashSet, VecDeque};
2977 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2978 buffer.push_back(buf_job("b")); buffer.push_back(buf_job("a")); let loaded: HashSet<String> = ["a".to_string()].into_iter().collect();
2981 let picked = pick_next_job(&mut buffer, &loaded, 0);
2982 assert_eq!(
2983 picked.request.model, "b",
2984 "max_deferrals=0 must force FIFO — head must win even when only the non-head model is loaded"
2985 );
2986 assert_eq!(buffer.len(), 1);
2987 assert_eq!(buffer.front().unwrap().job.request.model, "a");
2988 }
2989
2990 #[test]
2994 fn pick_next_job_max_deferrals_zero_with_empty_loaded_picks_head() {
2995 use std::collections::{HashSet, VecDeque};
2996 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
2997 buffer.push_back(buf_job("a")); buffer.push_back(buf_job("b"));
2999 let loaded: HashSet<String> = HashSet::new();
3000 let picked = pick_next_job(&mut buffer, &loaded, 0);
3001 assert_eq!(picked.request.model, "a");
3002 assert_eq!(buffer.len(), 1);
3003 assert_eq!(buffer.front().unwrap().job.request.model, "b");
3004 }
3005
3006 #[test]
3011 fn pick_next_job_picks_front_most_match_when_multiple_loaded() {
3012 use std::collections::{HashSet, VecDeque};
3013 let mut buffer: VecDeque<BufferedJob> = VecDeque::new();
3014 buffer.push_back(buf_job("a"));
3015 buffer.push_back(buf_job("b"));
3016 buffer.push_back(buf_job("a"));
3017 buffer.push_back(buf_job("b"));
3018 let loaded: HashSet<String> = ["a".to_string(), "b".to_string()].into_iter().collect();
3019 let picked = pick_next_job(&mut buffer, &loaded, 3);
3020 assert_eq!(
3021 picked.request.model, "a",
3022 "front-most match wins (the first `a`), not the loaded model with the most copies later in the buffer"
3023 );
3024 assert_eq!(buffer.len(), 3);
3028 let remaining: Vec<&str> = buffer
3029 .iter()
3030 .map(|b| b.job.request.model.as_str())
3031 .collect();
3032 assert_eq!(remaining, vec!["b", "a", "b"]);
3033 assert_eq!(buffer.front().unwrap().deferred, 0);
3034 }
3035
3036 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3040 async fn queue_dispatcher_reorders_interleaved_jobs_to_minimize_swaps() {
3041 let (worker, worker_rx) = test_worker(0, 8);
3042 {
3045 let mut cache = worker.model_cache.lock().unwrap();
3046 struct Engine(&'static str);
3047 impl mold_inference::InferenceEngine for Engine {
3048 fn generate(
3049 &mut self,
3050 _r: &GenerateRequest,
3051 ) -> anyhow::Result<mold_core::GenerateResponse> {
3052 unimplemented!()
3053 }
3054 fn model_name(&self) -> &str {
3055 self.0
3056 }
3057 fn is_loaded(&self) -> bool {
3058 true
3059 }
3060 fn load(&mut self) -> anyhow::Result<()> {
3061 Ok(())
3062 }
3063 }
3064 cache.insert(Box::new(Engine("a")), 0);
3065 }
3066
3067 let (job_tx, job_rx) = tokio::sync::mpsc::channel(8);
3068 let queue = QueueHandle::new(job_tx.clone());
3069 let state = crate::state::AppState::empty(
3070 mold_core::Config::default(),
3071 queue.clone(),
3072 Arc::new(GpuPool {
3073 workers: vec![worker.clone()],
3074 }),
3075 8,
3076 );
3077
3078 let mut result_rxs = Vec::new();
3081 for model in ["a", "b", "a", "b"] {
3082 let (tx, rx) = tokio::sync::oneshot::channel();
3083 let job = crate::state::GenerationJob {
3084 id: String::new(),
3085 request: fake_request(model),
3086 completion_payload: SseCompletionPayload::Full,
3087 progress_tx: None,
3088 result_tx: tx,
3089 output_dir: None,
3090 };
3091 queue.submit(job, 8).await.unwrap();
3092 result_rxs.push(rx);
3093 }
3094
3095 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3096
3097 let mut order = Vec::new();
3098 for _ in 0..4 {
3099 let dispatched = worker_rx
3100 .recv_timeout(std::time::Duration::from_secs(2))
3101 .expect("worker should receive the dispatched job");
3102 order.push(dispatched.model);
3103 }
3104 drop(job_tx);
3105 dispatcher.abort();
3106
3107 assert_eq!(
3108 order,
3109 vec![
3110 "a".to_string(),
3111 "a".to_string(),
3112 "b".to_string(),
3113 "b".to_string(),
3114 ],
3115 "lookahead reorder should batch all `a` jobs together before swapping to `b`"
3116 );
3117 }
3118
3119 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3127 async fn queue_dispatcher_honors_registry_reorder_in_real_dispatch() {
3128 let (worker, worker_rx) = test_worker(0, 8);
3129 let (job_tx, job_rx) = tokio::sync::mpsc::channel(8);
3130 let queue = QueueHandle::new(job_tx.clone());
3131 let state = crate::state::AppState::empty(
3132 mold_core::Config::default(),
3133 queue.clone(),
3134 Arc::new(GpuPool {
3135 workers: vec![worker.clone()],
3136 }),
3137 8,
3138 );
3139
3140 state.queue_pause.pause();
3144 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3145
3146 let mut result_rxs = Vec::new();
3149 for id in ["a", "b", "c"] {
3150 state
3151 .job_registry
3152 .register_job(id, format!("model-{id}"), None, None, None);
3153 let (tx, rx) = tokio::sync::oneshot::channel();
3154 let job = crate::state::GenerationJob {
3155 id: id.to_string(),
3156 request: fake_request(&format!("model-{id}")),
3157 completion_payload: SseCompletionPayload::Full,
3158 progress_tx: None,
3159 result_tx: tx,
3160 output_dir: None,
3161 };
3162 queue.submit(job, 8).await.unwrap();
3163 result_rxs.push(rx);
3164 }
3165
3166 state.job_registry.reorder_queued("c", 0).unwrap();
3169
3170 state.queue_pause.resume();
3173
3174 let mut order = Vec::new();
3175 for _ in 0..3 {
3176 let dispatched = worker_rx
3177 .recv_timeout(std::time::Duration::from_secs(2))
3178 .expect("worker should receive the dispatched job");
3179 order.push(dispatched.model);
3180 }
3181 drop(job_tx);
3182 dispatcher.abort();
3183
3184 assert_eq!(
3185 order,
3186 vec![
3187 "model-c".to_string(),
3188 "model-a".to_string(),
3189 "model-b".to_string(),
3190 ],
3191 "registry reorder must drive real dispatch order, not just the snapshot"
3192 );
3193 }
3194
3195 #[tokio::test]
3202 async fn top_up_buffer_never_exceeds_capacity() {
3203 use std::collections::VecDeque;
3204 let (job_tx, mut job_rx) = tokio::sync::mpsc::channel::<GenerationJob>(32);
3205
3206 for i in 0..10 {
3209 let (tx, _rx) = tokio::sync::oneshot::channel();
3210 let job = GenerationJob {
3211 id: String::new(),
3212 request: fake_request(&format!("model-{i}")),
3213 completion_payload: SseCompletionPayload::Full,
3214 progress_tx: None,
3215 result_tx: tx,
3216 output_dir: None,
3217 };
3218 job_tx.send(job).await.unwrap();
3219 }
3220
3221 let mut buffer: VecDeque<BufferedJob> = VecDeque::with_capacity(4);
3223 top_up_buffer(&mut buffer, &mut job_rx, 4);
3224 assert_eq!(
3225 buffer.len(),
3226 4,
3227 "top_up_buffer must cap at buffer_size, leaving the rest in the channel"
3228 );
3229
3230 while buffer.pop_front().is_some() {}
3233 top_up_buffer(&mut buffer, &mut job_rx, 4);
3234 assert_eq!(buffer.len(), 4);
3235 let names: Vec<&str> = buffer
3236 .iter()
3237 .map(|b| b.job.request.model.as_str())
3238 .collect();
3239 assert_eq!(
3240 names,
3241 vec!["model-4", "model-5", "model-6", "model-7"],
3242 "second top-up must drain the next FIFO window from the channel"
3243 );
3244
3245 drop(job_tx);
3248 while buffer.pop_front().is_some() {}
3249 top_up_buffer(&mut buffer, &mut job_rx, 4);
3250 assert_eq!(
3251 buffer.len(),
3252 2,
3253 "top_up_buffer drains the channel tail when fewer jobs than capacity remain"
3254 );
3255 let names: Vec<&str> = buffer
3256 .iter()
3257 .map(|b| b.job.request.model.as_str())
3258 .collect();
3259 assert_eq!(names, vec!["model-8", "model-9"]);
3260 }
3261
3262 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3268 async fn queue_dispatcher_dispatches_all_jobs_when_submission_exceeds_buffer() {
3269 let (worker, worker_rx) = test_worker(0, 4);
3270 let (job_tx, job_rx) = tokio::sync::mpsc::channel(32);
3271 let queue = QueueHandle::new(job_tx.clone());
3272 let state = crate::state::AppState::empty(
3273 mold_core::Config::default(),
3274 queue.clone(),
3275 Arc::new(GpuPool {
3276 workers: vec![worker.clone()],
3277 }),
3278 32,
3279 );
3280
3281 let drain_worker = worker.clone();
3287 let drainer = std::thread::spawn(move || {
3288 let mut order = Vec::new();
3289 while order.len() < 10 {
3290 match worker_rx.recv_timeout(std::time::Duration::from_secs(5)) {
3291 Ok(j) => {
3292 drain_worker.in_flight.fetch_sub(1, Ordering::SeqCst);
3293 order.push(j.model);
3294 }
3295 Err(e) => panic!("drain stalled at {:?}: {e:?}", order),
3296 }
3297 }
3298 order
3299 });
3300
3301 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3302
3303 let mut held_rxs = Vec::new();
3309 for i in 0..10 {
3310 let (tx, rx) = tokio::sync::oneshot::channel();
3311 held_rxs.push(rx);
3312 let job = crate::state::GenerationJob {
3313 id: String::new(),
3314 request: fake_request(&format!("model-{i}")),
3315 completion_payload: SseCompletionPayload::Full,
3316 progress_tx: None,
3317 result_tx: tx,
3318 output_dir: None,
3319 };
3320 queue.submit(job, 32).await.unwrap();
3321 }
3322
3323 let order = drainer.join().expect("drainer thread panic");
3324 drop(job_tx);
3325 dispatcher.abort();
3326
3327 let expected: Vec<String> = (0..10).map(|i| format!("model-{i}")).collect();
3328 assert_eq!(
3329 order, expected,
3330 "10 distinct jobs must come out in FIFO across buffer rotations"
3331 );
3332 }
3333
3334 static QUEUE_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
3336
3337 fn with_queue_env<R>(name: &str, value: Option<&str>, f: impl FnOnce() -> R) -> R {
3338 let _g = QUEUE_ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
3339 let prev = std::env::var(name).ok();
3340 match value {
3341 Some(v) => std::env::set_var(name, v),
3342 None => std::env::remove_var(name),
3343 }
3344 let out = f();
3345 match prev {
3346 Some(v) => std::env::set_var(name, v),
3347 None => std::env::remove_var(name),
3348 }
3349 out
3350 }
3351
3352 #[test]
3353 fn resolve_lookahead_buffer_uses_default_when_env_missing() {
3354 let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, None, resolve_lookahead_buffer);
3355 assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3356 }
3357
3358 #[test]
3359 fn resolve_lookahead_buffer_honors_env_within_range() {
3360 let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, Some("4"), resolve_lookahead_buffer);
3361 assert_eq!(n, 4);
3362 }
3363
3364 #[test]
3365 fn resolve_lookahead_buffer_falls_back_when_out_of_range() {
3366 let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, Some("0"), resolve_lookahead_buffer);
3368 assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3369 let n = with_queue_env(LOOKAHEAD_BUFFER_ENV, Some("999"), resolve_lookahead_buffer);
3370 assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3371 }
3372
3373 #[test]
3374 fn resolve_lookahead_buffer_falls_back_when_unparseable() {
3375 let n = with_queue_env(
3376 LOOKAHEAD_BUFFER_ENV,
3377 Some("not-a-number"),
3378 resolve_lookahead_buffer,
3379 );
3380 assert_eq!(n, DEFAULT_LOOKAHEAD_BUFFER);
3381 }
3382
3383 #[test]
3384 fn resolve_max_deferrals_uses_default_when_env_missing() {
3385 let n = with_queue_env(MAX_DEFERRALS_ENV, None, resolve_max_deferrals);
3386 assert_eq!(n, DEFAULT_MAX_DEFERRALS);
3387 }
3388
3389 #[test]
3390 fn resolve_max_deferrals_honors_env_within_range() {
3391 let n = with_queue_env(MAX_DEFERRALS_ENV, Some("0"), resolve_max_deferrals);
3393 assert_eq!(n, 0);
3394 let n = with_queue_env(MAX_DEFERRALS_ENV, Some("32"), resolve_max_deferrals);
3395 assert_eq!(n, 32);
3396 let n = with_queue_env(MAX_DEFERRALS_ENV, Some("5"), resolve_max_deferrals);
3397 assert_eq!(n, 5);
3398 }
3399
3400 #[test]
3401 fn resolve_max_deferrals_falls_back_when_out_of_range() {
3402 let n = with_queue_env(MAX_DEFERRALS_ENV, Some("999"), resolve_max_deferrals);
3403 assert_eq!(n, DEFAULT_MAX_DEFERRALS);
3404 }
3405
3406 #[test]
3407 fn resolve_max_deferrals_falls_back_when_unparseable() {
3408 let n = with_queue_env(
3409 MAX_DEFERRALS_ENV,
3410 Some("not-a-number"),
3411 resolve_max_deferrals,
3412 );
3413 assert_eq!(n, DEFAULT_MAX_DEFERRALS);
3414 }
3415
3416 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3417 async fn queue_dispatcher_honors_explicit_placement_gpu() {
3418 let (worker0, rx0) = test_worker(0, 1);
3419 let (worker1, rx1) = test_worker(1, 1);
3420 let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3421 let queue = QueueHandle::new(job_tx.clone());
3422 let state = crate::state::AppState::empty(
3423 mold_core::Config::default(),
3424 queue.clone(),
3425 Arc::new(GpuPool {
3426 workers: vec![worker0, worker1],
3427 }),
3428 8,
3429 );
3430
3431 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state));
3432
3433 let mut request = fake_request("flux-dev:q4");
3434 request.placement = Some(mold_core::types::DevicePlacement {
3435 text_encoders: mold_core::types::DeviceRef::Auto,
3436 advanced: Some(mold_core::types::AdvancedPlacement {
3437 transformer: mold_core::types::DeviceRef::gpu(1),
3438 ..mold_core::types::AdvancedPlacement::default()
3439 }),
3440 });
3441
3442 let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3443 let job = crate::state::GenerationJob {
3444 id: String::new(),
3445 request,
3446 completion_payload: SseCompletionPayload::Full,
3447 progress_tx: None,
3448 result_tx,
3449 output_dir: None,
3450 };
3451 let _position = queue.submit(job, 8).await.unwrap();
3452
3453 let dispatched = rx1
3454 .recv_timeout(std::time::Duration::from_secs(1))
3455 .expect("explicit placement should route to gpu 1");
3456 assert_eq!(dispatched.model, "flux-dev:q4");
3457 assert!(rx0.try_recv().is_err(), "gpu 0 should not receive the job");
3458
3459 drop(job_tx);
3460 dispatcher.abort();
3461 }
3462
3463 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3464 async fn queue_dispatcher_records_auto_selected_gpu_before_worker_starts() {
3465 let (worker0, rx0) = test_worker(0, 1);
3466 let (worker1, rx1) = test_worker(1, 1);
3467 let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3468 let queue = QueueHandle::new(job_tx.clone());
3469 let state = crate::state::AppState::empty(
3470 mold_core::Config::default(),
3471 queue.clone(),
3472 Arc::new(GpuPool {
3473 workers: vec![worker0, worker1],
3474 }),
3475 8,
3476 );
3477 state.job_registry.register("auto-job", "flux-dev:q4");
3478
3479 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3480
3481 let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3482 let job = crate::state::GenerationJob {
3483 id: "auto-job".to_string(),
3484 request: fake_request("flux-dev:q4"),
3485 completion_payload: SseCompletionPayload::Full,
3486 progress_tx: None,
3487 result_tx,
3488 output_dir: None,
3489 };
3490 let _position = queue.submit(job, 8).await.unwrap();
3491
3492 let (dispatched, ordinal) = match rx0.recv_timeout(std::time::Duration::from_secs(1)) {
3493 Ok(job) => (job, 0),
3494 Err(_) => (
3495 rx1.recv_timeout(std::time::Duration::from_secs(1))
3496 .expect("auto job should dispatch to one GPU"),
3497 1,
3498 ),
3499 };
3500 assert_eq!(dispatched.model, "flux-dev:q4");
3501 assert_eq!(
3502 state.job_registry.target_gpu("auto-job"),
3503 Some(Some(ordinal))
3504 );
3505
3506 drop(job_tx);
3507 dispatcher.abort();
3508 }
3509
3510 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3511 async fn paused_dispatcher_holds_new_jobs_until_resumed() {
3512 let (worker0, rx0) = test_worker(0, 1);
3513 let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3514 let queue = QueueHandle::new(job_tx.clone());
3515 let state = crate::state::AppState::empty(
3516 mold_core::Config::default(),
3517 queue.clone(),
3518 Arc::new(GpuPool {
3519 workers: vec![worker0],
3520 }),
3521 8,
3522 );
3523
3524 assert!(state.queue_pause.pause());
3526 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3527
3528 let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3529 let job = crate::state::GenerationJob {
3530 id: "paused-job".to_string(),
3531 request: fake_request("flux-dev:q4"),
3532 completion_payload: SseCompletionPayload::Full,
3533 progress_tx: None,
3534 result_tx,
3535 output_dir: None,
3536 };
3537 let _position = queue.submit(job, 8).await.unwrap();
3538
3539 assert!(
3541 rx0.recv_timeout(std::time::Duration::from_millis(200))
3542 .is_err(),
3543 "paused dispatcher must not hand a job to a worker"
3544 );
3545
3546 assert!(state.queue_pause.resume());
3548 let dispatched = rx0
3549 .recv_timeout(std::time::Duration::from_secs(1))
3550 .expect("resumed dispatcher should dispatch the queued job");
3551 assert_eq!(dispatched.model, "flux-dev:q4");
3552
3553 drop(job_tx);
3554 dispatcher.abort();
3555 }
3556
3557 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
3558 async fn pause_while_dispatcher_is_parked_on_an_empty_queue_still_holds_the_next_job() {
3559 let (worker0, rx0) = test_worker(0, 1);
3564 let (job_tx, job_rx) = tokio::sync::mpsc::channel(4);
3565 let queue = QueueHandle::new(job_tx.clone());
3566 let state = crate::state::AppState::empty(
3567 mold_core::Config::default(),
3568 queue.clone(),
3569 Arc::new(GpuPool {
3570 workers: vec![worker0],
3571 }),
3572 8,
3573 );
3574
3575 let dispatcher = tokio::spawn(run_queue_dispatcher(job_rx, state.clone()));
3577 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
3578
3579 assert!(state.queue_pause.pause());
3581 let (result_tx, _result_rx) = tokio::sync::oneshot::channel();
3582 let job = crate::state::GenerationJob {
3583 id: "parked-job".to_string(),
3584 request: fake_request("flux-dev:q4"),
3585 completion_payload: SseCompletionPayload::Full,
3586 progress_tx: None,
3587 result_tx,
3588 output_dir: None,
3589 };
3590 let _position = queue.submit(job, 8).await.unwrap();
3591
3592 assert!(
3593 rx0.recv_timeout(std::time::Duration::from_millis(200))
3594 .is_err(),
3595 "a job arriving while paused must not wake straight into dispatch"
3596 );
3597
3598 assert!(state.queue_pause.resume());
3599 let dispatched = rx0
3600 .recv_timeout(std::time::Duration::from_secs(1))
3601 .expect("resume should release the held job");
3602 assert_eq!(dispatched.model, "flux-dev:q4");
3603
3604 drop(job_tx);
3605 dispatcher.abort();
3606 }
3607}
3608
3609#[cfg(test)]
3610mod queue_pause_tests {
3611 use super::QueuePause;
3612 use std::time::Duration;
3613
3614 #[test]
3615 fn pause_and_resume_report_state_transitions() {
3616 let gate = QueuePause::new();
3617 assert!(!gate.is_paused());
3618 assert!(gate.pause(), "first pause flips state");
3619 assert!(gate.is_paused());
3620 assert!(!gate.pause(), "second pause is a no-op transition");
3621 assert!(gate.resume(), "first resume flips state");
3622 assert!(!gate.is_paused());
3623 assert!(!gate.resume(), "second resume is a no-op transition");
3624 }
3625
3626 #[tokio::test]
3627 async fn wait_if_paused_returns_immediately_when_not_paused() {
3628 let gate = QueuePause::new();
3629 tokio::time::timeout(Duration::from_secs(1), gate.wait_if_paused())
3631 .await
3632 .expect("wait_if_paused must not block when the gate is open");
3633 }
3634
3635 #[tokio::test]
3636 async fn wait_if_paused_blocks_until_resumed() {
3637 let gate = QueuePause::new();
3638 assert!(gate.pause());
3639
3640 let waiter = {
3641 let gate = gate.clone();
3642 tokio::spawn(async move { gate.wait_if_paused().await })
3643 };
3644
3645 tokio::time::sleep(Duration::from_millis(50)).await;
3647 assert!(!waiter.is_finished(), "waiter must block while paused");
3648
3649 assert!(gate.resume());
3651 tokio::time::timeout(Duration::from_secs(1), waiter)
3652 .await
3653 .expect("waiter must unblock within the timeout after resume")
3654 .expect("waiter task must not panic");
3655 }
3656
3657 #[tokio::test]
3658 async fn resume_wakes_every_gated_waiter() {
3659 let gate = QueuePause::new();
3661 assert!(gate.pause());
3662
3663 let waiters: Vec<_> = (0..3)
3664 .map(|_| {
3665 let gate = gate.clone();
3666 tokio::spawn(async move { gate.wait_if_paused().await })
3667 })
3668 .collect();
3669
3670 tokio::time::sleep(Duration::from_millis(50)).await;
3671 assert!(gate.resume());
3672
3673 for waiter in waiters {
3674 tokio::time::timeout(Duration::from_secs(1), waiter)
3675 .await
3676 .expect("every gated waiter must wake on a single resume")
3677 .expect("waiter task must not panic");
3678 }
3679 }
3680}