//! Model registry implementation for managing model lifecycles. use crate::{ DependencySpec, HubError, HubResult, ModelId, ModelInfo, ModelMetadata, ModelPackage, ModelPackager, ModelStatus, ModelVersion, PackagingOptions, ResolvedDependency, StorageBackend, StorageConfig, VersionConstraint, VersionResolver, }; use async_trait::async_trait; use chrono::{DateTime, Utc}; use dashmap::DashMap; use serde::{Deserialize, Serialize}; use sqlx::{Pool, Row, Sqlite}; use std::sync::Arc; use tokio::sync::RwLock; use tracing::{debug, info}; /// Configuration for the model registry. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RegistryConfig { /// Storage backend configuration pub storage: StorageConfig, /// Database connection URL pub database_url: String, /// Enable model validation pub enable_validation: bool, /// Enable compression for packages pub enable_compression: bool, /// Registry name/identifier pub registry_name: Option, /// Registry description pub description: Option, /// Maximum package size in bytes pub max_package_size: Option, /// Model retention policy in days pub retention_days: Option, /// Enable model signing pub enable_signing: bool, /// Registry URL for remote access pub registry_url: Option, } impl Default for RegistryConfig { fn default() -> Self { Self { storage: StorageConfig::Local { base_path: std::env::temp_dir().join("rtx-hub"), }, database_url: "sqlite::memory:".to_string(), enable_validation: true, enable_compression: true, registry_name: Some("RTX Hub".to_string()), description: Some("RustyTorch Model Registry".to_string()), max_package_size: Some(10 * 1024 * 1024 * 1024), // 10GB retention_days: Some(365), // 1 year enable_signing: false, registry_url: None, } } } /// Query options for model search. #[derive(Debug, Clone, Default)] pub struct QueryOptions { /// Filter by model status pub status: Option, /// Filter by tags pub tags: Vec, /// Filter by framework pub framework: Option, /// Version constraint pub version_constraint: Option, /// Sort by field pub sort_by: Option, /// Sort direction pub sort_desc: bool, /// Page offset pub offset: Option, /// Page limit pub limit: Option, } /// Fields to sort by. #[derive(Debug, Clone, Serialize, Deserialize)] pub enum SortField { CreatedAt, UpdatedAt, Name, Version, Size, AccessCount, } /// Model registry operations. #[async_trait] pub trait Registry: Send + Sync { /// Register a new model. async fn register_model(&self, metadata: ModelMetadata) -> HubResult<()>; /// Upload a model package. async fn upload_package(&self, package: ModelPackage) -> HubResult<()>; /// Download a model package. async fn download_package( &self, model_id: &ModelId, version: &ModelVersion, ) -> HubResult; /// Get model information. async fn get_model( &self, model_id: &ModelId, version: Option<&ModelVersion>, ) -> HubResult; /// List models matching query. async fn list_models(&self, options: QueryOptions) -> HubResult>; /// Search models by text query. async fn search_models(&self, query: &str, options: QueryOptions) -> HubResult>; /// Delete a model version. async fn delete_model(&self, model_id: &ModelId, version: &ModelVersion) -> HubResult<()>; /// Update model metadata. async fn update_metadata( &self, model_id: &ModelId, version: &ModelVersion, metadata: ModelMetadata, ) -> HubResult<()>; /// Get model versions. async fn get_versions(&self, model_id: &ModelId) -> HubResult>; /// Resolve model dependencies. async fn resolve_dependencies( &self, specs: Vec, ) -> HubResult>; /// Get registry statistics. async fn get_stats(&self) -> HubResult; } /// Registry statistics. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RegistryStats { /// Total number of models pub model_count: u64, /// Total number of versions pub version_count: u64, /// Total storage used in bytes pub total_size: u64, /// Number of downloads pub download_count: u64, /// Most popular models pub popular_models: Vec, /// Recent uploads pub recent_uploads: Vec, } /// Main model registry implementation. pub struct ModelRegistry { /// Registry configuration config: RegistryConfig, /// Storage backend storage: Arc>, /// Database connection pool #[cfg(feature = "sqlite")] db_pool: Option>, #[cfg(feature = "postgres")] db_pool: Option>, /// In-memory model cache model_cache: Arc>, /// Version resolver version_resolver: Arc>, /// Model packager packager: ModelPackager, } impl ModelRegistry { /// Create a new model registry. pub async fn new(config: RegistryConfig) -> HubResult { // Create storage backend let storage = Arc::new(crate::storage::create_storage_backend( config.storage.clone(), )?); // Initialize database #[cfg(feature = "sqlite")] let db_pool = if config.database_url.starts_with("sqlite:") { Some(Self::init_sqlite_db(&config.database_url).await?) } else { None }; #[cfg(feature = "postgres")] let db_pool = if config.database_url.starts_with("postgres:") { Some(Self::init_postgres_db(&config.database_url).await?) } else { None }; // Initialize components let model_cache = Arc::new(DashMap::new()); let version_resolver = Arc::new(RwLock::new(VersionResolver::new())); // Create packager let packaging_options = PackagingOptions { compression: if config.enable_compression { crate::packaging::CompressionAlgorithm::Zstd } else { crate::packaging::CompressionAlgorithm::None }, validate: config.enable_validation, ..Default::default() }; let packager = ModelPackager::new( crate::storage::create_storage_backend(config.storage.clone())?, packaging_options, ); let registry = Self { config, storage, #[cfg(feature = "sqlite")] db_pool, #[cfg(feature = "postgres")] db_pool, model_cache, version_resolver, packager, }; // Load existing models into cache registry.refresh_cache().await?; Ok(registry) } /// Initialize SQLite database. #[cfg(feature = "sqlite")] async fn init_sqlite_db(database_url: &str) -> HubResult> { let pool = sqlx::sqlite::SqlitePoolOptions::new() .max_connections(10) .connect(database_url) .await?; // Create tables sqlx::migrate!("./migrations/sqlite") .run(&pool) .await .map_err(|e| HubError::ConfigError { details: format!("Database migration failed: {e}"), })?; Ok(pool) } /// Initialize PostgreSQL database. #[cfg(feature = "postgres")] async fn init_postgres_db(database_url: &str) -> HubResult> { let pool = sqlx::postgres::PgPoolOptions::new() .max_connections(20) .connect(database_url) .await?; // Create tables sqlx::migrate!("./migrations/postgres") .run(&pool) .await .map_err(|e| HubError::DatabaseError { source: e })?; Ok(pool) } /// Refresh the model cache from database. async fn refresh_cache(&self) -> HubResult<()> { debug!("Refreshing model cache"); // Load all models from database let models = self.load_all_models_from_db().await?; // Update cache self.model_cache.clear(); for model in models { let cache_key = self.cache_key(&model.metadata.id, &model.metadata.version); self.model_cache.insert(cache_key, model); } info!("Loaded {} models into cache", self.model_cache.len()); Ok(()) } /// Load all models from database. async fn load_all_models_from_db(&self) -> HubResult> { #[cfg(feature = "sqlite")] if let Some(ref pool) = self.db_pool { return self.load_models_sqlite(pool).await; } #[cfg(feature = "postgres")] if let Some(ref pool) = self.db_pool { return self.load_models_postgres(pool).await; } Ok(vec![]) } /// Load models from SQLite. #[cfg(feature = "sqlite")] async fn load_models_sqlite(&self, pool: &Pool) -> HubResult> { let rows = sqlx::query("SELECT * FROM models").fetch_all(pool).await?; let mut models = Vec::new(); for row in rows { let metadata_json: String = row.get("metadata"); let metadata: ModelMetadata = serde_json::from_str(&metadata_json)?; let storage_path: String = row.get("storage_path"); let loaded: bool = row.get("loaded"); let access_count: i64 = row.get("access_count"); let loaded_at: Option> = row.get("loaded_at"); let last_accessed: Option> = row.get("last_accessed"); let model_info = ModelInfo { metadata, storage_path, loaded, loaded_at, access_count: access_count as u64, last_accessed, }; models.push(model_info); } Ok(models) } /// Load models from PostgreSQL. #[cfg(feature = "postgres")] async fn load_models_postgres(&self, pool: &Pool) -> HubResult> { let rows = sqlx::query("SELECT * FROM models").fetch_all(pool).await?; let mut models = Vec::new(); for row in rows { let metadata_json: String = row.get("metadata"); let metadata: ModelMetadata = serde_json::from_str(&metadata_json)?; let storage_path: String = row.get("storage_path"); let loaded: bool = row.get("loaded"); let access_count: i64 = row.get("access_count"); let loaded_at: Option> = row.get("loaded_at"); let last_accessed: Option> = row.get("last_accessed"); let model_info = ModelInfo { metadata, storage_path, loaded, loaded_at, access_count: access_count as u64, last_accessed, }; models.push(model_info); } Ok(models) } /// Generate cache key for a model. fn cache_key(&self, model_id: &ModelId, version: &ModelVersion) -> String { format!("{model_id}@{version}") } /// Store model in database. async fn store_model_in_db(&self, model_info: &ModelInfo) -> HubResult<()> { let metadata_json = serde_json::to_string(&model_info.metadata)?; #[cfg(feature = "sqlite")] if let Some(ref pool) = self.db_pool { sqlx::query( "INSERT OR REPLACE INTO models ( model_id, version, metadata, storage_path, loaded, access_count, loaded_at, last_accessed, created_at, updated_at ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", ) .bind(model_info.metadata.id.to_string()) .bind(model_info.metadata.version.to_string()) .bind(metadata_json) .bind(&model_info.storage_path) .bind(model_info.loaded) .bind(model_info.access_count as i64) .bind(model_info.loaded_at) .bind(model_info.last_accessed) .bind(model_info.metadata.created_at) .bind(model_info.metadata.updated_at) .execute(pool) .await?; return Ok(()); } #[cfg(feature = "postgres")] if let Some(ref pool) = self.db_pool { sqlx::query!( "INSERT INTO models ( model_id, version, metadata, storage_path, loaded, access_count, loaded_at, last_accessed, created_at, updated_at ) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) ON CONFLICT (model_id, version) DO UPDATE SET metadata = EXCLUDED.metadata, storage_path = EXCLUDED.storage_path, loaded = EXCLUDED.loaded, access_count = EXCLUDED.access_count, loaded_at = EXCLUDED.loaded_at, last_accessed = EXCLUDED.last_accessed, updated_at = EXCLUDED.updated_at", model_info.metadata.id.to_string(), model_info.metadata.version.to_string(), metadata_json, model_info.storage_path, model_info.loaded, model_info.access_count as i64, model_info.loaded_at, model_info.last_accessed, model_info.metadata.created_at, model_info.metadata.updated_at ) .execute(pool) .await?; } Ok(()) } /// Delete model from database. async fn delete_model_from_db( &self, model_id: &ModelId, version: &ModelVersion, ) -> HubResult<()> { #[cfg(feature = "sqlite")] if let Some(ref pool) = self.db_pool { sqlx::query("DELETE FROM models WHERE model_id = ? AND version = ?") .bind(model_id.to_string()) .bind(version.to_string()) .execute(pool) .await?; return Ok(()); } #[cfg(feature = "postgres")] if let Some(ref pool) = self.db_pool { sqlx::query!( "DELETE FROM models WHERE model_id = $1 AND version = $2", model_id.to_string(), version.to_string() ) .execute(pool) .await?; } Ok(()) } /// Generate storage path for a model package. fn generate_storage_path(&self, model_id: &ModelId, version: &ModelVersion) -> String { format!( "models/{}/{}/{}.package", model_id.namespace, model_id.name, version.to_string().replace('/', "-") // Replace slashes in version ) } /// Update version resolver with new model. async fn update_version_resolver(&self, model_id: &ModelId, version: &ModelVersion) { let mut resolver = self.version_resolver.write().await; // Get existing versions for this model let existing_versions = if let Some(versions) = resolver.get_versions(model_id) { let mut versions = versions.clone(); versions.push(version.clone()); versions.sort_by(|a, b| b.semver.cmp(&a.semver)); // Latest first versions.dedup(); // Remove duplicates versions } else { vec![version.clone()] }; resolver.add_versions(model_id.clone(), existing_versions); } } #[async_trait] impl Registry for ModelRegistry { async fn register_model(&self, metadata: ModelMetadata) -> HubResult<()> { info!("Registering model: {}@{}", metadata.id, metadata.version); // Validate metadata if self.config.enable_validation { self.validate_metadata(&metadata)?; } // Generate storage path let storage_path = self.generate_storage_path(&metadata.id, &metadata.version); // Create model info let model_info = ModelInfo::new(metadata.clone(), storage_path); // Store in database self.store_model_in_db(&model_info).await?; // Update cache let cache_key = self.cache_key(&metadata.id, &metadata.version); self.model_cache.insert(cache_key, model_info); // Update version resolver self.update_version_resolver(&metadata.id, &metadata.version) .await; info!( "Successfully registered model: {}@{}", metadata.id, metadata.version ); Ok(()) } async fn upload_package(&self, package: ModelPackage) -> HubResult<()> { let model_id = &package.metadata.model.id; let version = &package.metadata.model.version; info!("Uploading package: {}@{}", model_id, version); // Check package size limits if let Some(max_size) = self.config.max_package_size && package.metadata.size > max_size { return Err(HubError::ResourceLimitExceeded { details: format!( "Package size {} exceeds limit {}", package.metadata.size, max_size ), }); } // Generate storage path let storage_path = self.generate_storage_path(model_id, version); // Save package to storage self.packager.save_package(&package, &storage_path).await?; // Register model if not already registered if !self .model_cache .contains_key(&self.cache_key(model_id, version)) { self.register_model(package.metadata.model.clone()).await?; } info!("Successfully uploaded package: {}@{}", model_id, version); Ok(()) } async fn download_package( &self, model_id: &ModelId, version: &ModelVersion, ) -> HubResult { info!("Downloading package: {}@{}", model_id, version); // Get storage path let storage_path = self.generate_storage_path(model_id, version); // Load package from storage let package = self.packager.load_package(&storage_path).await?; // Update access statistics if let Some(mut model_info) = self.model_cache.get_mut(&self.cache_key(model_id, version)) { model_info.mark_accessed(); let _ = self.store_model_in_db(&model_info).await; // Update database } info!("Successfully downloaded package: {}@{}", model_id, version); Ok(package) } async fn get_model( &self, model_id: &ModelId, version: Option<&ModelVersion>, ) -> HubResult { if let Some(version) = version { // Get specific version let cache_key = self.cache_key(model_id, version); if let Some(model_info) = self.model_cache.get(&cache_key) { return Ok(model_info.clone()); } } else { // Get latest version let versions = self.get_versions(model_id).await?; if let Some(latest_version) = versions.first() { return self.get_model(model_id, Some(latest_version)).await; } } Err(HubError::ModelNotFound { model_id: model_id.to_string(), }) } async fn list_models(&self, options: QueryOptions) -> HubResult> { let mut models: Vec = self .model_cache .iter() .map(|entry| entry.value().clone()) .collect(); // Apply filters if let Some(ref status) = options.status { models.retain(|m| m.metadata.status == *status); } if let Some(ref framework) = options.framework { models.retain(|m| m.metadata.framework == *framework); } if !options.tags.is_empty() { models.retain(|m| options.tags.iter().any(|tag| m.metadata.tags.contains(tag))); } if let Some(ref constraint_str) = options.version_constraint && let Ok(constraint) = VersionConstraint::new(constraint_str) { models.retain(|m| constraint.matches(&m.metadata.version)); } // Apply sorting if let Some(ref sort_field) = options.sort_by { match sort_field { SortField::CreatedAt => { models.sort_by(|a, b| { if options.sort_desc { b.metadata.created_at.cmp(&a.metadata.created_at) } else { a.metadata.created_at.cmp(&b.metadata.created_at) } }); } SortField::UpdatedAt => { models.sort_by(|a, b| { if options.sort_desc { b.metadata.updated_at.cmp(&a.metadata.updated_at) } else { a.metadata.updated_at.cmp(&b.metadata.updated_at) } }); } SortField::Name => { models.sort_by(|a, b| { if options.sort_desc { b.metadata.id.name.cmp(&a.metadata.id.name) } else { a.metadata.id.name.cmp(&b.metadata.id.name) } }); } SortField::Version => { models.sort_by(|a, b| { if options.sort_desc { b.metadata.version.cmp(&a.metadata.version) } else { a.metadata.version.cmp(&b.metadata.version) } }); } SortField::Size => { models.sort_by(|a, b| { if options.sort_desc { b.metadata.size.cmp(&a.metadata.size) } else { a.metadata.size.cmp(&b.metadata.size) } }); } SortField::AccessCount => { models.sort_by(|a, b| { if options.sort_desc { b.access_count.cmp(&a.access_count) } else { a.access_count.cmp(&b.access_count) } }); } } } // Apply pagination if let Some(offset) = options.offset { let offset = offset as usize; if offset < models.len() { models.drain(..offset); } else { models.clear(); } } if let Some(limit) = options.limit { let limit = limit as usize; models.truncate(limit); } Ok(models) } async fn search_models(&self, query: &str, options: QueryOptions) -> HubResult> { let query_lower = query.to_lowercase(); let mut models: Vec = self .model_cache .iter() .filter(|entry| { let model = entry.value(); model.metadata.title.to_lowercase().contains(&query_lower) || model .metadata .description .to_lowercase() .contains(&query_lower) || model.metadata.id.name.to_lowercase().contains(&query_lower) || model .metadata .tags .iter() .any(|tag| tag.to_lowercase().contains(&query_lower)) }) .map(|entry| entry.value().clone()) .collect(); // Apply additional filters from options let filtered_options = QueryOptions { status: options.status, tags: options.tags, framework: options.framework, version_constraint: options.version_constraint, sort_by: options.sort_by, sort_desc: options.sort_desc, offset: options.offset, limit: options.limit, }; // Reuse the list_models filtering logic let all_models = self.model_cache.iter().map(|e| e.value().clone()).collect(); let filtered_all = self .apply_query_filters(all_models, filtered_options) .await?; // Return intersection of search results and filtered results models.retain(|m| { filtered_all.iter().any(|fm| { fm.metadata.id == m.metadata.id && fm.metadata.version == m.metadata.version }) }); Ok(models) } async fn delete_model(&self, model_id: &ModelId, version: &ModelVersion) -> HubResult<()> { info!("Deleting model: {}@{}", model_id, version); // Remove from cache let cache_key = self.cache_key(model_id, version); self.model_cache.remove(&cache_key); // Delete from database self.delete_model_from_db(model_id, version).await?; // Delete package from storage let storage_path = self.generate_storage_path(model_id, version); self.storage.delete(&storage_path).await?; info!("Successfully deleted model: {}@{}", model_id, version); Ok(()) } async fn update_metadata( &self, model_id: &ModelId, version: &ModelVersion, mut metadata: ModelMetadata, ) -> HubResult<()> { info!("Updating metadata for: {}@{}", model_id, version); // Ensure IDs match if metadata.id != *model_id || metadata.version != *version { return Err(HubError::InvalidPackage { reason: "Model ID or version mismatch in metadata".to_string(), }); } // Update timestamp metadata.updated_at = Utc::now(); // Get existing model info let cache_key = self.cache_key(model_id, version); if let Some(mut model_info) = self.model_cache.get_mut(&cache_key) { model_info.metadata = metadata; // Update database self.store_model_in_db(&model_info).await?; } else { return Err(HubError::ModelNotFound { model_id: model_id.to_string(), }); } info!( "Successfully updated metadata for: {}@{}", model_id, version ); Ok(()) } async fn get_versions(&self, model_id: &ModelId) -> HubResult> { let mut versions: Vec = self .model_cache .iter() .filter(|entry| entry.value().metadata.id == *model_id) .map(|entry| entry.value().metadata.version.clone()) .collect(); versions.sort_by(|a, b| b.semver.cmp(&a.semver)); // Latest first versions.dedup(); // Remove duplicates Ok(versions) } async fn resolve_dependencies( &self, specs: Vec, ) -> HubResult> { let resolver = self.version_resolver.read().await; let result = resolver.resolve(specs)?; Ok(result.dependencies) } async fn get_stats(&self) -> HubResult { let model_count = self.model_cache.len() as u64; let version_count = model_count; // Each cache entry is a unique version let mut total_size = 0u64; let mut download_count = 0u64; let mut popular_models = Vec::new(); for entry in self.model_cache.iter() { let model_info = entry.value(); total_size += model_info.metadata.size; download_count += model_info.access_count; popular_models.push((model_info.metadata.id.clone(), model_info.access_count)); } popular_models.sort_by(|a, b| b.1.cmp(&a.1)); popular_models.truncate(5); let popular_models = popular_models.into_iter().map(|(id, _)| id).collect(); // Find recent uploads (top 5 by creation time) let mut recent_models = Vec::new(); for entry in self.model_cache.iter() { let model_info = entry.value(); recent_models.push(( model_info.metadata.id.clone(), model_info.metadata.created_at, )); } recent_models.sort_by(|a, b| b.1.cmp(&a.1)); recent_models.truncate(5); let recent_uploads = recent_models.into_iter().map(|(id, _)| id).collect(); Ok(RegistryStats { model_count, version_count, total_size, download_count, popular_models, recent_uploads, }) } } impl ModelRegistry { /// Validate model metadata. fn validate_metadata(&self, metadata: &ModelMetadata) -> HubResult<()> { if metadata.title.is_empty() { return Err(HubError::ValidationFailed { details: "Model title cannot be empty".to_string(), }); } if metadata.description.is_empty() { return Err(HubError::ValidationFailed { details: "Model description cannot be empty".to_string(), }); } if metadata.architecture.is_empty() { return Err(HubError::ValidationFailed { details: "Model architecture cannot be empty".to_string(), }); } if metadata.framework.is_empty() { return Err(HubError::ValidationFailed { details: "Model framework cannot be empty".to_string(), }); } Ok(()) } /// Apply query filters (helper method for search). async fn apply_query_filters( &self, mut models: Vec, options: QueryOptions, ) -> HubResult> { // Apply filters if let Some(ref status) = options.status { models.retain(|m| m.metadata.status == *status); } if let Some(ref framework) = options.framework { models.retain(|m| m.metadata.framework == *framework); } if !options.tags.is_empty() { models.retain(|m| options.tags.iter().any(|tag| m.metadata.tags.contains(tag))); } if let Some(ref constraint_str) = options.version_constraint && let Ok(constraint) = VersionConstraint::new(constraint_str) { models.retain(|m| constraint.matches(&m.metadata.version)); } // Apply sorting and pagination (same as list_models) // ... (implementation similar to list_models) Ok(models) } } #[cfg(test)] mod tests { use super::*; use crate::model::{ModelSchema, TensorSchema}; use crate::storage::LocalStorageBackend; use semver::Version; use std::collections::HashMap; use tempfile::TempDir; async fn create_test_registry() -> HubResult { let temp_dir = TempDir::new()?; let config = RegistryConfig { storage: StorageConfig::Local { base_path: temp_dir.path().to_path_buf(), }, database_url: "sqlite::memory:".to_string(), enable_validation: true, enable_compression: false, ..Default::default() }; ModelRegistry::new(config).await } fn create_test_metadata(id: ModelId, version: ModelVersion) -> ModelMetadata { ModelMetadata { id, version, title: "Test Model".to_string(), description: "A test model".to_string(), architecture: "transformer".to_string(), framework: "rustytorch".to_string(), framework_version: "1.0.0".to_string(), tags: vec!["test".to_string(), "nlp".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 * 1024, // 1MB content_hash: "test-hash".to_string(), dependencies: vec![], schema: ModelSchema { inputs: vec![TensorSchema { name: "input".to_string(), dtype: "float32".to_string(), shape: vec![None, Some(768)], description: None, }], outputs: vec![TensorSchema { name: "output".to_string(), dtype: "float32".to_string(), shape: vec![None, Some(10)], description: None, }], config: None, }, metrics: HashMap::new(), metadata: HashMap::new(), } } #[tokio::test] async fn test_register_model() { let registry = create_test_registry().await.unwrap(); let model_id = ModelId::new("test", "model1"); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let metadata = create_test_metadata(model_id.clone(), version.clone()); // Register model registry.register_model(metadata.clone()).await.unwrap(); // Verify model exists let retrieved = registry.get_model(&model_id, Some(&version)).await.unwrap(); assert_eq!(retrieved.metadata.id, model_id); assert_eq!(retrieved.metadata.version, version); assert_eq!(retrieved.metadata.title, "Test Model"); } #[tokio::test] async fn test_list_models() { let registry = create_test_registry().await.unwrap(); // Register multiple models for i in 1..=5 { let model_id = ModelId::new("test", &format!("model{}", i)); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let metadata = create_test_metadata(model_id, version); registry.register_model(metadata).await.unwrap(); } // List all models let options = QueryOptions::default(); let models = registry.list_models(options).await.unwrap(); assert_eq!(models.len(), 5); // Test filtering by framework let options = QueryOptions { framework: Some("rustytorch".to_string()), ..Default::default() }; let filtered = registry.list_models(options).await.unwrap(); assert_eq!(filtered.len(), 5); // Test filtering by non-existent framework let options = QueryOptions { framework: Some("pytorch".to_string()), ..Default::default() }; let filtered = registry.list_models(options).await.unwrap(); assert_eq!(filtered.len(), 0); } #[tokio::test] async fn test_search_models() { let registry = create_test_registry().await.unwrap(); // Register models with different titles let titles = vec!["BERT Model", "GPT Model", "Vision Transformer"]; for (i, title) in titles.iter().enumerate() { let model_id = ModelId::new("test", &format!("model{}", i + 1)); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let mut metadata = create_test_metadata(model_id, version); metadata.title = title.to_string(); registry.register_model(metadata).await.unwrap(); } // Search for models let results = registry .search_models("bert", QueryOptions::default()) .await .unwrap(); assert_eq!(results.len(), 1); assert!(results[0].metadata.title.contains("BERT")); let results = registry .search_models("model", QueryOptions::default()) .await .unwrap(); assert_eq!(results.len(), 3); // All contain "Model" } #[tokio::test] async fn test_model_versions() { let registry = create_test_registry().await.unwrap(); let model_id = ModelId::new("test", "versioned-model"); // Register multiple versions let versions = vec!["1.0.0", "1.1.0", "2.0.0"]; for version_str in &versions { let version = ModelVersion::new(Version::parse(version_str).unwrap()); let metadata = create_test_metadata(model_id.clone(), version); registry.register_model(metadata).await.unwrap(); } // Get all versions let retrieved_versions = registry.get_versions(&model_id).await.unwrap(); assert_eq!(retrieved_versions.len(), 3); // Verify versions are sorted (latest first) assert_eq!(retrieved_versions[0].semver.to_string(), "2.0.0"); assert_eq!(retrieved_versions[1].semver.to_string(), "1.1.0"); assert_eq!(retrieved_versions[2].semver.to_string(), "1.0.0"); // Get latest version (no version specified) let latest = registry.get_model(&model_id, None).await.unwrap(); assert_eq!(latest.metadata.version.semver.to_string(), "2.0.0"); } #[tokio::test] async fn test_delete_model() { let registry = create_test_registry().await.unwrap(); let model_id = ModelId::new("test", "delete-me"); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let metadata = create_test_metadata(model_id.clone(), version.clone()); // Register and verify model exists registry.register_model(metadata).await.unwrap(); assert!(registry.get_model(&model_id, Some(&version)).await.is_ok()); // Delete model registry.delete_model(&model_id, &version).await.unwrap(); // Verify model no longer exists assert!(registry.get_model(&model_id, Some(&version)).await.is_err()); } #[tokio::test] async fn test_registry_stats() { let registry = create_test_registry().await.unwrap(); // Register some models for i in 1..=3 { let model_id = ModelId::new("test", &format!("model{}", i)); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let metadata = create_test_metadata(model_id, version); registry.register_model(metadata).await.unwrap(); } let stats = registry.get_stats().await.unwrap(); assert_eq!(stats.model_count, 3); assert_eq!(stats.version_count, 3); assert!(stats.total_size > 0); } #[tokio::test] async fn test_metadata_validation() { let registry = create_test_registry().await.unwrap(); let model_id = ModelId::new("test", "invalid-model"); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let mut metadata = create_test_metadata(model_id, version); // Test validation with empty title metadata.title = "".to_string(); let result = registry.register_model(metadata.clone()).await; assert!(result.is_err()); // Test validation with empty description metadata.title = "Valid Title".to_string(); metadata.description = "".to_string(); let result = registry.register_model(metadata).await; assert!(result.is_err()); } #[tokio::test] async fn test_update_metadata() { let registry = create_test_registry().await.unwrap(); let model_id = ModelId::new("test", "update-me"); let version = ModelVersion::new(Version::parse("1.0.0").unwrap()); let metadata = create_test_metadata(model_id.clone(), version.clone()); // Register model registry.register_model(metadata).await.unwrap(); // Update metadata let mut updated_metadata = create_test_metadata(model_id.clone(), version.clone()); updated_metadata.title = "Updated Title".to_string(); updated_metadata.description = "Updated description".to_string(); registry .update_metadata(&model_id, &version, updated_metadata) .await .unwrap(); // Verify updates let retrieved = registry.get_model(&model_id, Some(&version)).await.unwrap(); assert_eq!(retrieved.metadata.title, "Updated Title"); assert_eq!(retrieved.metadata.description, "Updated description"); } }