use std::process::Command;
use crate::tasks::generate_image::{GenerateImageRequest, GenerateImageResult};
use crate::InferenceError;
const CLI_BINARY: &str = "mflux-generate";
const DEFAULT_MODEL: &str = "mlx-community/Flux-1.lite-8B-MLX-Q4";
const DEFAULT_BASE_MODEL: &str = "dev";
const NATIVE_SERVED_MODELS: &[&str] = &["mlx-community/Flux-1.lite-8B-MLX-Q4"];
pub fn native_backend_serves(model: Option<&str>) -> bool {
match model {
None => true,
Some(m) => NATIVE_SERVED_MODELS.contains(&m),
}
}
const BUILTIN_BASE_MODELS: &[&str] = &[
"dev",
"schnell",
"krea-dev",
"dev-krea",
"qwen",
"fibo",
"fibo-lite",
"fibo-edit",
"fibo-edit-rmbg",
"z-image",
"z-image-turbo",
"flux2-klein-4b",
"flux2-klein-9b",
"flux2-klein-base-4b",
"flux2-klein-base-9b",
];
fn is_builtin_base_model(model: &str) -> bool {
BUILTIN_BASE_MODELS.contains(&model)
}
fn cli_path() -> std::path::PathBuf {
if let Some(managed) = managed_cli() {
if runs(&managed) {
return managed;
}
}
std::path::PathBuf::from(CLI_BINARY)
}
fn managed_cli() -> Option<std::path::PathBuf> {
let home = std::env::var_os("HOME")?;
let p = std::path::Path::new(&home)
.join(".car")
.join("visual-runtime")
.join("bin")
.join(CLI_BINARY);
p.is_file().then_some(p)
}
fn runs(bin: &std::path::Path) -> bool {
Command::new(bin)
.arg("--help")
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.status()
.map(|s| s.success())
.unwrap_or(false)
}
pub fn is_available() -> bool {
Command::new(cli_path())
.arg("--help")
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.status()
.map(|s| s.success())
.unwrap_or(false)
}
pub fn generate_image(req: &GenerateImageRequest) -> Result<GenerateImageResult, InferenceError> {
let output_path = req
.output_path
.clone()
.unwrap_or_else(|| "output.png".to_string());
let mut cmd = Command::new(cli_path());
let model = req.model.as_deref().unwrap_or(DEFAULT_MODEL);
if is_builtin_base_model(model) {
cmd.arg("--base-model").arg(model);
} else {
cmd.arg("--base-model")
.arg(DEFAULT_BASE_MODEL)
.arg("--model")
.arg(model);
}
cmd.arg("--prompt")
.arg(&req.prompt)
.arg("--output")
.arg(&output_path);
if let Some(w) = req.width {
cmd.arg("--width").arg(w.to_string());
}
if let Some(h) = req.height {
cmd.arg("--height").arg(h.to_string());
}
if let Some(s) = req.steps {
cmd.arg("--steps").arg(s.to_string());
}
if let Some(g) = req.guidance {
cmd.arg("--guidance").arg(g.to_string());
}
if let Some(seed) = req.seed {
cmd.arg("--seed").arg(seed.to_string());
}
tracing::info!(prompt = %req.prompt, output = %output_path, "external mflux: invoking");
let output = cmd.output().map_err(|e| {
InferenceError::InferenceFailed(format!(
"failed to spawn `{CLI_BINARY}`: {e}. \
Install with `uv pip install mflux` and put its venv's bin on PATH."
))
})?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
return Err(InferenceError::InferenceFailed(format!(
"mflux-generate exited with status {}: stderr={stderr}",
output.status
)));
}
Ok(GenerateImageResult {
image_path: output_path,
media_type: "image/png".to_string(),
model_used: Some(format!(
"external:{}",
req.model.as_deref().unwrap_or(DEFAULT_MODEL)
)),
})
}
#[cfg(test)]
mod model_selection_tests {
use super::*;
#[test]
fn mflux_builtin_identifiers_are_recognized() {
for m in [
"flux2-klein-4b",
"flux2-klein-9b",
"flux2-klein-base-4b",
"flux2-klein-base-9b",
"krea-dev",
"qwen",
"z-image-turbo",
"schnell",
"dev",
] {
assert!(
is_builtin_base_model(m),
"`{m}` is an mflux base model and must be passed as --base-model"
);
}
}
#[test]
fn huggingface_checkpoints_are_not_builtins() {
for m in [
"mlx-community/Flux-1.lite-8B-MLX-Q4",
"AITRADER/FLUX2-klein-4B-mlx-4bit",
"some-org/some-model",
] {
assert!(!is_builtin_base_model(m), "`{m}` is a checkpoint path");
}
assert!(!is_builtin_base_model(DEFAULT_MODEL));
}
#[test]
fn cli_path_prefers_the_managed_runtime_when_present() {
let p = cli_path();
if p.is_absolute() {
assert!(
p.ends_with(std::path::Path::new("visual-runtime/bin").join(CLI_BINARY)),
"an absolute resolution must be the managed runtime: {}",
p.display()
);
} else {
assert_eq!(p, std::path::PathBuf::from(CLI_BINARY));
}
}
}
#[cfg(test)]
mod dispatch_tests {
use super::*;
#[test]
fn architectures_without_a_rust_backend_route_external() {
for m in [
"flux2-klein-4b",
"flux2-klein-9b",
"krea-dev",
"qwen",
"z-image-turbo",
"AITRADER/FLUX2-klein-4B-mlx-4bit",
] {
assert!(
!native_backend_serves(Some(m)),
"`{m}` has no Rust implementation and must route to mflux"
);
}
}
#[test]
fn the_implemented_checkpoint_stays_native() {
assert!(native_backend_serves(Some(DEFAULT_MODEL)));
assert!(
native_backend_serves(None),
"CAR's default model is the one the native backend implements"
);
}
#[test]
fn dispatch_is_independent_of_the_environment() {
let before = (
native_backend_serves(None),
native_backend_serves(Some("flux2-klein-4b")),
);
for key in ["CAR_IMAGE_BACKEND", "CAR_VIDEO_BACKEND"] {
unsafe { std::env::set_var(key, "external") };
}
let after = (
native_backend_serves(None),
native_backend_serves(Some("flux2-klein-4b")),
);
for key in ["CAR_IMAGE_BACKEND", "CAR_VIDEO_BACKEND"] {
unsafe { std::env::remove_var(key) };
}
assert_eq!(
before, after,
"backend choice must not read the environment"
);
}
}