use std::collections::BTreeMap;
use std::fs::File;
use std::io::Read;
use std::path::Path;
use serde_json::Value;
pub fn single_file_bundles_vae(path: &Path) -> std::io::Result<bool> {
let mut file = File::open(path)?;
let mut len_buf = [0u8; 8];
file.read_exact(&mut len_buf)?;
let header_len = u64::from_le_bytes(len_buf) as usize;
let mut header_buf = vec![0u8; header_len];
file.read_exact(&mut header_buf)?;
let header: BTreeMap<String, Value> = serde_json::from_slice(&header_buf).map_err(|e| {
std::io::Error::other(format!(
"parse safetensors header at {}: {e}",
path.display()
))
})?;
Ok(header.keys().any(|k| {
k != "__metadata__"
&& (k.starts_with("encoder.conv_in")
|| k.starts_with("first_stage_model.encoder.")
|| k.starts_with("vae.encoder."))
}))
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::PathBuf;
fn temp_safetensors(name: &str) -> PathBuf {
let mut path = std::env::temp_dir();
path.push(format!(
"mold-probe-{}-{}-{}.safetensors",
name,
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos(),
));
path
}
fn write_safetensors_with_keys(path: &Path, keys: &[&str]) {
use std::io::Write;
let mut header = serde_json::Map::new();
for key in keys {
header.insert(
(*key).to_string(),
serde_json::json!({
"dtype": "F32",
"shape": [1],
"data_offsets": [0, 4],
}),
);
}
let header_json = serde_json::to_vec(&serde_json::Value::Object(header)).unwrap();
let mut f = File::create(path).expect("create fixture");
f.write_all(&(header_json.len() as u64).to_le_bytes())
.unwrap();
f.write_all(&header_json).unwrap();
f.write_all(&[0u8; 4]).unwrap(); }
#[test]
fn true_for_bundled_diffusers_prefix() {
let path = temp_safetensors("flux-vae-diffusers");
write_safetensors_with_keys(
&path,
&[
"double_blocks.0.img_attn.proj.weight",
"encoder.conv_in.weight",
"decoder.conv_out.weight",
],
);
let bundled = single_file_bundles_vae(&path).expect("probe must not error");
assert!(
bundled,
"diffusers-style `encoder.conv_in.weight` must mark the file as VAE-bundled"
);
let _ = std::fs::remove_file(path);
}
#[test]
fn false_for_unet_only() {
let path = temp_safetensors("flux-unet-only");
write_safetensors_with_keys(
&path,
&[
"double_blocks.0.img_attn.proj.weight",
"double_blocks.0.img_attn.norm.query_norm.scale",
"single_blocks.0.linear1.weight",
"img_in.weight",
"txt_in.weight",
"final_layer.linear.weight",
],
);
let bundled = single_file_bundles_vae(&path).expect("probe must not error");
assert!(
!bundled,
"transformer-only checkpoint (no encoder.conv_in / first_stage_model / vae prefix) \
must NOT be marked as VAE-bundled"
);
let _ = std::fs::remove_file(path);
}
#[test]
fn handles_a1111_prefix() {
let path = temp_safetensors("flux-vae-a1111");
write_safetensors_with_keys(
&path,
&[
"model.diffusion_model.double_blocks.0.img_attn.proj.weight",
"first_stage_model.encoder.conv_in.weight",
"first_stage_model.decoder.conv_out.weight",
],
);
let bundled = single_file_bundles_vae(&path).expect("probe must not error");
assert!(
bundled,
"A1111 `first_stage_model.encoder.*` prefix must mark the file as VAE-bundled"
);
let _ = std::fs::remove_file(path);
}
#[test]
fn handles_pruner_prefix() {
let path = temp_safetensors("flux-vae-pruned");
write_safetensors_with_keys(
&path,
&[
"double_blocks.0.img_attn.proj.weight",
"vae.encoder.conv_in.weight",
],
);
let bundled = single_file_bundles_vae(&path).expect("probe must not error");
assert!(
bundled,
"pruner-style `vae.encoder.*` prefix must mark the file as VAE-bundled"
);
let _ = std::fs::remove_file(path);
}
#[test]
fn io_error_on_missing_file() {
let missing = std::env::temp_dir().join("mold-probe-flux-vae-missing.safetensors");
let _ = std::fs::remove_file(&missing); let err = single_file_bundles_vae(&missing).unwrap_err();
assert_eq!(err.kind(), std::io::ErrorKind::NotFound);
}
}