use std::path::{Path, PathBuf};
pub const ALIASES: [(&str, &str, &str); 3] = [
("laya", "convaiinnovations/laya", ""),
("laya-multilingual", "convaiinnovations/laya", "multilingual"),
("laya-typed-decisions", "convaiinnovations/laya", "typed-decisions"),
];
pub const FILES: [&str; 5] = [
"rl_agent_config.json",
"encoder/config.json",
"tokenizer/tokenizer.json",
"tokenizer/tokenizer_config.json",
"model.safetensors",
];
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct HubRef {
pub repo: String,
pub subfolder: String,
}
impl HubRef {
#[must_use]
pub fn parse(name: &str) -> Option<HubRef> {
if let Some((_, repo, sub)) = ALIASES.iter().find(|a| a.0 == name) {
return Some(HubRef { repo: (*repo).into(), subfolder: (*sub).into() });
}
let rest = name.strip_prefix("hf://")?;
let mut parts = rest.splitn(3, '/');
let (org, repo) = (parts.next()?, parts.next()?);
if org.is_empty() || repo.is_empty() {
return None;
}
let subfolder = parts.next().unwrap_or("").trim_matches('/').to_string();
Some(HubRef { repo: format!("{org}/{repo}"), subfolder })
}
#[must_use]
pub fn repo_dir(&self, cache: &Path) -> PathBuf {
cache.join(format!("models--{}", self.repo.replace('/', "--")))
}
#[must_use]
pub fn local(&self, cache: &Path) -> Option<PathBuf> {
let repo = self.repo_dir(cache);
let snapshots = repo.join("snapshots");
let has = |snap: &Path| {
let dir = snap.join(&self.subfolder);
dir.join("model.safetensors").is_file().then_some(dir)
};
if let Ok(rev) = std::fs::read_to_string(repo.join("refs/main"))
&& let Some(dir) = has(&snapshots.join(rev.trim()))
{
return Some(dir);
}
std::fs::read_dir(&snapshots)
.ok()?
.filter_map(|e| {
let e = e.ok()?;
let dir = has(&e.path())?;
Some((e.metadata().and_then(|m| m.modified()).ok()?, dir))
})
.max_by_key(|(t, _)| *t)
.map(|(_, dir)| dir)
}
}
#[must_use]
pub fn cache_dir() -> PathBuf {
let var = |k: &str| std::env::var_os(k).filter(|v| !v.is_empty()).map(PathBuf::from);
if let Some(d) = var("HF_HUB_CACHE") {
return d;
}
if let Some(d) = var("HF_HOME") {
return d.join("hub");
}
let home = var("HOME").or_else(|| var("USERPROFILE")).unwrap_or_else(|| PathBuf::from("."));
home.join(".cache/huggingface/hub")
}
pub fn resolve(name: &str) -> Result<PathBuf, String> {
let path = Path::new(name);
if path.exists() {
return Ok(path.to_path_buf());
}
let Some(r) = HubRef::parse(name) else {
let known: Vec<&str> = ALIASES.iter().map(|a| a.0).collect();
return Err(format!(
"no model {name:?}: not a path, an hf:// reference or one of {}",
known.join(", ")
));
};
let cache = cache_dir();
r.local(&cache)
.ok_or_else(|| format!("{name} is not in {}, run kime pull {name} first", cache.display()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn names() {
let r = HubRef::parse("laya-multilingual").unwrap();
assert_eq!(
(r.repo.as_str(), r.subfolder.as_str()),
("convaiinnovations/laya", "multilingual")
);
let r = HubRef::parse("hf://org/repo/a/b/").unwrap();
assert_eq!((r.repo.as_str(), r.subfolder.as_str()), ("org/repo", "a/b"));
assert_eq!(HubRef::parse("hf://org"), None);
assert_eq!(HubRef::parse("kime-v1-s-en"), None);
assert_eq!(r.repo_dir(Path::new("/c")), Path::new("/c/models--org--repo"));
}
#[test]
fn subfolders_pulled_at_different_commits() {
let cache = std::env::temp_dir().join(format!("kime-hub-{}", std::process::id()));
let repo = cache.join("models--convaiinnovations--laya");
for (rev, sub) in [("old", ""), ("old", "multilingual"), ("new", "typed-decisions")] {
let dir = repo.join("snapshots").join(rev).join(sub);
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(dir.join("model.safetensors"), b"").unwrap();
}
std::fs::create_dir_all(repo.join("refs")).unwrap();
std::fs::write(repo.join("refs/main"), "new").unwrap();
let at = |name: &str| HubRef::parse(name).unwrap().local(&cache);
assert_eq!(at("laya-typed-decisions"), Some(repo.join("snapshots/new/typed-decisions")));
assert_eq!(at("laya-multilingual"), Some(repo.join("snapshots/old/multilingual")));
assert_eq!(at("laya"), Some(repo.join("snapshots/old/")));
let _ = std::fs::remove_dir_all(cache);
}
}