1211 lines
40 KiB
Rust
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");
|
|
}
|
|
}
|