Files
rustytorch/crates/production/rtx-hub/src/registry.rs
T
2026-03-04 00:08:42 +00:00

1211 lines
40 KiB
Rust

//! 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<String>,
/// Registry description
pub description: Option<String>,
/// Maximum package size in bytes
pub max_package_size: Option<u64>,
/// Model retention policy in days
pub retention_days: Option<u32>,
/// Enable model signing
pub enable_signing: bool,
/// Registry URL for remote access
pub registry_url: Option<String>,
}
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<ModelStatus>,
/// Filter by tags
pub tags: Vec<String>,
/// Filter by framework
pub framework: Option<String>,
/// Version constraint
pub version_constraint: Option<String>,
/// Sort by field
pub sort_by: Option<SortField>,
/// Sort direction
pub sort_desc: bool,
/// Page offset
pub offset: Option<u32>,
/// Page limit
pub limit: Option<u32>,
}
/// 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<ModelPackage>;
/// Get model information.
async fn get_model(
&self,
model_id: &ModelId,
version: Option<&ModelVersion>,
) -> HubResult<ModelInfo>;
/// List models matching query.
async fn list_models(&self, options: QueryOptions) -> HubResult<Vec<ModelInfo>>;
/// Search models by text query.
async fn search_models(&self, query: &str, options: QueryOptions) -> HubResult<Vec<ModelInfo>>;
/// 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<Vec<ModelVersion>>;
/// Resolve model dependencies.
async fn resolve_dependencies(
&self,
specs: Vec<DependencySpec>,
) -> HubResult<Vec<ResolvedDependency>>;
/// Get registry statistics.
async fn get_stats(&self) -> HubResult<RegistryStats>;
}
/// 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<ModelId>,
/// Recent uploads
pub recent_uploads: Vec<ModelId>,
}
/// Main model registry implementation.
pub struct ModelRegistry {
/// Registry configuration
config: RegistryConfig,
/// Storage backend
storage: Arc<Box<dyn StorageBackend>>,
/// Database connection pool
#[cfg(feature = "sqlite")]
db_pool: Option<Pool<Sqlite>>,
#[cfg(feature = "postgres")]
db_pool: Option<Pool<Postgres>>,
/// In-memory model cache
model_cache: Arc<DashMap<String, ModelInfo>>,
/// Version resolver
version_resolver: Arc<RwLock<VersionResolver>>,
/// Model packager
packager: ModelPackager,
}
impl ModelRegistry {
/// Create a new model registry.
pub async fn new(config: RegistryConfig) -> HubResult<Self> {
// 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<Pool<Sqlite>> {
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<Pool<Postgres>> {
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<Vec<ModelInfo>> {
#[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<Sqlite>) -> HubResult<Vec<ModelInfo>> {
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<DateTime<Utc>> = row.get("loaded_at");
let last_accessed: Option<DateTime<Utc>> = 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<Postgres>) -> HubResult<Vec<ModelInfo>> {
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<DateTime<Utc>> = row.get("loaded_at");
let last_accessed: Option<DateTime<Utc>> = 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<ModelPackage> {
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<ModelInfo> {
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<Vec<ModelInfo>> {
let mut models: Vec<ModelInfo> = 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<Vec<ModelInfo>> {
let query_lower = query.to_lowercase();
let mut models: Vec<ModelInfo> = 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<Vec<ModelVersion>> {
let mut versions: Vec<ModelVersion> = 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<DependencySpec>,
) -> HubResult<Vec<ResolvedDependency>> {
let resolver = self.version_resolver.read().await;
let result = resolver.resolve(specs)?;
Ok(result.dependencies)
}
async fn get_stats(&self) -> HubResult<RegistryStats> {
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<ModelInfo>,
options: QueryOptions,
) -> HubResult<Vec<ModelInfo>> {
// 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<ModelRegistry> {
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");
}
}