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