use anyhow::{Context, Result, bail};
use serde::Deserialize;
use std::path::{Path, PathBuf};
use super::ModelVariant;
pub const MANIFEST_FILE: &str = "manifest.toml";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelManifest {
pub architecture: ModelVariant,
pub files: ManifestFiles,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ManifestFiles {
pub encoder: String,
pub encoder_int8: Option<String>,
pub decoder: Option<String>,
pub joint: Option<String>,
pub vocab: String,
}
#[derive(Debug, Deserialize)]
struct RawManifest {
architecture: String,
files: RawFiles,
}
#[derive(Debug, Deserialize)]
struct RawFiles {
encoder: String,
#[serde(default)]
encoder_int8: Option<String>,
#[serde(default)]
decoder: Option<String>,
#[serde(default)]
joint: Option<String>,
vocab: String,
}
impl ModelManifest {
pub fn load(dir: &Path) -> Result<Option<Self>> {
let path = dir.join(MANIFEST_FILE);
if !path.is_file() {
return Ok(None);
}
let text = std::fs::read_to_string(&path)
.with_context(|| format!("failed to read model manifest {}", path.display()))?;
Self::parse(&text)
.with_context(|| format!("invalid model manifest {}", path.display()))
.map(Some)
}
pub fn parse(text: &str) -> Result<Self> {
let raw: RawManifest =
toml::from_str(text).context("failed to parse model manifest TOML")?;
Self::from_raw(raw)
}
fn from_raw(raw: RawManifest) -> Result<Self> {
let architecture: ModelVariant = raw
.architecture
.parse()
.map_err(|e: String| anyhow::anyhow!("invalid architecture: {e}"))?;
let encoder = normalize_required_basename("encoder", &raw.files.encoder)?;
let vocab = normalize_required_basename("vocab", &raw.files.vocab)?;
let encoder_int8 = normalize_optional_basename("encoder_int8", raw.files.encoder_int8)?;
let decoder = normalize_optional_basename("decoder", raw.files.decoder)?;
let joint = normalize_optional_basename("joint", raw.files.joint)?;
if !architecture.is_ctc() {
if decoder.is_none() {
bail!(
"manifest files.decoder is required for architecture '{}'",
architecture.as_str()
);
}
if joint.is_none() {
bail!(
"manifest files.joint is required for architecture '{}'",
architecture.as_str()
);
}
}
Ok(Self {
architecture,
files: ManifestFiles {
encoder,
encoder_int8,
decoder,
joint,
vocab,
},
})
}
pub fn preferred_encoder_path(&self, dir: &Path) -> PathBuf {
if let Some(ref int8_name) = self.files.encoder_int8 {
let int8 = dir.join(int8_name);
if int8.exists() {
return int8;
}
}
dir.join(&self.files.encoder)
}
pub fn prefers_int8(&self, dir: &Path) -> bool {
self.files
.encoder_int8
.as_ref()
.is_some_and(|name| dir.join(name).exists())
}
pub fn decoder_path(&self, dir: &Path) -> Option<PathBuf> {
self.files.decoder.as_ref().map(|name| dir.join(name))
}
pub fn joint_path(&self, dir: &Path) -> Option<PathBuf> {
self.files.joint.as_ref().map(|name| dir.join(name))
}
pub fn vocab_path(&self, dir: &Path) -> PathBuf {
dir.join(&self.files.vocab)
}
}
fn normalize_required_basename(field: &str, value: &str) -> Result<String> {
let trimmed = value.trim();
if trimmed.is_empty() {
bail!("manifest files.{field} must be a non-empty basename");
}
validate_basename(field, trimmed)?;
Ok(trimmed.to_string())
}
fn normalize_optional_basename(field: &str, value: Option<String>) -> Result<Option<String>> {
let Some(value) = value else {
return Ok(None);
};
let trimmed = value.trim();
if trimmed.is_empty() {
return Ok(None);
}
validate_basename(field, trimmed)?;
Ok(Some(trimmed.to_string()))
}
fn validate_basename(field: &str, value: &str) -> Result<()> {
if value.contains('/') || value.contains('\\') || value.contains("..") {
bail!(
"manifest files.{field} must be a basename (got '{value}'); \
paths and '..' are not allowed"
);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
#[test]
fn test_parse_valid_rnnt_manifest() {
let text = r#"
architecture = "rnnt"
[files]
encoder = "v3_rnnt_encoder.onnx"
encoder_int8 = "v3_rnnt_encoder_int8.onnx"
decoder = "v3_rnnt_decoder.onnx"
joint = "v3_rnnt_joint.onnx"
vocab = "v3_vocab.txt"
"#;
let m = ModelManifest::parse(text).expect("valid manifest");
assert_eq!(m.architecture, ModelVariant::Rnnt);
assert_eq!(m.files.encoder, "v3_rnnt_encoder.onnx");
assert_eq!(
m.files.encoder_int8.as_deref(),
Some("v3_rnnt_encoder_int8.onnx")
);
assert_eq!(m.files.decoder.as_deref(), Some("v3_rnnt_decoder.onnx"));
assert_eq!(m.files.joint.as_deref(), Some("v3_rnnt_joint.onnx"));
assert_eq!(m.files.vocab, "v3_vocab.txt");
}
#[test]
fn test_parse_valid_ctc_manifest_without_decoder_joint() {
let text = r#"
architecture = "ml_ctc"
[files]
encoder = "multilingual_ctc.onnx"
encoder_int8 = "multilingual_ctc.int8.onnx"
vocab = "multilingual_vocab.txt"
"#;
let m = ModelManifest::parse(text).expect("ctc manifest");
assert_eq!(m.architecture, ModelVariant::MlCtc);
assert!(m.files.decoder.is_none());
assert!(m.files.joint.is_none());
}
#[test]
fn test_parse_empty_decoder_joint_treated_as_absent() {
let text = r#"
architecture = "ml_ctc_large"
[files]
encoder = "multilingual_large_ctc.onnx"
decoder = ""
joint = ""
vocab = "multilingual_vocab.txt"
"#;
let m = ModelManifest::parse(text).expect("empty decoder/joint ok for ctc");
assert!(m.files.decoder.is_none());
assert!(m.files.joint.is_none());
}
#[test]
fn test_load_missing_file_returns_none() {
let dir = tempfile::tempdir().expect("tempdir");
let loaded = ModelManifest::load(dir.path()).expect("missing is ok");
assert!(loaded.is_none());
}
#[test]
fn test_load_valid_manifest_from_disk() {
let dir = tempfile::tempdir().expect("tempdir");
fs::write(
dir.path().join(MANIFEST_FILE),
r#"
architecture = "e2e_rnnt"
[files]
encoder = "custom_encoder.onnx"
decoder = "custom_decoder.onnx"
joint = "custom_joint.onnx"
vocab = "custom_vocab.txt"
"#,
)
.unwrap();
let m = ModelManifest::load(dir.path())
.expect("load")
.expect("present");
assert_eq!(m.architecture, ModelVariant::E2eRnnt);
assert_eq!(m.files.encoder, "custom_encoder.onnx");
}
#[test]
fn test_invalid_architecture_rejected() {
let text = r#"
architecture = "whisper"
[files]
encoder = "e.onnx"
decoder = "d.onnx"
joint = "j.onnx"
vocab = "v.txt"
"#;
let err = ModelManifest::parse(text).expect_err("whisper must fail");
let msg = format!("{err:#}");
assert!(
msg.contains("architecture") || msg.contains("whisper"),
"error should mention architecture: {msg}"
);
}
#[test]
fn test_rnnt_requires_decoder_and_joint() {
let text = r#"
architecture = "rnnt"
[files]
encoder = "e.onnx"
vocab = "v.txt"
"#;
let err = ModelManifest::parse(text).expect_err("decoder required");
assert!(
format!("{err:#}").contains("decoder"),
"expected decoder error, got {err:#}"
);
}
#[test]
fn test_reject_path_separators_in_basenames() {
let text = r#"
architecture = "rnnt"
[files]
encoder = "../escape.onnx"
decoder = "d.onnx"
joint = "j.onnx"
vocab = "v.txt"
"#;
let err = ModelManifest::parse(text).expect_err("path sep rejected");
assert!(
format!("{err:#}").contains("basename"),
"expected basename error, got {err:#}"
);
}
#[test]
fn test_preferred_encoder_path_prefers_int8_when_present() {
let dir = tempfile::tempdir().expect("tempdir");
fs::write(dir.path().join("enc.onnx"), b"fp32").unwrap();
fs::write(dir.path().join("enc_int8.onnx"), b"int8").unwrap();
let m = ModelManifest {
architecture: ModelVariant::Rnnt,
files: ManifestFiles {
encoder: "enc.onnx".into(),
encoder_int8: Some("enc_int8.onnx".into()),
decoder: Some("d.onnx".into()),
joint: Some("j.onnx".into()),
vocab: "v.txt".into(),
},
};
assert_eq!(
m.preferred_encoder_path(dir.path()).file_name().unwrap(),
"enc_int8.onnx"
);
assert!(m.prefers_int8(dir.path()));
}
#[test]
fn test_preferred_encoder_path_without_int8_file_is_not_int8() {
let dir = tempfile::tempdir().expect("tempdir");
fs::write(dir.path().join("enc.onnx"), b"fp32").unwrap();
let m = ModelManifest {
architecture: ModelVariant::Rnnt,
files: ManifestFiles {
encoder: "enc.onnx".into(),
encoder_int8: Some("enc_int8.onnx".into()),
decoder: Some("d.onnx".into()),
joint: Some("j.onnx".into()),
vocab: "v.txt".into(),
},
};
assert!(!m.prefers_int8(dir.path()));
assert_eq!(
m.preferred_encoder_path(dir.path()).file_name().unwrap(),
"enc.onnx"
);
}
#[test]
fn test_load_invalid_manifest_returns_error() {
let dir = tempfile::tempdir().expect("tempdir");
fs::write(dir.path().join(MANIFEST_FILE), "architecture = [\n").unwrap();
let err = ModelManifest::load(dir.path()).expect_err("invalid toml");
assert!(
format!("{err:#}").contains("manifest"),
"expected manifest context, got {err:#}"
);
}
}