use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use ffai_core::error::{Error, Result};
use serde::Deserialize;
use sha2::{Digest, Sha256};
#[derive(Debug, Clone)]
pub struct ResolvedModel {
pub name: String,
pub license: String,
pub files: BTreeMap<String, PathBuf>,
}
impl ResolvedModel {
pub fn file(&self, name: &str) -> Result<&Path> {
self.files
.get(name)
.map(PathBuf::as_path)
.ok_or_else(|| Error::Model(format!("model `{}` has no file `{name}`", self.name)))
}
}
fn hub_download(repo: &str, filename: &str, cache_only: bool) -> Result<PathBuf> {
let (owner, name) = repo
.split_once('/')
.ok_or_else(|| Error::Model(format!("hf_repo `{repo}` is not in `owner/name` form")))?;
let client = hf_hub::HFClientSync::new()
.map_err(|e| Error::Model(format!("hugging face client init failed: {e}")))?;
client
.model(owner, name)
.download_file()
.filename(filename)
.local_files_only(cache_only)
.send()
.map_err(|e| Error::Model(format!("{repo}/{filename}: {e}")))
}
fn verify_checksum(path: &Path, file: &ModelFile) -> Result<()> {
let Some(expected) = &file.sha256 else {
return Ok(());
};
let bytes = std::fs::read(path)?;
let actual: String = Sha256::digest(&bytes).iter().map(|b| format!("{b:02x}")).collect();
if actual != expected.to_ascii_lowercase() {
return Err(Error::Model(format!(
"checksum mismatch for {}: manifest says {expected}, file is {actual}",
path.display()
)));
}
Ok(())
}
#[derive(Debug, Clone, Deserialize)]
pub struct ModelFile {
pub name: String,
pub sha256: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct ModelManifest {
pub name: String,
pub task: String,
pub description: Option<String>,
pub license: String,
pub hf_repo: Option<String>,
#[serde(default)]
pub files: Vec<ModelFile>,
}
impl ModelManifest {
pub fn from_toml(text: &str) -> Result<Self> {
toml::from_str(text).map_err(|e| Error::Model(format!("bad manifest: {e}")))
}
pub fn load(path: &Path) -> Result<Self> {
let text = std::fs::read_to_string(path)?;
Self::from_toml(&text)
}
pub fn cache_path(&self) -> PathBuf {
cache_dir().join("models").join(&self.name)
}
pub fn is_cached(&self) -> bool {
!self.files.is_empty() && self.files.iter().all(|f| self.local_path(&f.name).is_some())
}
pub fn local_path(&self, name: &str) -> Option<PathBuf> {
let manual = self.cache_path().join(name);
if manual.exists() {
return Some(manual);
}
hub_download(self.hf_repo.as_ref()?, name, true).ok()
}
pub fn fetch(&self) -> Result<ResolvedModel> {
let mut files = BTreeMap::new();
for file in &self.files {
if let Some(path) = self.local_path(&file.name) {
verify_checksum(&path, file)?;
files.insert(file.name.clone(), path);
continue;
}
let repo = self.hf_repo.as_ref().ok_or_else(|| {
Error::Model(format!(
"model `{}` declares no hf_repo and `{}` is not present under {}",
self.name,
file.name,
self.cache_path().display()
))
})?;
let path = hub_download(repo, &file.name, false)?;
verify_checksum(&path, file)?;
files.insert(file.name.clone(), path);
}
Ok(ResolvedModel { name: self.name.clone(), license: self.license.clone(), files })
}
}
pub fn load_dir(dir: &Path) -> Result<Vec<ModelManifest>> {
let mut out = Vec::new();
for entry in std::fs::read_dir(dir)? {
let path = entry?.path();
if path.extension().and_then(|e| e.to_str()) == Some("toml") {
out.push(ModelManifest::load(&path)?);
}
}
out.sort_by(|a, b| a.task.cmp(&b.task).then_with(|| a.name.cmp(&b.name)));
Ok(out)
}
pub fn cache_dir() -> PathBuf {
if let Ok(dir) = std::env::var("FFAI_CACHE") {
return PathBuf::from(dir);
}
dirs::cache_dir()
.unwrap_or_else(std::env::temp_dir)
.join("ffai")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn manifest_parses_and_surfaces_license() {
let m = ModelManifest::from_toml(
r#"
name = "whisper-tiny"
task = "asr"
license = "Apache-2.0"
hf_repo = "openai/whisper-tiny"
[[files]]
name = "model.safetensors"
"#,
)
.unwrap();
assert_eq!(m.name, "whisper-tiny");
assert_eq!(m.license, "Apache-2.0");
assert_eq!(m.files.len(), 1);
assert!(!m.is_cached());
}
}