test_model_registry_new and test_register_and_retrieve_model no longer ignored. foxhunt-postgres docker container available locally. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1349 lines
45 KiB
Rust
1349 lines
45 KiB
Rust
//! Model Registry and Versioning System
|
|
//!
|
|
//! This module provides comprehensive model versioning with metadata tracking,
|
|
//! PostgreSQL storage, and API endpoints for model management.
|
|
//!
|
|
//! # Features
|
|
//!
|
|
//! - **Model Versioning**: Semantic versioning (v1.0.0) for all ML models
|
|
//! - **Metadata Tracking**: Training metrics, hyperparameters, data sources
|
|
//! - **PostgreSQL Storage**: Persistent version history with TimescaleDB
|
|
//! - **Production Tags**: Mark models as production/experimental/archived
|
|
//! - **S3 Integration**: Store model artifacts with checksums
|
|
//! - **Query API**: Query models by version, type, tag, date range
|
|
//!
|
|
//! # Example
|
|
//!
|
|
//! ```rust,no_run
|
|
//! use ml::model_registry::{ModelRegistry, ModelVersionMetadata};
|
|
//! use ml::ModelType;
|
|
//!
|
|
//! # async fn example() -> Result<(), Box<dyn std::error::Error>> {
|
|
//! let registry = ModelRegistry::new("postgresql://...", "s3://bucket/").await?;
|
|
//!
|
|
//! // Register a new model version
|
|
//! let metadata = ModelVersionMetadata {
|
|
//! model_id: "dqn-v1.0.0".to_owned(),
|
|
//! model_type: ModelType::DQN,
|
|
//! version: "1.0.0".to_owned(),
|
|
//! hyperparameters: serde_json::json!({
|
|
//! "epochs": 500,
|
|
//! "batch_size": 128,
|
|
//! "learning_rate": 0.0001
|
|
//! }),
|
|
//! metrics: serde_json::json!({
|
|
//! "final_loss": 0.001,
|
|
//! "best_epoch": 487,
|
|
//! "training_time_seconds": 168
|
|
//! }),
|
|
//! data_source: "databento_2024_Q4".to_owned(),
|
|
//! s3_location: "s3://foxhunt-ml-models/dqn/1.0.0/".to_owned(),
|
|
//! checksum: "sha256:abc123...".to_owned(),
|
|
//! training_date: chrono::Utc::now(),
|
|
//! is_production: true,
|
|
//! is_experimental: false,
|
|
//! is_archived: false,
|
|
//! };
|
|
//!
|
|
//! registry.register_version(&metadata).await?;
|
|
//!
|
|
//! // Query models by version
|
|
//! let model = registry.get_model_by_version("dqn-v1.0.0").await?;
|
|
//!
|
|
//! // Get production models
|
|
//! let production_models = registry.get_production_models().await?;
|
|
//!
|
|
//! # Ok(())
|
|
//! # }
|
|
//! ```
|
|
|
|
use crate::{MLError, MLResult, ModelType};
|
|
use chrono::{DateTime, Utc};
|
|
use serde::{Deserialize, Serialize};
|
|
use sqlx::{postgres::PgPoolOptions, PgPool, Row};
|
|
use std::collections::HashMap;
|
|
use std::sync::Arc;
|
|
use tokio::sync::RwLock;
|
|
|
|
/// Checkpoint loading utilities
|
|
pub mod checkpoint_loader;
|
|
|
|
/// Model version metadata for tracking ML model versions
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct ModelVersionMetadata {
|
|
/// Unique model identifier (e.g., "dqn-v1.0.0")
|
|
pub model_id: String,
|
|
|
|
/// Type of the model
|
|
pub model_type: ModelType,
|
|
|
|
/// Semantic version (e.g., "1.0.0")
|
|
pub version: String,
|
|
|
|
/// Training date and time
|
|
pub training_date: DateTime<Utc>,
|
|
|
|
/// Hyperparameters used for training
|
|
pub hyperparameters: serde_json::Value,
|
|
|
|
/// Training and validation metrics
|
|
pub metrics: serde_json::Value,
|
|
|
|
/// Data source identifier (e.g., "databento_2024_Q4")
|
|
pub data_source: String,
|
|
|
|
/// S3 location for model artifacts
|
|
pub s3_location: String,
|
|
|
|
/// SHA-256 checksum of model artifacts
|
|
pub checksum: String,
|
|
|
|
/// Production status flag
|
|
pub is_production: bool,
|
|
|
|
/// Experimental status flag
|
|
pub is_experimental: bool,
|
|
|
|
/// Archived status flag
|
|
pub is_archived: bool,
|
|
|
|
/// Additional metadata
|
|
#[serde(default)]
|
|
pub metadata: HashMap<String, String>,
|
|
|
|
/// Created timestamp
|
|
#[serde(default = "Utc::now")]
|
|
pub created_at: DateTime<Utc>,
|
|
|
|
/// Last updated timestamp
|
|
#[serde(default = "Utc::now")]
|
|
pub updated_at: DateTime<Utc>,
|
|
}
|
|
|
|
impl ModelVersionMetadata {
|
|
/// Create new model version metadata
|
|
pub fn new(
|
|
model_id: String,
|
|
model_type: ModelType,
|
|
version: String,
|
|
data_source: String,
|
|
s3_location: String,
|
|
) -> Self {
|
|
Self {
|
|
model_id,
|
|
model_type,
|
|
version,
|
|
training_date: Utc::now(),
|
|
hyperparameters: serde_json::json!({}),
|
|
metrics: serde_json::json!({}),
|
|
data_source,
|
|
s3_location,
|
|
checksum: String::new(),
|
|
is_production: false,
|
|
is_experimental: true,
|
|
is_archived: false,
|
|
metadata: HashMap::new(),
|
|
created_at: Utc::now(),
|
|
updated_at: Utc::now(),
|
|
}
|
|
}
|
|
|
|
/// Mark as production model
|
|
pub fn mark_production(mut self) -> Self {
|
|
self.is_production = true;
|
|
self.is_experimental = false;
|
|
self.updated_at = Utc::now();
|
|
self
|
|
}
|
|
|
|
/// Mark as experimental model
|
|
pub fn mark_experimental(mut self) -> Self {
|
|
self.is_production = false;
|
|
self.is_experimental = true;
|
|
self.updated_at = Utc::now();
|
|
self
|
|
}
|
|
|
|
/// Mark as archived model
|
|
pub fn mark_archived(mut self) -> Self {
|
|
self.is_archived = true;
|
|
self.updated_at = Utc::now();
|
|
self
|
|
}
|
|
|
|
/// Add hyperparameter
|
|
pub fn add_hyperparameter(&mut self, key: &str, value: serde_json::Value) {
|
|
if let Some(obj) = self.hyperparameters.as_object_mut() {
|
|
obj.insert(key.to_string(), value);
|
|
}
|
|
}
|
|
|
|
/// Add metric
|
|
pub fn add_metric(&mut self, key: &str, value: serde_json::Value) {
|
|
if let Some(obj) = self.metrics.as_object_mut() {
|
|
obj.insert(key.to_string(), value);
|
|
}
|
|
}
|
|
|
|
/// Add metadata
|
|
pub fn add_metadata(&mut self, key: &str, value: String) {
|
|
self.metadata.insert(key.to_string(), value);
|
|
}
|
|
|
|
/// Set checksum
|
|
pub fn set_checksum(&mut self, checksum: String) {
|
|
self.checksum = checksum;
|
|
}
|
|
}
|
|
|
|
/// Model registry for managing ML model versions
|
|
#[derive(Debug, Clone)]
|
|
pub struct ModelRegistry {
|
|
/// PostgreSQL connection pool
|
|
db_pool: PgPool,
|
|
|
|
/// S3 bucket base path
|
|
s3_base_path: String,
|
|
|
|
/// In-memory cache for fast lookups
|
|
cache: Arc<RwLock<HashMap<String, ModelVersionMetadata>>>,
|
|
}
|
|
|
|
impl ModelRegistry {
|
|
/// Create new model registry
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `database_url` - PostgreSQL connection URL
|
|
/// * `s3_base_path` - S3 bucket base path (e.g., "s3://foxhunt-ml-models/")
|
|
///
|
|
/// # Returns
|
|
///
|
|
/// Returns a `Result<ModelRegistry>` with the initialized registry
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns an error if database connection or schema creation fails
|
|
pub async fn new(database_url: &str, s3_base_path: &str) -> MLResult<Self> {
|
|
let db_pool = PgPoolOptions::new()
|
|
.max_connections(5)
|
|
.connect(database_url)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to connect to database: {}", e)))?;
|
|
|
|
// Create schema if not exists
|
|
Self::ensure_schema(&db_pool).await?;
|
|
|
|
Ok(Self {
|
|
db_pool,
|
|
s3_base_path: s3_base_path.to_string(),
|
|
cache: Arc::new(RwLock::new(HashMap::new())),
|
|
})
|
|
}
|
|
|
|
/// Ensure database schema exists
|
|
async fn ensure_schema(pool: &PgPool) -> MLResult<()> {
|
|
// Create table
|
|
sqlx::query(
|
|
"
|
|
CREATE TABLE IF NOT EXISTS ml_model_versions (
|
|
id SERIAL PRIMARY KEY,
|
|
model_id VARCHAR(255) NOT NULL UNIQUE,
|
|
model_type VARCHAR(50) NOT NULL,
|
|
version VARCHAR(50) NOT NULL,
|
|
training_date TIMESTAMPTZ NOT NULL,
|
|
hyperparameters JSONB NOT NULL DEFAULT '{}'::jsonb,
|
|
metrics JSONB NOT NULL DEFAULT '{}'::jsonb,
|
|
data_source VARCHAR(255) NOT NULL,
|
|
s3_location TEXT NOT NULL,
|
|
checksum VARCHAR(255) NOT NULL,
|
|
is_production BOOLEAN NOT NULL DEFAULT false,
|
|
is_experimental BOOLEAN NOT NULL DEFAULT true,
|
|
is_archived BOOLEAN NOT NULL DEFAULT false,
|
|
metadata JSONB NOT NULL DEFAULT '{}'::jsonb,
|
|
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
CONSTRAINT unique_model_version UNIQUE (model_type, version)
|
|
)
|
|
",
|
|
)
|
|
.execute(pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to create schema: {}", e)))?;
|
|
|
|
// Create indexes (each as separate statement)
|
|
let indexes = vec![
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_model_type ON ml_model_versions(model_type)",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_version ON ml_model_versions(version)",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_training_date ON ml_model_versions(training_date DESC)",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_is_production ON ml_model_versions(is_production) WHERE is_production = true",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_is_experimental ON ml_model_versions(is_experimental) WHERE is_experimental = true",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_is_archived ON ml_model_versions(is_archived) WHERE is_archived = false",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_metadata_gin ON ml_model_versions USING GIN (metadata)",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_hyperparameters_gin ON ml_model_versions USING GIN (hyperparameters)",
|
|
"CREATE INDEX IF NOT EXISTS idx_ml_model_versions_metrics_gin ON ml_model_versions USING GIN (metrics)",
|
|
];
|
|
|
|
for index_query in indexes {
|
|
sqlx::query(index_query)
|
|
.execute(pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to create index: {}", e)))?;
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Register a new model version
|
|
pub async fn register_version(&self, metadata: &ModelVersionMetadata) -> MLResult<()> {
|
|
let query = "
|
|
INSERT INTO ml_model_versions (
|
|
model_id, model_type, version, training_date,
|
|
hyperparameters, metrics, data_source, s3_location,
|
|
checksum, is_production, is_experimental, is_archived, metadata
|
|
) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13)
|
|
ON CONFLICT (model_id)
|
|
DO UPDATE SET
|
|
training_date = EXCLUDED.training_date,
|
|
hyperparameters = EXCLUDED.hyperparameters,
|
|
metrics = EXCLUDED.metrics,
|
|
data_source = EXCLUDED.data_source,
|
|
s3_location = EXCLUDED.s3_location,
|
|
checksum = EXCLUDED.checksum,
|
|
is_production = EXCLUDED.is_production,
|
|
is_experimental = EXCLUDED.is_experimental,
|
|
is_archived = EXCLUDED.is_archived,
|
|
metadata = EXCLUDED.metadata,
|
|
updated_at = NOW()
|
|
";
|
|
|
|
let model_type_str = format!("{:?}", metadata.model_type);
|
|
|
|
sqlx::query(query)
|
|
.bind(&metadata.model_id)
|
|
.bind(&model_type_str)
|
|
.bind(&metadata.version)
|
|
.bind(metadata.training_date)
|
|
.bind(&metadata.hyperparameters)
|
|
.bind(&metadata.metrics)
|
|
.bind(&metadata.data_source)
|
|
.bind(&metadata.s3_location)
|
|
.bind(&metadata.checksum)
|
|
.bind(metadata.is_production)
|
|
.bind(metadata.is_experimental)
|
|
.bind(metadata.is_archived)
|
|
.bind(serde_json::to_value(&metadata.metadata).unwrap_or_default())
|
|
.execute(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to register version: {}", e)))?;
|
|
|
|
// Update cache
|
|
self.cache
|
|
.write()
|
|
.await
|
|
.insert(metadata.model_id.clone(), metadata.clone());
|
|
|
|
tracing::info!("Registered model version: {}", metadata.model_id);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Get model by version
|
|
pub async fn get_model_by_version(&self, model_id: &str) -> MLResult<ModelVersionMetadata> {
|
|
// Check cache first
|
|
if let Some(metadata) = self.cache.read().await.get(model_id) {
|
|
return Ok(metadata.clone());
|
|
}
|
|
|
|
// Query database
|
|
let query = "
|
|
SELECT model_id, model_type, version, training_date,
|
|
hyperparameters, metrics, data_source, s3_location,
|
|
checksum, is_production, is_experimental, is_archived,
|
|
metadata, created_at, updated_at
|
|
FROM ml_model_versions
|
|
WHERE model_id = $1
|
|
";
|
|
|
|
let row = sqlx::query(query)
|
|
.bind(model_id)
|
|
.fetch_one(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelNotFound(format!("Model {} not found: {}", model_id, e)))?;
|
|
|
|
let metadata = self.row_to_metadata(row)?;
|
|
|
|
// Update cache
|
|
self.cache
|
|
.write()
|
|
.await
|
|
.insert(model_id.to_string(), metadata.clone());
|
|
|
|
Ok(metadata)
|
|
}
|
|
|
|
/// Get all production models
|
|
pub async fn get_production_models(&self) -> MLResult<Vec<ModelVersionMetadata>> {
|
|
let query = "
|
|
SELECT model_id, model_type, version, training_date,
|
|
hyperparameters, metrics, data_source, s3_location,
|
|
checksum, is_production, is_experimental, is_archived,
|
|
metadata, created_at, updated_at
|
|
FROM ml_model_versions
|
|
WHERE is_production = true AND is_archived = false
|
|
ORDER BY training_date DESC
|
|
";
|
|
|
|
let rows = sqlx::query(query)
|
|
.fetch_all(&self.db_pool)
|
|
.await
|
|
.map_err(|e| {
|
|
MLError::ModelError(format!("Failed to query production models: {}", e))
|
|
})?;
|
|
|
|
let mut models = Vec::new();
|
|
for row in rows {
|
|
models.push(self.row_to_metadata(row)?);
|
|
}
|
|
|
|
Ok(models)
|
|
}
|
|
|
|
/// Get all experimental models
|
|
pub async fn get_experimental_models(&self) -> MLResult<Vec<ModelVersionMetadata>> {
|
|
let query = "
|
|
SELECT model_id, model_type, version, training_date,
|
|
hyperparameters, metrics, data_source, s3_location,
|
|
checksum, is_production, is_experimental, is_archived,
|
|
metadata, created_at, updated_at
|
|
FROM ml_model_versions
|
|
WHERE is_experimental = true AND is_archived = false
|
|
ORDER BY training_date DESC
|
|
";
|
|
|
|
let rows = sqlx::query(query)
|
|
.fetch_all(&self.db_pool)
|
|
.await
|
|
.map_err(|e| {
|
|
MLError::ModelError(format!("Failed to query experimental models: {}", e))
|
|
})?;
|
|
|
|
let mut models = Vec::new();
|
|
for row in rows {
|
|
models.push(self.row_to_metadata(row)?);
|
|
}
|
|
|
|
Ok(models)
|
|
}
|
|
|
|
/// Get models by type
|
|
pub async fn get_models_by_type(
|
|
&self,
|
|
model_type: ModelType,
|
|
) -> MLResult<Vec<ModelVersionMetadata>> {
|
|
let query = "
|
|
SELECT model_id, model_type, version, training_date,
|
|
hyperparameters, metrics, data_source, s3_location,
|
|
checksum, is_production, is_experimental, is_archived,
|
|
metadata, created_at, updated_at
|
|
FROM ml_model_versions
|
|
WHERE model_type = $1 AND is_archived = false
|
|
ORDER BY training_date DESC
|
|
";
|
|
|
|
let model_type_str = format!("{:?}", model_type);
|
|
|
|
let rows = sqlx::query(query)
|
|
.bind(&model_type_str)
|
|
.fetch_all(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to query models by type: {}", e)))?;
|
|
|
|
let mut models = Vec::new();
|
|
for row in rows {
|
|
models.push(self.row_to_metadata(row)?);
|
|
}
|
|
|
|
Ok(models)
|
|
}
|
|
|
|
/// Get models by date range
|
|
pub async fn get_models_by_date_range(
|
|
&self,
|
|
start_date: DateTime<Utc>,
|
|
end_date: DateTime<Utc>,
|
|
) -> MLResult<Vec<ModelVersionMetadata>> {
|
|
let query = "
|
|
SELECT model_id, model_type, version, training_date,
|
|
hyperparameters, metrics, data_source, s3_location,
|
|
checksum, is_production, is_experimental, is_archived,
|
|
metadata, created_at, updated_at
|
|
FROM ml_model_versions
|
|
WHERE training_date >= $1 AND training_date <= $2
|
|
ORDER BY training_date DESC
|
|
";
|
|
|
|
let rows = sqlx::query(query)
|
|
.bind(start_date)
|
|
.bind(end_date)
|
|
.fetch_all(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to query models by date: {}", e)))?;
|
|
|
|
let mut models = Vec::new();
|
|
for row in rows {
|
|
models.push(self.row_to_metadata(row)?);
|
|
}
|
|
|
|
Ok(models)
|
|
}
|
|
|
|
/// Mark model as production
|
|
pub async fn mark_production(&self, model_id: &str) -> MLResult<()> {
|
|
let query = "
|
|
UPDATE ml_model_versions
|
|
SET is_production = true, is_experimental = false, updated_at = NOW()
|
|
WHERE model_id = $1
|
|
";
|
|
|
|
sqlx::query(query)
|
|
.bind(model_id)
|
|
.execute(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to mark production: {}", e)))?;
|
|
|
|
// Invalidate cache
|
|
self.cache.write().await.remove(model_id);
|
|
|
|
tracing::info!("Marked model as production: {}", model_id);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Mark model as experimental
|
|
pub async fn mark_experimental(&self, model_id: &str) -> MLResult<()> {
|
|
let query = "
|
|
UPDATE ml_model_versions
|
|
SET is_production = false, is_experimental = true, updated_at = NOW()
|
|
WHERE model_id = $1
|
|
";
|
|
|
|
sqlx::query(query)
|
|
.bind(model_id)
|
|
.execute(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to mark experimental: {}", e)))?;
|
|
|
|
// Invalidate cache
|
|
self.cache.write().await.remove(model_id);
|
|
|
|
tracing::info!("Marked model as experimental: {}", model_id);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Archive model
|
|
pub async fn archive_model(&self, model_id: &str) -> MLResult<()> {
|
|
let query = "
|
|
UPDATE ml_model_versions
|
|
SET is_archived = true, updated_at = NOW()
|
|
WHERE model_id = $1
|
|
";
|
|
|
|
sqlx::query(query)
|
|
.bind(model_id)
|
|
.execute(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to archive model: {}", e)))?;
|
|
|
|
// Invalidate cache
|
|
self.cache.write().await.remove(model_id);
|
|
|
|
tracing::info!("Archived model: {}", model_id);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Delete model version (permanent)
|
|
pub async fn delete_version(&self, model_id: &str) -> MLResult<()> {
|
|
let query = "
|
|
DELETE FROM ml_model_versions
|
|
WHERE model_id = $1
|
|
";
|
|
|
|
sqlx::query(query)
|
|
.bind(model_id)
|
|
.execute(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to delete version: {}", e)))?;
|
|
|
|
// Invalidate cache
|
|
self.cache.write().await.remove(model_id);
|
|
|
|
tracing::warn!("Deleted model version: {}", model_id);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Get registry statistics
|
|
pub async fn get_statistics(&self) -> MLResult<RegistryStatistics> {
|
|
let query = "
|
|
SELECT
|
|
COUNT(*) FILTER (WHERE is_production = true AND is_archived = false) as production_count,
|
|
COUNT(*) FILTER (WHERE is_experimental = true AND is_archived = false) as experimental_count,
|
|
COUNT(*) FILTER (WHERE is_archived = true) as archived_count,
|
|
COUNT(*) as total_count,
|
|
COUNT(DISTINCT model_type) as model_types_count,
|
|
MAX(training_date) as latest_training_date,
|
|
MIN(training_date) as earliest_training_date
|
|
FROM ml_model_versions
|
|
";
|
|
|
|
let row = sqlx::query(query)
|
|
.fetch_one(&self.db_pool)
|
|
.await
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get statistics: {}", e)))?;
|
|
|
|
Ok(RegistryStatistics {
|
|
production_count: row.try_get::<i64, _>("production_count").unwrap_or(0) as usize,
|
|
experimental_count: row.try_get::<i64, _>("experimental_count").unwrap_or(0) as usize,
|
|
archived_count: row.try_get::<i64, _>("archived_count").unwrap_or(0) as usize,
|
|
total_count: row.try_get::<i64, _>("total_count").unwrap_or(0) as usize,
|
|
model_types_count: row.try_get::<i64, _>("model_types_count").unwrap_or(0) as usize,
|
|
latest_training_date: row.try_get("latest_training_date").ok(),
|
|
earliest_training_date: row.try_get("earliest_training_date").ok(),
|
|
})
|
|
}
|
|
|
|
/// Convert database row to metadata
|
|
fn row_to_metadata(&self, row: sqlx::postgres::PgRow) -> MLResult<ModelVersionMetadata> {
|
|
let model_type_str: String = row
|
|
.try_get("model_type")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get model_type: {}", e)))?;
|
|
|
|
let model_type = ModelType::from_str(&model_type_str).ok_or_else(|| {
|
|
MLError::ModelError(format!("Invalid model type: {}", model_type_str))
|
|
})?;
|
|
|
|
let metadata_json: serde_json::Value = row
|
|
.try_get("metadata")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get metadata: {}", e)))?;
|
|
|
|
let metadata_map: HashMap<String, String> =
|
|
serde_json::from_value(metadata_json).unwrap_or_default();
|
|
|
|
Ok(ModelVersionMetadata {
|
|
model_id: row
|
|
.try_get("model_id")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get model_id: {}", e)))?,
|
|
model_type,
|
|
version: row
|
|
.try_get("version")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get version: {}", e)))?,
|
|
training_date: row
|
|
.try_get("training_date")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get training_date: {}", e)))?,
|
|
hyperparameters: row.try_get("hyperparameters").map_err(|e| {
|
|
MLError::ModelError(format!("Failed to get hyperparameters: {}", e))
|
|
})?,
|
|
metrics: row
|
|
.try_get("metrics")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get metrics: {}", e)))?,
|
|
data_source: row
|
|
.try_get("data_source")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get data_source: {}", e)))?,
|
|
s3_location: row
|
|
.try_get("s3_location")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get s3_location: {}", e)))?,
|
|
checksum: row
|
|
.try_get("checksum")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get checksum: {}", e)))?,
|
|
is_production: row
|
|
.try_get("is_production")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get is_production: {}", e)))?,
|
|
is_experimental: row.try_get("is_experimental").map_err(|e| {
|
|
MLError::ModelError(format!("Failed to get is_experimental: {}", e))
|
|
})?,
|
|
is_archived: row
|
|
.try_get("is_archived")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get is_archived: {}", e)))?,
|
|
metadata: metadata_map,
|
|
created_at: row
|
|
.try_get("created_at")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get created_at: {}", e)))?,
|
|
updated_at: row
|
|
.try_get("updated_at")
|
|
.map_err(|e| MLError::ModelError(format!("Failed to get updated_at: {}", e)))?,
|
|
})
|
|
}
|
|
}
|
|
|
|
/// Registry statistics
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct RegistryStatistics {
|
|
/// Number of production models
|
|
pub production_count: usize,
|
|
/// Number of experimental models
|
|
pub experimental_count: usize,
|
|
/// Number of archived models
|
|
pub archived_count: usize,
|
|
/// Total number of models
|
|
pub total_count: usize,
|
|
/// Number of distinct model types
|
|
pub model_types_count: usize,
|
|
/// Latest training date
|
|
pub latest_training_date: Option<DateTime<Utc>>,
|
|
/// Earliest training date
|
|
pub earliest_training_date: Option<DateTime<Utc>>,
|
|
}
|
|
|
|
/// A rollback event recorded in the audit log
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct RollbackEvent {
|
|
/// Model name that was rolled back
|
|
pub model_name: String,
|
|
/// Version that was active before the rollback
|
|
pub from_version: String,
|
|
/// Version that became active after the rollback
|
|
pub to_version: String,
|
|
/// Timestamp of the rollback
|
|
pub timestamp: DateTime<Utc>,
|
|
/// Optional reason for the rollback
|
|
pub reason: Option<String>,
|
|
}
|
|
|
|
/// Version entry stored per model in the in-memory registry
|
|
#[derive(Debug, Clone)]
|
|
struct VersionEntry {
|
|
/// All registered versions (ordered by registration time)
|
|
versions: Vec<ModelVersionMetadata>,
|
|
/// Index into `versions` for the currently active version
|
|
active_index: usize,
|
|
}
|
|
|
|
/// In-memory model registry with version tracking and rollback support.
|
|
///
|
|
/// This registry does not require a database. It stores all model versions
|
|
/// in memory, tracks which version is "active" for each model name, and
|
|
/// records rollback events in an audit log.
|
|
#[derive(Debug, Clone)]
|
|
pub struct InMemoryModelRegistry {
|
|
/// Per-model version entries keyed by model name (e.g. "dqn", "ppo")
|
|
models: Arc<RwLock<HashMap<String, VersionEntry>>>,
|
|
/// Audit log of rollback events
|
|
rollback_log: Arc<RwLock<Vec<RollbackEvent>>>,
|
|
}
|
|
|
|
impl InMemoryModelRegistry {
|
|
/// Create a new empty in-memory model registry.
|
|
pub fn new() -> Self {
|
|
Self {
|
|
models: Arc::new(RwLock::new(HashMap::new())),
|
|
rollback_log: Arc::new(RwLock::new(Vec::new())),
|
|
}
|
|
}
|
|
|
|
/// Register a new model version.
|
|
///
|
|
/// The first version registered for a given model name automatically becomes
|
|
/// the active version. Subsequent registrations are stored but do not change
|
|
/// the active version (use [`rollback_to_version`] or [`promote_version`] for that).
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `model_name` - Logical model name (e.g. "dqn", "ppo")
|
|
/// * `metadata` - Full version metadata
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns `MLError::ModelError` if a version with the same version string
|
|
/// is already registered for this model name.
|
|
pub async fn register_version(
|
|
&self,
|
|
model_name: &str,
|
|
metadata: ModelVersionMetadata,
|
|
) -> MLResult<()> {
|
|
let mut models = self.models.write().await;
|
|
let entry = models
|
|
.entry(model_name.to_string())
|
|
.or_insert_with(|| VersionEntry {
|
|
versions: Vec::new(),
|
|
active_index: 0,
|
|
});
|
|
|
|
// Check for duplicate version strings
|
|
let version_str = metadata.version.clone();
|
|
for existing in &entry.versions {
|
|
if existing.version == version_str {
|
|
return Err(MLError::ModelError(format!(
|
|
"Version {} already registered for model {}",
|
|
version_str, model_name
|
|
)));
|
|
}
|
|
}
|
|
|
|
entry.versions.push(metadata);
|
|
|
|
// First version auto-becomes active (active_index is already 0)
|
|
// Subsequent versions do not change active_index
|
|
|
|
tracing::info!(
|
|
model_name = model_name,
|
|
version = version_str.as_str(),
|
|
total_versions = entry.versions.len(),
|
|
"Registered model version"
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Roll back a model to a previously registered version.
|
|
///
|
|
/// This sets the specified version as the active version and records
|
|
/// a rollback event in the audit log.
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `model_name` - Logical model name
|
|
/// * `target_version` - Semantic version string to roll back to
|
|
/// * `reason` - Optional reason for the rollback
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns `MLError::ModelNotFound` if the model name or target version
|
|
/// is not found in the registry.
|
|
pub async fn rollback_to_version(
|
|
&self,
|
|
model_name: &str,
|
|
target_version: &str,
|
|
reason: Option<String>,
|
|
) -> MLResult<()> {
|
|
let mut models = self.models.write().await;
|
|
let entry = models.get_mut(model_name).ok_or_else(|| {
|
|
MLError::ModelNotFound(format!("Model {} not found in registry", model_name))
|
|
})?;
|
|
|
|
// Find the target version index
|
|
let target_index = entry
|
|
.versions
|
|
.iter()
|
|
.position(|v| v.version == target_version)
|
|
.ok_or_else(|| {
|
|
MLError::ModelNotFound(format!(
|
|
"Version {} not found for model {}",
|
|
target_version, model_name
|
|
))
|
|
})?;
|
|
|
|
let from_version = entry
|
|
.versions
|
|
.get(entry.active_index)
|
|
.map(|v| v.version.clone())
|
|
.unwrap_or_default();
|
|
|
|
if entry.active_index == target_index {
|
|
tracing::warn!(
|
|
model_name = model_name,
|
|
version = target_version,
|
|
"Rollback requested to already-active version (no-op)"
|
|
);
|
|
return Ok(());
|
|
}
|
|
|
|
entry.active_index = target_index;
|
|
|
|
tracing::info!(
|
|
model_name = model_name,
|
|
from_version = from_version.as_str(),
|
|
to_version = target_version,
|
|
reason = reason.as_deref().unwrap_or("none"),
|
|
"Rolled back model version"
|
|
);
|
|
|
|
// Record rollback event
|
|
let event = RollbackEvent {
|
|
model_name: model_name.to_string(),
|
|
from_version,
|
|
to_version: target_version.to_string(),
|
|
timestamp: Utc::now(),
|
|
reason,
|
|
};
|
|
// Drop models lock before acquiring rollback_log lock to avoid deadlock
|
|
drop(models);
|
|
self.rollback_log.write().await.push(event);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Promote a version to be the active version (same as rollback but with
|
|
/// clearer semantics for forward version changes).
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns `MLError::ModelNotFound` if the model or version is not found.
|
|
pub async fn promote_version(
|
|
&self,
|
|
model_name: &str,
|
|
target_version: &str,
|
|
) -> MLResult<()> {
|
|
self.rollback_to_version(model_name, target_version, Some("promoted".to_owned()))
|
|
.await
|
|
}
|
|
|
|
/// List all registered versions for a model, ordered by registration time.
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `model_name` - Logical model name
|
|
///
|
|
/// # Returns
|
|
///
|
|
/// A vector of `(version_string, is_active)` tuples.
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns `MLError::ModelNotFound` if the model name is not found.
|
|
pub async fn list_versions(
|
|
&self,
|
|
model_name: &str,
|
|
) -> MLResult<Vec<(String, bool)>> {
|
|
let models = self.models.read().await;
|
|
let entry = models.get(model_name).ok_or_else(|| {
|
|
MLError::ModelNotFound(format!("Model {} not found in registry", model_name))
|
|
})?;
|
|
|
|
let result = entry
|
|
.versions
|
|
.iter()
|
|
.enumerate()
|
|
.map(|(i, v)| (v.version.clone(), i == entry.active_index))
|
|
.collect();
|
|
|
|
Ok(result)
|
|
}
|
|
|
|
/// Get the active version metadata for a model.
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns `MLError::ModelNotFound` if the model is not found or has no versions.
|
|
pub async fn get_active_version(
|
|
&self,
|
|
model_name: &str,
|
|
) -> MLResult<ModelVersionMetadata> {
|
|
let models = self.models.read().await;
|
|
let entry = models.get(model_name).ok_or_else(|| {
|
|
MLError::ModelNotFound(format!("Model {} not found in registry", model_name))
|
|
})?;
|
|
|
|
entry
|
|
.versions
|
|
.get(entry.active_index)
|
|
.cloned()
|
|
.ok_or_else(|| {
|
|
MLError::ModelNotFound(format!("No versions registered for model {}", model_name))
|
|
})
|
|
}
|
|
|
|
/// Get a specific version's metadata for a model.
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns `MLError::ModelNotFound` if the model or version is not found.
|
|
pub async fn get_version(
|
|
&self,
|
|
model_name: &str,
|
|
version: &str,
|
|
) -> MLResult<ModelVersionMetadata> {
|
|
let models = self.models.read().await;
|
|
let entry = models.get(model_name).ok_or_else(|| {
|
|
MLError::ModelNotFound(format!("Model {} not found in registry", model_name))
|
|
})?;
|
|
|
|
entry
|
|
.versions
|
|
.iter()
|
|
.find(|v| v.version == version)
|
|
.cloned()
|
|
.ok_or_else(|| {
|
|
MLError::ModelNotFound(format!(
|
|
"Version {} not found for model {}",
|
|
version, model_name
|
|
))
|
|
})
|
|
}
|
|
|
|
/// Get all model names in the registry.
|
|
pub async fn list_models(&self) -> Vec<String> {
|
|
self.models.read().await.keys().cloned().collect()
|
|
}
|
|
|
|
/// Get the rollback audit log.
|
|
pub async fn get_rollback_log(&self) -> Vec<RollbackEvent> {
|
|
self.rollback_log.read().await.clone()
|
|
}
|
|
|
|
/// Get rollback events for a specific model.
|
|
pub async fn get_rollback_log_for_model(&self, model_name: &str) -> Vec<RollbackEvent> {
|
|
self.rollback_log
|
|
.read()
|
|
.await
|
|
.iter()
|
|
.filter(|e| e.model_name == model_name)
|
|
.cloned()
|
|
.collect()
|
|
}
|
|
}
|
|
|
|
impl Default for InMemoryModelRegistry {
|
|
fn default() -> Self {
|
|
Self::new()
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[tokio::test]
|
|
async fn test_model_registry_new() {
|
|
let registry = ModelRegistry::new(
|
|
"postgresql://foxhunt:foxhunt_dev_password@localhost:5432/foxhunt",
|
|
"s3://foxhunt-ml-models/",
|
|
)
|
|
.await;
|
|
|
|
assert!(registry.is_ok());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_register_and_retrieve_model() {
|
|
let registry = ModelRegistry::new(
|
|
"postgresql://foxhunt:foxhunt_dev_password@localhost:5432/foxhunt",
|
|
"s3://foxhunt-ml-models/",
|
|
)
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"dqn-test-v1.0.0".to_owned(),
|
|
ModelType::DQN,
|
|
"1.0.0".to_owned(),
|
|
"test_data".to_owned(),
|
|
"s3://foxhunt-ml-models/dqn/1.0.0/".to_owned(),
|
|
);
|
|
|
|
metadata.add_hyperparameter("epochs", serde_json::json!(500));
|
|
metadata.add_metric("final_loss", serde_json::json!(0.001));
|
|
metadata.set_checksum("sha256:test123".to_owned());
|
|
|
|
// Register
|
|
registry.register_version(&metadata).await.unwrap();
|
|
|
|
// Retrieve
|
|
let retrieved = registry
|
|
.get_model_by_version("dqn-test-v1.0.0")
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(retrieved.model_id, "dqn-test-v1.0.0");
|
|
assert_eq!(retrieved.version, "1.0.0");
|
|
}
|
|
|
|
// ---- In-memory registry tests (no database required) ----
|
|
|
|
fn make_version(model_id: &str, version: &str) -> ModelVersionMetadata {
|
|
ModelVersionMetadata::new(
|
|
model_id.to_string(),
|
|
ModelType::DQN,
|
|
version.to_string(),
|
|
"test_data".to_owned(),
|
|
format!("s3://models/{}/{}/", model_id, version),
|
|
)
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_register_and_get_active() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
|
|
let active = registry.get_active_version("dqn").await.unwrap();
|
|
assert_eq!(active.version, "1.0.0");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_first_version_is_active() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
let v2 = make_version("dqn-v2", "2.0.0");
|
|
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
registry.register_version("dqn", v2).await.unwrap();
|
|
|
|
// First registered version stays active
|
|
let active = registry.get_active_version("dqn").await.unwrap();
|
|
assert_eq!(active.version, "1.0.0");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_duplicate_version_rejected() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
let v1_dup = make_version("dqn-v1-dup", "1.0.0");
|
|
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
let result = registry.register_version("dqn", v1_dup).await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_list_versions() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
let v2 = make_version("dqn-v2", "2.0.0");
|
|
let v3 = make_version("dqn-v3", "3.0.0");
|
|
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
registry.register_version("dqn", v2).await.unwrap();
|
|
registry.register_version("dqn", v3).await.unwrap();
|
|
|
|
let versions = registry.list_versions("dqn").await.unwrap();
|
|
assert_eq!(versions.len(), 3);
|
|
assert_eq!(versions.first().map(|v| v.0.as_str()), Some("1.0.0"));
|
|
assert_eq!(versions.first().map(|v| v.1), Some(true)); // active
|
|
assert_eq!(versions.get(1).map(|v| v.0.as_str()), Some("2.0.0"));
|
|
assert_eq!(versions.get(1).map(|v| v.1), Some(false)); // not active
|
|
assert_eq!(versions.get(2).map(|v| v.0.as_str()), Some("3.0.0"));
|
|
assert_eq!(versions.get(2).map(|v| v.1), Some(false)); // not active
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_list_versions_unknown_model() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
let result = registry.list_versions("nonexistent").await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_rollback_to_version() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
let v2 = make_version("dqn-v2", "2.0.0");
|
|
let v3 = make_version("dqn-v3", "3.0.0");
|
|
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
registry.register_version("dqn", v2).await.unwrap();
|
|
registry.register_version("dqn", v3).await.unwrap();
|
|
|
|
// Promote to v3 first
|
|
registry.promote_version("dqn", "3.0.0").await.unwrap();
|
|
let active = registry.get_active_version("dqn").await.unwrap();
|
|
assert_eq!(active.version, "3.0.0");
|
|
|
|
// Roll back to v1
|
|
registry
|
|
.rollback_to_version("dqn", "1.0.0", Some("regression in v3".to_owned()))
|
|
.await
|
|
.unwrap();
|
|
|
|
let active = registry.get_active_version("dqn").await.unwrap();
|
|
assert_eq!(active.version, "1.0.0");
|
|
|
|
// Check rollback log
|
|
let log = registry.get_rollback_log().await;
|
|
// Two events: promote to v3 and rollback to v1
|
|
assert_eq!(log.len(), 2);
|
|
|
|
let last = log.get(1);
|
|
assert!(last.is_some());
|
|
if let Some(event) = last {
|
|
assert_eq!(event.model_name, "dqn");
|
|
assert_eq!(event.from_version, "3.0.0");
|
|
assert_eq!(event.to_version, "1.0.0");
|
|
assert_eq!(event.reason.as_deref(), Some("regression in v3"));
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_rollback_unknown_model() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
let result = registry
|
|
.rollback_to_version("nonexistent", "1.0.0", None)
|
|
.await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_rollback_unknown_version() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
|
|
let result = registry
|
|
.rollback_to_version("dqn", "99.0.0", None)
|
|
.await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_rollback_to_same_version_is_noop() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
|
|
// Rolling back to already-active version succeeds silently
|
|
registry
|
|
.rollback_to_version("dqn", "1.0.0", None)
|
|
.await
|
|
.unwrap();
|
|
|
|
// No rollback event recorded for no-op
|
|
let log = registry.get_rollback_log().await;
|
|
assert!(log.is_empty());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_get_version() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
let v2 = make_version("dqn-v2", "2.0.0");
|
|
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
registry.register_version("dqn", v2).await.unwrap();
|
|
|
|
let retrieved = registry.get_version("dqn", "2.0.0").await.unwrap();
|
|
assert_eq!(retrieved.version, "2.0.0");
|
|
assert_eq!(retrieved.model_id, "dqn-v2");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_get_version_not_found() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
|
|
let result = registry.get_version("dqn", "99.0.0").await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_list_models() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let dqn = make_version("dqn-v1", "1.0.0");
|
|
let ppo = make_version("ppo-v1", "1.0.0");
|
|
|
|
registry.register_version("dqn", dqn).await.unwrap();
|
|
registry.register_version("ppo", ppo).await.unwrap();
|
|
|
|
let mut models = registry.list_models().await;
|
|
models.sort();
|
|
assert_eq!(models, vec!["dqn", "ppo"]);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_rollback_log_per_model() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let dqn_v1 = make_version("dqn-v1", "1.0.0");
|
|
let dqn_v2 = make_version("dqn-v2", "2.0.0");
|
|
let ppo_v1 = make_version("ppo-v1", "1.0.0");
|
|
let ppo_v2 = make_version("ppo-v2", "2.0.0");
|
|
|
|
registry.register_version("dqn", dqn_v1).await.unwrap();
|
|
registry.register_version("dqn", dqn_v2).await.unwrap();
|
|
registry.register_version("ppo", ppo_v1).await.unwrap();
|
|
registry.register_version("ppo", ppo_v2).await.unwrap();
|
|
|
|
// Roll back both
|
|
registry
|
|
.rollback_to_version("dqn", "2.0.0", None)
|
|
.await
|
|
.unwrap();
|
|
registry
|
|
.rollback_to_version("ppo", "2.0.0", None)
|
|
.await
|
|
.unwrap();
|
|
|
|
// Filter by model
|
|
let dqn_log = registry.get_rollback_log_for_model("dqn").await;
|
|
assert_eq!(dqn_log.len(), 1);
|
|
assert_eq!(dqn_log.first().map(|e| e.model_name.as_str()), Some("dqn"));
|
|
|
|
let ppo_log = registry.get_rollback_log_for_model("ppo").await;
|
|
assert_eq!(ppo_log.len(), 1);
|
|
assert_eq!(ppo_log.first().map(|e| e.model_name.as_str()), Some("ppo"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_default_trait() {
|
|
let registry = InMemoryModelRegistry::default();
|
|
let models = registry.list_models().await;
|
|
assert!(models.is_empty());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_multiple_rollbacks() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
let v2 = make_version("dqn-v2", "2.0.0");
|
|
let v3 = make_version("dqn-v3", "3.0.0");
|
|
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
registry.register_version("dqn", v2).await.unwrap();
|
|
registry.register_version("dqn", v3).await.unwrap();
|
|
|
|
// v1 -> v3 -> v2 -> v1 -> v3
|
|
registry.promote_version("dqn", "3.0.0").await.unwrap();
|
|
registry
|
|
.rollback_to_version("dqn", "2.0.0", None)
|
|
.await
|
|
.unwrap();
|
|
registry
|
|
.rollback_to_version("dqn", "1.0.0", None)
|
|
.await
|
|
.unwrap();
|
|
registry.promote_version("dqn", "3.0.0").await.unwrap();
|
|
|
|
let active = registry.get_active_version("dqn").await.unwrap();
|
|
assert_eq!(active.version, "3.0.0");
|
|
|
|
let log = registry.get_rollback_log().await;
|
|
assert_eq!(log.len(), 4);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_inmemory_version_list_reflects_active_after_rollback() {
|
|
let registry = InMemoryModelRegistry::new();
|
|
|
|
let v1 = make_version("dqn-v1", "1.0.0");
|
|
let v2 = make_version("dqn-v2", "2.0.0");
|
|
|
|
registry.register_version("dqn", v1).await.unwrap();
|
|
registry.register_version("dqn", v2).await.unwrap();
|
|
|
|
// Initially v1 is active
|
|
let versions = registry.list_versions("dqn").await.unwrap();
|
|
assert_eq!(versions.first().map(|v| v.1), Some(true));
|
|
assert_eq!(versions.get(1).map(|v| v.1), Some(false));
|
|
|
|
// Promote v2
|
|
registry.promote_version("dqn", "2.0.0").await.unwrap();
|
|
|
|
let versions = registry.list_versions("dqn").await.unwrap();
|
|
assert_eq!(versions.first().map(|v| v.1), Some(false));
|
|
assert_eq!(versions.get(1).map(|v| v.1), Some(true));
|
|
}
|
|
}
|