use anyhow::Result;
use std::path::PathBuf;
#[cfg(any(feature = "gcs", feature = "tls-rustls"))]
pub(crate) fn ensure_crypto_provider() -> Result<()> {
if rustls::crypto::CryptoProvider::get_default().is_some() {
return Ok(());
}
match rustls::crypto::ring::default_provider().install_default() {
Ok(()) => Ok(()),
Err(_) if rustls::crypto::CryptoProvider::get_default().is_some() => Ok(()),
Err(_) => anyhow::bail!("Failed to install rustls ring CryptoProvider"),
}
}
#[cfg(not(any(feature = "gcs", feature = "tls-rustls")))]
pub(crate) fn ensure_crypto_provider() -> Result<()> {
Ok(())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelDownloadOutcome {
pub path: PathBuf,
pub resolved_revision: Option<String>,
}
pub fn reject_unsupported_revision(provider_name: &str, revision: Option<&str>) -> Result<()> {
match revision {
Some(revision) => anyhow::bail!(
"Provider '{provider_name}' does not support pinned revisions (requested '{revision}')"
),
None => Ok(()),
}
}
#[async_trait::async_trait]
pub trait ModelProviderTrait: Send + Sync {
async fn download_model(
&self,
model_name: &str,
cache_path: Option<PathBuf>,
ignore_weights: bool,
) -> Result<PathBuf>;
async fn download_model_revision(
&self,
model_name: &str,
cache_path: Option<PathBuf>,
ignore_weights: bool,
revision: Option<&str>,
) -> Result<ModelDownloadOutcome> {
reject_unsupported_revision(self.provider_name(), revision)?;
let path = self
.download_model(model_name, cache_path, ignore_weights)
.await?;
Ok(ModelDownloadOutcome {
path,
resolved_revision: None,
})
}
fn supports_revisions(&self) -> bool {
false
}
async fn resolve_revision(
&self,
_model_name: &str,
_cache_dir: Option<PathBuf>,
revision: Option<&str>,
) -> Result<Option<String>> {
reject_unsupported_revision(self.provider_name(), revision)?;
Ok(None)
}
async fn record_local_revision(
&self,
_model_name: &str,
_cache_dir: &std::path::Path,
_requested_revision: Option<&str>,
_commit: &str,
) -> Result<()> {
Ok(())
}
async fn delete_model(&self, model_name: &str, cache_dir: PathBuf) -> Result<()>;
async fn delete_model_revision(
&self,
model_name: &str,
cache_dir: PathBuf,
revision: Option<&str>,
) -> Result<()> {
reject_unsupported_revision(self.provider_name(), revision)?;
self.delete_model(model_name, cache_dir).await
}
async fn get_model_path(&self, model_name: &str, cache_dir: PathBuf) -> Result<PathBuf>;
async fn get_model_path_revision(
&self,
model_name: &str,
cache_dir: PathBuf,
revision: Option<&str>,
) -> Result<PathBuf> {
reject_unsupported_revision(self.provider_name(), revision)?;
self.get_model_path(model_name, cache_dir).await
}
fn canonical_model_name(&self, model_name: &str) -> Result<String> {
Ok(model_name.to_string())
}
fn provider_name(&self) -> &'static str;
fn is_ignored(filename: &str) -> bool
where
Self: Sized,
{
const DEFAULT_IGNORED: [&str; 1] = ["README.md"];
let name = std::path::Path::new(filename)
.file_name()
.and_then(|s| s.to_str())
.unwrap_or(filename);
name.starts_with('.') || DEFAULT_IGNORED.contains(&name)
}
fn is_image(path: &std::path::Path) -> bool
where
Self: Sized,
{
path.extension().is_some_and(|ext| {
ext.eq_ignore_ascii_case("png")
|| ext.eq_ignore_ascii_case("jpg")
|| ext.eq_ignore_ascii_case("jpeg")
|| ext.eq_ignore_ascii_case("gif")
|| ext.eq_ignore_ascii_case("webp")
|| ext.eq_ignore_ascii_case("svg")
|| ext.eq_ignore_ascii_case("ico")
|| ext.eq_ignore_ascii_case("bmp")
|| ext.eq_ignore_ascii_case("tiff")
|| ext.eq_ignore_ascii_case("tif")
})
}
fn is_weight_file(filename: &str) -> bool
where
Self: Sized,
{
is_weight_file(filename)
}
}
pub fn is_weight_file(filename: &str) -> bool {
filename.ends_with(".bin")
|| filename.ends_with(".safetensors")
|| filename.ends_with(".h5")
|| filename.ends_with(".msgpack")
|| filename.ends_with(".ckpt.index")
|| filename.ends_with(".iop")
|| filename.ends_with(".gas")
}
#[cfg(feature = "gcs")]
pub mod gcs;
pub mod huggingface;
pub(crate) mod lock_file;
pub mod ngc;
pub mod s3;
pub use gcs::GcsProvider;
pub use huggingface::HuggingFaceProvider;
pub use ngc::NgcProvider;
pub use s3::S3Provider;
#[cfg(not(feature = "gcs"))]
pub mod gcs {
use super::ModelProviderTrait;
use crate::cache::{ModelInfo, ProviderCache};
use anyhow::Result;
use std::path::{Path, PathBuf};
const FEATURE_DISABLED: &str = "GCS support is disabled; rebuild with the `gcs` feature";
pub struct GcsProvider;
pub struct GcsProviderCache;
#[async_trait::async_trait]
impl ModelProviderTrait for GcsProvider {
async fn download_model(
&self,
_model_name: &str,
_cache_dir: Option<PathBuf>,
_ignore_weights: bool,
) -> Result<PathBuf> {
anyhow::bail!(FEATURE_DISABLED)
}
async fn delete_model(&self, _model_name: &str, _cache_dir: PathBuf) -> Result<()> {
anyhow::bail!(FEATURE_DISABLED)
}
async fn get_model_path(&self, _model_name: &str, _cache_dir: PathBuf) -> Result<PathBuf> {
anyhow::bail!(FEATURE_DISABLED)
}
fn provider_name(&self) -> &'static str {
"GCS"
}
}
impl ProviderCache for GcsProviderCache {
fn clear_model(&self, _cache_root: &Path, _model_name: &str) -> Result<()> {
anyhow::bail!(FEATURE_DISABLED)
}
fn resolve_model_path(
&self,
_cache_root: &Path,
_model_name: &str,
_revision: Option<&str>,
) -> Result<PathBuf> {
anyhow::bail!(FEATURE_DISABLED)
}
fn list_models(&self, _cache_root: &Path) -> Result<Vec<ModelInfo>> {
Ok(Vec::new())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::Path;
#[test]
fn test_is_image_function() {
assert!(HuggingFaceProvider::is_image(Path::new("test.png")));
assert!(HuggingFaceProvider::is_image(Path::new("test.PNG")));
assert!(HuggingFaceProvider::is_image(Path::new("test.jpg")));
assert!(HuggingFaceProvider::is_image(Path::new("test.JPG")));
assert!(HuggingFaceProvider::is_image(Path::new("test.jpeg")));
assert!(HuggingFaceProvider::is_image(Path::new("test.JPEG")));
assert!(!HuggingFaceProvider::is_image(Path::new("test.txt")));
assert!(!HuggingFaceProvider::is_image(Path::new("test.py")));
assert!(!HuggingFaceProvider::is_image(Path::new("test")));
assert!(!HuggingFaceProvider::is_image(Path::new("test.model")));
}
#[test]
fn test_ignored_files() {
assert!(HuggingFaceProvider::is_ignored(".gitattributes"));
assert!(HuggingFaceProvider::is_ignored(".gitignore"));
assert!(HuggingFaceProvider::is_ignored(".gitkeep"));
assert!(HuggingFaceProvider::is_ignored(".hidden"));
assert!(HuggingFaceProvider::is_ignored("subdir/.gitkeep"));
assert!(HuggingFaceProvider::is_ignored("a/b/.hidden"));
assert!(HuggingFaceProvider::is_ignored("README.md"));
assert!(HuggingFaceProvider::is_ignored("subdir/README.md"));
assert!(!HuggingFaceProvider::is_ignored("model.bin"));
assert!(!HuggingFaceProvider::is_ignored("tokenizer.json"));
assert!(!HuggingFaceProvider::is_ignored("config.json"));
}
#[test]
fn test_is_weight_file() {
assert!(HuggingFaceProvider::is_weight_file("model.bin"));
assert!(HuggingFaceProvider::is_weight_file("model.safetensors"));
assert!(HuggingFaceProvider::is_weight_file("model.h5"));
assert!(HuggingFaceProvider::is_weight_file("model.msgpack"));
assert!(HuggingFaceProvider::is_weight_file("model.ckpt.index"));
assert!(HuggingFaceProvider::is_weight_file("model.iop"));
assert!(HuggingFaceProvider::is_weight_file("model.gas"));
assert!(!HuggingFaceProvider::is_weight_file("tokenizer.json"));
assert!(!HuggingFaceProvider::is_weight_file("config.json"));
assert!(!HuggingFaceProvider::is_weight_file("README.md"));
}
#[test]
fn test_canonical_model_name_default_preserves_input() {
let provider = HuggingFaceProvider;
let canonical = provider.canonical_model_name("test/model");
assert!(
canonical
.as_ref()
.is_ok_and(|model_name| model_name == "test/model"),
"Expected canonical model name, got {canonical:?}"
);
}
#[cfg(feature = "gcs")]
#[test]
fn test_canonical_model_name_gcs_trims_trailing_slash() {
let provider = GcsProvider;
let canonical = provider.canonical_model_name("gs://test-bucket/org/model/rev-1/");
assert!(
canonical
.as_ref()
.is_ok_and(|model_name| model_name == "gs://test-bucket/org/model/rev-1"),
"Expected canonical model name, got {canonical:?}"
);
}
#[cfg(feature = "gcs")]
#[test]
fn test_gcs_rejects_path_traversal_segments() {
let provider = GcsProvider;
let escapes = [
"gs://bucket/org/../../../etc/passwd",
"gs://bucket/org/./model",
"gs://bucket/org//model",
r"gs://bucket/org\..\etc\passwd",
"gs://bucket/C:/windows/system32",
];
for model_name in escapes {
let result = provider.canonical_model_name(model_name);
assert!(
result.is_err(),
"Expected '{model_name}' to be rejected, got {result:?}"
);
}
}
}