Files
foxhunt/crates/ml/src/model_registry.rs
jgrusewski fd865ccf8a fix: unignore 2 PostgreSQL model registry tests — docker running
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>
2026-03-18 17:20:23 +01:00

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));
}
}