//! Pretrained model loading API similar to HuggingFace's from_pretrained(). //! //! This module provides a convenient API for loading pretrained models from various //! sources including local cache, remote registries, and mirrors. use crate::{HubError, HubResult, ModelId, ModelPackage, ModelRegistry, Registry}; use async_trait::async_trait; use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::path::{Path, PathBuf}; use tokio::fs; use tracing::{debug, info}; /// Configuration for loading pretrained models. #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct PretrainedConfig { /// Git-like revision/tag (e.g., "main", "v1.0.0", commit hash) pub revision: Option, /// Local cache directory (defaults to ~/.cache/rtx-hub) pub cache_dir: Option, /// Force re-download even if cached pub force_download: bool, /// Resume interrupted downloads pub resume_download: bool, /// Only use local files (offline mode) pub local_files_only: bool, /// Authentication token for private models pub use_auth_token: Option, /// Alternative mirror endpoint pub mirror: Option, } impl PretrainedConfig { /// Create a new configuration with defaults. pub fn new() -> Self { Self::default() } /// Set the revision/tag. pub fn with_revision(mut self, revision: impl Into) -> Self { self.revision = Some(revision.into()); self } /// Set the cache directory. pub fn with_cache_dir(mut self, cache_dir: impl Into) -> Self { self.cache_dir = Some(cache_dir.into()); self } /// Enable force download. pub fn with_force_download(mut self, force: bool) -> Self { self.force_download = force; self } /// Enable resume download. pub fn with_resume_download(mut self, resume: bool) -> Self { self.resume_download = resume; self } /// Enable local files only mode. pub fn with_local_files_only(mut self, local_only: bool) -> Self { self.local_files_only = local_only; self } /// Set authentication token. pub fn with_auth_token(mut self, token: impl Into) -> Self { self.use_auth_token = Some(token.into()); self } /// Set mirror endpoint. pub fn with_mirror(mut self, mirror: impl Into) -> Self { self.mirror = Some(mirror.into()); self } /// Get the cache directory, using default if not set. pub fn get_cache_dir(&self) -> PathBuf { self.cache_dir.clone().unwrap_or_else(|| { dirs::home_dir() .unwrap_or_else(|| PathBuf::from(".")) .join(".cache") .join("rtx-hub") }) } } /// Trait for loading models from pretrained checkpoints. #[async_trait] pub trait FromPretrained: Sized { /// Configuration type for this model. type Config; /// Load a model from a pretrained checkpoint. async fn from_pretrained(model_id: &str, config: PretrainedConfig) -> HubResult; } /// Model file information. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ModelFile { /// File path within the model directory pub path: String, /// File size in bytes pub size: u64, /// File hash for integrity verification pub hash: String, /// Last modified timestamp pub modified: DateTime, } /// Loader for pretrained models from registry. pub struct PretrainedLoader { /// Model registry registry: std::sync::Arc, /// Local cache manager cache: LocalCache, } impl PretrainedLoader { /// Create a new pretrained loader. pub fn new(registry: std::sync::Arc, cache_dir: PathBuf) -> Self { let cache = LocalCache::new(cache_dir); Self { registry, cache } } /// Resolve model ID to storage path. pub async fn resolve_model( &self, model_id: &str, config: &PretrainedConfig, ) -> HubResult { let parsed_id = ModelId::parse(model_id)?; // Determine version from revision or use latest let version = if let Some(ref revision) = config.revision { // Try to parse as semver first if let Ok(semver) = semver::Version::parse(revision) { Some(crate::ModelVersion::new(semver)) } else { // Otherwise get latest version None } } else { None }; // Get model info from registry let model_info = self .registry .get_model(&parsed_id, version.as_ref()) .await?; // Check if cached let cache_path = self .cache .cache_path(model_id, &model_info.metadata.version.to_string()); if config.local_files_only && !self.cache.is_cached(&cache_path).await { return Err(HubError::ModelNotFound { model_id: format!("{} (local files only mode)", model_id), }); } Ok(cache_path) } /// Download model if not cached. pub async fn download_if_needed( &self, model_id: &str, config: &PretrainedConfig, ) -> HubResult { let parsed_id = ModelId::parse(model_id)?; let cache_path = self.resolve_model(model_id, config).await?; // Check if we need to download let needs_download = config.force_download || !self.cache.is_cached(&cache_path).await; if needs_download && !config.local_files_only { info!("Downloading model: {}", model_id); // Get model version let model_info = self.registry.get_model(&parsed_id, None).await?; let version = &model_info.metadata.version; // Download package from registry let package = self.registry.download_package(&parsed_id, version).await?; // Extract to cache self.cache.store_package(&cache_path, &package).await?; debug!("Model downloaded and cached at: {}", cache_path.display()); } else { debug!("Using cached model at: {}", cache_path.display()); } Ok(cache_path) } /// Load model configuration from cache. pub async fn load_config(&self, model_path: &Path) -> HubResult { let config_path = model_path.join("config.json"); if !config_path.exists() { return Err(HubError::InvalidPackage { reason: "Model config not found".to_string(), }); } let config_data = fs::read_to_string(&config_path).await?; let config: serde_json::Value = serde_json::from_str(&config_data)?; Ok(config) } /// Load tokenizer from cache. pub async fn load_tokenizer(&self, model_path: &Path) -> HubResult { let tokenizer_path = model_path.join("tokenizer.json"); if !tokenizer_path.exists() { return Err(HubError::InvalidPackage { reason: "Tokenizer not found".to_string(), }); } let tokenizer_data = fs::read_to_string(&tokenizer_path).await?; let tokenizer: serde_json::Value = serde_json::from_str(&tokenizer_data)?; Ok(tokenizer) } /// Get list of all model files. pub async fn get_model_files(&self, model_path: &Path) -> HubResult> { let mut files = Vec::new(); let mut entries = fs::read_dir(model_path).await?; while let Some(entry) = entries.next_entry().await? { let path = entry.path(); let metadata = fs::metadata(&path).await?; if metadata.is_file() { let relative_path = path .strip_prefix(model_path) .map_err(|e| HubError::InvalidPackage { reason: e.to_string(), })? .to_string_lossy() .to_string(); files.push(ModelFile { path: relative_path, size: metadata.len(), hash: String::new(), // TODO: Calculate hash modified: DateTime::::from(metadata.modified()?), }); } } Ok(files) } } /// Local cache manager for pretrained models. pub struct LocalCache { /// Base cache directory cache_dir: PathBuf, /// Cache metadata metadata: HashMap, } /// Cache entry metadata. #[derive(Debug, Clone, Serialize, Deserialize)] struct CacheEntry { /// Model ID model_id: String, /// Model version version: String, /// Cache path path: PathBuf, /// Size in bytes size: u64, /// Last accessed timestamp last_accessed: DateTime, /// Number of accesses access_count: u64, } impl LocalCache { /// Create a new local cache manager. pub fn new(cache_dir: PathBuf) -> Self { Self { cache_dir, metadata: HashMap::new(), } } /// Get cache path for a model. pub fn cache_path(&self, model_id: &str, version: &str) -> PathBuf { self.cache_dir .join("models") .join(model_id.replace('/', "--")) .join(version) } /// Check if model is cached. pub async fn is_cached(&self, cache_path: &Path) -> bool { cache_path.exists() && cache_path.is_dir() } /// Store package to cache. pub async fn store_package(&self, cache_path: &Path, package: &ModelPackage) -> HubResult<()> { // Create cache directory fs::create_dir_all(cache_path).await?; // Write package metadata let metadata_path = cache_path.join("metadata.json"); let metadata_json = serde_json::to_string_pretty(&package.metadata)?; fs::write(&metadata_path, metadata_json).await?; // Write manifest let manifest_path = cache_path.join("manifest.json"); let manifest_json = serde_json::to_string_pretty(&package.manifest)?; fs::write(&manifest_path, manifest_json).await?; debug!("Package stored to cache at: {}", cache_path.display()); Ok(()) } /// Evict least recently used models. pub async fn evict_lru(&mut self, max_size: u64) -> HubResult<()> { let current_size = self.get_cache_size().await?; if current_size <= max_size { return Ok(()); } // Sort by last accessed time let mut entries: Vec<_> = self.metadata.values().cloned().collect(); entries.sort_by_key(|e| e.last_accessed); let mut freed_size = 0u64; for entry in entries { if current_size - freed_size <= max_size { break; } // Remove cache directory if entry.path.exists() { fs::remove_dir_all(&entry.path).await?; freed_size += entry.size; debug!( "Evicted model {} from cache (freed {} bytes)", entry.model_id, entry.size ); } // Remove from metadata self.metadata.remove(&entry.model_id); } Ok(()) } /// Get total cache size in bytes. pub async fn get_cache_size(&self) -> HubResult { let mut total_size = 0u64; if !self.cache_dir.exists() { return Ok(0); } let mut entries = fs::read_dir(&self.cache_dir).await?; while let Some(entry) = entries.next_entry().await? { let path = entry.path(); if path.is_dir() { total_size += self.calculate_dir_size(&path).await?; } } Ok(total_size) } /// Clear all cached models. pub async fn clear_cache(&mut self) -> HubResult<()> { if self.cache_dir.exists() { fs::remove_dir_all(&self.cache_dir).await?; fs::create_dir_all(&self.cache_dir).await?; } self.metadata.clear(); info!("Cache cleared"); Ok(()) } /// Calculate directory size recursively. fn calculate_dir_size<'a>( &'a self, dir: &'a Path, ) -> std::pin::Pin> + 'a>> { Box::pin(async move { let mut total_size = 0u64; let mut entries = fs::read_dir(dir).await?; while let Some(entry) = entries.next_entry().await? { let path = entry.path(); let metadata = fs::metadata(&path).await?; if metadata.is_file() { total_size += metadata.len(); } else if metadata.is_dir() { total_size += self.calculate_dir_size(&path).await?; } } Ok(total_size) }) } } #[cfg(test)] mod tests { use super::*; use crate::model::ModelSchema; use crate::{ModelMetadata, ModelStatus, ModelVersion, RegistryConfig, StorageConfig}; use semver::Version; use std::collections::HashMap; use tempfile::TempDir; #[test] fn test_pretrained_config_builder() { let config = PretrainedConfig::new() .with_revision("v1.0.0") .with_force_download(true) .with_local_files_only(false) .with_auth_token("test-token"); assert_eq!(config.revision, Some("v1.0.0".to_string())); assert!(config.force_download); assert!(!config.local_files_only); assert_eq!(config.use_auth_token, Some("test-token".to_string())); } #[test] fn test_pretrained_config_cache_dir() { let config = PretrainedConfig::new(); let cache_dir = config.get_cache_dir(); assert!(cache_dir.to_string_lossy().contains("rtx-hub")); let custom_dir = PathBuf::from("/tmp/custom-cache"); let config = PretrainedConfig::new().with_cache_dir(custom_dir.clone()); assert_eq!(config.get_cache_dir(), custom_dir); } #[tokio::test] async fn test_local_cache_path() { let temp_dir = TempDir::new().unwrap(); let cache = LocalCache::new(temp_dir.path().to_path_buf()); let cache_path = cache.cache_path("namespace/model", "1.0.0"); assert!(cache_path.to_string_lossy().contains("namespace--model")); assert!(cache_path.to_string_lossy().contains("1.0.0")); } #[tokio::test] async fn test_local_cache_is_cached() { let temp_dir = TempDir::new().unwrap(); let cache = LocalCache::new(temp_dir.path().to_path_buf()); let cache_path = cache.cache_path("test/model", "1.0.0"); assert!(!cache.is_cached(&cache_path).await); // Create cache directory fs::create_dir_all(&cache_path).await.unwrap(); assert!(cache.is_cached(&cache_path).await); } #[tokio::test] async fn test_local_cache_clear() { let temp_dir = TempDir::new().unwrap(); let mut cache = LocalCache::new(temp_dir.path().to_path_buf()); // Create some cache directories let cache_path1 = cache.cache_path("test/model1", "1.0.0"); let cache_path2 = cache.cache_path("test/model2", "1.0.0"); fs::create_dir_all(&cache_path1).await.unwrap(); fs::create_dir_all(&cache_path2).await.unwrap(); // Clear cache cache.clear_cache().await.unwrap(); assert!(!cache.is_cached(&cache_path1).await); assert!(!cache.is_cached(&cache_path2).await); } #[tokio::test] async fn test_local_cache_size() { let temp_dir = TempDir::new().unwrap(); let cache = LocalCache::new(temp_dir.path().to_path_buf()); let initial_size = cache.get_cache_size().await.unwrap(); assert_eq!(initial_size, 0); // Create a test file let cache_path = cache.cache_path("test/model", "1.0.0"); fs::create_dir_all(&cache_path).await.unwrap(); let test_file = cache_path.join("test.txt"); fs::write(&test_file, b"Hello, world!").await.unwrap(); let size = cache.get_cache_size().await.unwrap(); assert!(size > 0); } async fn create_test_registry() -> (ModelRegistry, TempDir) { let temp_dir = TempDir::new().unwrap(); let config = RegistryConfig { storage: StorageConfig::Local { base_path: temp_dir.path().to_path_buf(), }, database_url: "sqlite::memory:".to_string(), enable_validation: false, enable_compression: false, ..Default::default() }; let registry = ModelRegistry::new(config).await.unwrap(); (registry, temp_dir) } fn create_test_metadata(id: ModelId, version: ModelVersion) -> ModelMetadata { ModelMetadata { id, version, title: "Test Model".to_string(), description: "Test Description".to_string(), architecture: "transformer".to_string(), framework: "rustytorch".to_string(), framework_version: "1.0.0".to_string(), tags: vec!["test".to_string()], author: "Test Author".to_string(), license: Some("MIT".to_string()), created_at: Utc::now(), updated_at: Utc::now(), status: ModelStatus::Available, size: 1024, content_hash: "test-hash".to_string(), dependencies: vec![], schema: ModelSchema { inputs: vec![], outputs: vec![], config: None, }, metrics: HashMap::new(), metadata: HashMap::new(), } } #[tokio::test] async fn test_pretrained_loader_resolve_model() { let (registry, _temp_dir) = create_test_registry().await; let cache_dir = TempDir::new().unwrap(); // Register a test model let model_id = ModelId::new("test", "model"); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let metadata = create_test_metadata(model_id.clone(), version.clone()); registry.register_model(metadata).await.unwrap(); let loader = PretrainedLoader::new( std::sync::Arc::new(registry), cache_dir.path().to_path_buf(), ); // Resolve model let config = PretrainedConfig::new(); let resolved_path = loader.resolve_model("test/model", &config).await.unwrap(); assert!(resolved_path.to_string_lossy().contains("test--model")); assert!(resolved_path.to_string_lossy().contains("1.0.0")); } #[tokio::test] async fn test_pretrained_loader_local_files_only() { let (registry, _temp_dir) = create_test_registry().await; let cache_dir = TempDir::new().unwrap(); // Register a test model let model_id = ModelId::new("test", "model"); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let metadata = create_test_metadata(model_id.clone(), version.clone()); registry.register_model(metadata).await.unwrap(); let loader = PretrainedLoader::new( std::sync::Arc::new(registry), cache_dir.path().to_path_buf(), ); // Try to resolve with local_files_only let config = PretrainedConfig::new().with_local_files_only(true); let result = loader.resolve_model("test/model", &config).await; // Should fail because model is not cached assert!(result.is_err()); } #[tokio::test] async fn test_pretrained_loader_get_model_files() { let temp_dir = TempDir::new().unwrap(); let cache_dir = TempDir::new().unwrap(); let (registry, _) = create_test_registry().await; let loader = PretrainedLoader::new( std::sync::Arc::new(registry), cache_dir.path().to_path_buf(), ); // Create test files let model_path = temp_dir.path().join("model"); fs::create_dir_all(&model_path).await.unwrap(); fs::write(model_path.join("config.json"), b"{}") .await .unwrap(); fs::write(model_path.join("weights.bin"), b"fake weights") .await .unwrap(); // Get model files let files = loader.get_model_files(&model_path).await.unwrap(); assert_eq!(files.len(), 2); assert!(files.iter().any(|f| f.path == "config.json")); assert!(files.iter().any(|f| f.path == "weights.bin")); } #[tokio::test] async fn test_pretrained_loader_load_config() { let temp_dir = TempDir::new().unwrap(); let cache_dir = TempDir::new().unwrap(); let (registry, _) = create_test_registry().await; let loader = PretrainedLoader::new( std::sync::Arc::new(registry), cache_dir.path().to_path_buf(), ); // Create test config let model_path = temp_dir.path().join("model"); fs::create_dir_all(&model_path).await.unwrap(); let config_json = r#"{"model_type": "transformer", "hidden_size": 768}"#; fs::write(model_path.join("config.json"), config_json) .await .unwrap(); // Load config let config = loader.load_config(&model_path).await.unwrap(); assert_eq!(config["model_type"], "transformer"); assert_eq!(config["hidden_size"], 768); } #[tokio::test] async fn test_pretrained_loader_load_tokenizer() { let temp_dir = TempDir::new().unwrap(); let cache_dir = TempDir::new().unwrap(); let (registry, _) = create_test_registry().await; let loader = PretrainedLoader::new( std::sync::Arc::new(registry), cache_dir.path().to_path_buf(), ); // Create test tokenizer let model_path = temp_dir.path().join("model"); fs::create_dir_all(&model_path).await.unwrap(); let tokenizer_json = r#"{"vocab_size": 50257}"#; fs::write(model_path.join("tokenizer.json"), tokenizer_json) .await .unwrap(); // Load tokenizer let tokenizer = loader.load_tokenizer(&model_path).await.unwrap(); assert_eq!(tokenizer["vocab_size"], 50257); } #[tokio::test] async fn test_local_cache_evict_lru() { let temp_dir = TempDir::new().unwrap(); let mut cache = LocalCache::new(temp_dir.path().to_path_buf()); // Create actual directories first let model1_path = temp_dir.path().join("models").join("model1"); let model2_path = temp_dir.path().join("models").join("model2"); fs::create_dir_all(&model1_path).await.unwrap(); fs::create_dir_all(&model2_path).await.unwrap(); // Write some data to make them have size fs::write(model1_path.join("data.bin"), vec![0u8; 1000]) .await .unwrap(); fs::write(model2_path.join("data.bin"), vec![0u8; 2000]) .await .unwrap(); // Add cache entries let now = Utc::now(); cache.metadata.insert( "model1".to_string(), CacheEntry { model_id: "model1".to_string(), version: "1.0.0".to_string(), path: model1_path.clone(), size: 1000, last_accessed: now - chrono::Duration::days(10), access_count: 1, }, ); cache.metadata.insert( "model2".to_string(), CacheEntry { model_id: "model2".to_string(), version: "1.0.0".to_string(), path: model2_path.clone(), size: 2000, last_accessed: now - chrono::Duration::days(5), access_count: 2, }, ); // Initial state: both models exist assert!(model1_path.exists()); assert!(model2_path.exists()); // Evict LRU with max size 2500 (no eviction needed) cache.evict_lru(5000).await.unwrap(); // Both should still exist assert!(model1_path.exists()); assert!(model2_path.exists()); // Test that metadata is tracked correctly assert_eq!(cache.metadata.len(), 2); } }