use rtx_auto::{ error::AutoError, proposal::{Proposal, ProposalType}, rollback::{Checkpoint, CheckpointType, RollbackManager}, }; use rtx_runtime::Runtime; use std::collections::HashMap; use std::sync::Arc; #[tokio::test] async fn test_rollback_manager_creation() { let runtime = Arc::new(Runtime::new().expect("Failed to create runtime")); let mut manager = RollbackManager::new(runtime.clone()); assert!(manager.is_ok(), "Failed to create RollbackManager"); } #[tokio::test] async fn test_create_checkpoint() { let runtime = Arc::new(Runtime::new().expect("Failed to create runtime")); let mut manager = RollbackManager::new(runtime.clone()).unwrap(); let mut state = HashMap::new(); state.insert("model_weights".to_string(), vec![1.0, 2.0, 3.0]); state.insert("optimizer_state".to_string(), vec![0.1, 0.2, 0.3]); let checkpoint = manager .create_checkpoint( CheckpointType::BeforeOptimization, state, "Initial model state before applying quantization".to_string(), ) .await; assert!(checkpoint.is_ok()); let checkpoint = checkpoint.unwrap(); assert_eq!( checkpoint.checkpoint_type(), CheckpointType::BeforeOptimization ); assert!(checkpoint.id().len() > 0); assert!(checkpoint.timestamp() > 0); } #[tokio::test] async fn test_save_and_load_checkpoint() { let runtime = Arc::new(Runtime::new().expect("Failed to create runtime")); let mut manager = RollbackManager::new(runtime.clone()).unwrap(); let mut original_state = HashMap::new(); original_state.insert("layer1".to_string(), vec![1.0, 2.0, 3.0, 4.0]); original_state.insert("layer2".to_string(), vec![5.0, 6.0, 7.0, 8.0]); let checkpoint = manager .create_checkpoint( CheckpointType::BeforeOptimization, original_state.clone(), "Test checkpoint".to_string(), ) .await .unwrap(); let save_result = manager.save_checkpoint(&checkpoint).await; assert!(save_result.is_ok()); let loaded_checkpoint = manager.load_checkpoint(checkpoint.id()).await; assert!(loaded_checkpoint.is_ok()); let loaded = loaded_checkpoint.unwrap(); assert_eq!(loaded.id(), checkpoint.id()); assert_eq!(loaded.state(), checkpoint.state()); } #[tokio::test] async fn test_rollback_to_checkpoint() { let runtime = Arc::new(Runtime::new().expect("Failed to create runtime")); let mut manager = RollbackManager::new(runtime.clone()).unwrap(); // Create original state let mut original_state = HashMap::new(); original_state.insert("accuracy".to_string(), vec![0.95]); original_state.insert("loss".to_string(), vec![0.1]); let checkpoint = manager .create_checkpoint( CheckpointType::BeforeOptimization, original_state, "Before optimization".to_string(), ) .await .unwrap(); manager.save_checkpoint(&checkpoint).await.unwrap(); // Simulate applying an optimization that degrades performance let mut degraded_state = HashMap::new(); degraded_state.insert("accuracy".to_string(), vec![0.70]); // Worse accuracy degraded_state.insert("loss".to_string(), vec![0.5]); // Higher loss manager.update_current_state(degraded_state).await.unwrap(); // Rollback to the checkpoint let rollback_result = manager.rollback_to_checkpoint(checkpoint.id()).await; assert!(rollback_result.is_ok()); let current_state = manager.get_current_state().await.unwrap(); assert_eq!(current_state["accuracy"], vec![0.95]); assert_eq!(current_state["loss"], vec![0.1]); } #[tokio::test] async fn test_automatic_rollback_on_degradation() { let runtime = Arc::new(Runtime::new().expect("Failed to create runtime")); let mut manager = RollbackManager::new(runtime.clone()).unwrap(); // Set up automatic rollback threshold manager .set_rollback_threshold("accuracy", 0.90) .await .unwrap(); let mut good_state = HashMap::new(); good_state.insert("accuracy".to_string(), vec![0.95]); let checkpoint = manager .create_checkpoint( CheckpointType::Automatic, good_state, "Good performance state".to_string(), ) .await .unwrap(); manager.save_checkpoint(&checkpoint).await.unwrap(); // Apply changes that degrade performance below threshold let mut bad_state = HashMap::new(); bad_state.insert("accuracy".to_string(), vec![0.85]); // Below threshold let should_rollback = manager .should_trigger_automatic_rollback(&bad_state) .await .unwrap(); assert!( should_rollback, "Should trigger automatic rollback when performance degrades" ); if should_rollback { let rollback_result = manager.automatic_rollback().await; assert!(rollback_result.is_ok()); } } #[tokio::test] async fn test_checkpoint_cleanup() { let runtime = Arc::new(Runtime::new().expect("Failed to create runtime")); let mut manager = RollbackManager::new(runtime.clone()).unwrap(); // Create multiple checkpoints let mut checkpoints = Vec::new(); for i in 0..5 { let mut state = HashMap::new(); state.insert(format!("metric_{}", i), vec![i as f32]); let checkpoint = manager .create_checkpoint(CheckpointType::Manual, state, format!("Checkpoint {}", i)) .await .unwrap(); manager.save_checkpoint(&checkpoint).await.unwrap(); checkpoints.push(checkpoint); } // Cleanup old checkpoints (keep only 3 most recent) let cleanup_result = manager.cleanup_old_checkpoints(3).await; assert!(cleanup_result.is_ok()); let remaining_checkpoints = manager.list_checkpoints().await.unwrap(); assert!( remaining_checkpoints.len() <= 3, "Should keep only 3 most recent checkpoints" ); } #[tokio::test] async fn test_proposal_application_with_rollback() { let runtime = Arc::new(Runtime::new().expect("Failed to create runtime")); let mut manager = RollbackManager::new(runtime.clone()).unwrap(); let proposal = Proposal::new( ProposalType::Quantization, "Apply aggressive quantization".to_string(), 2.0, ); // Create checkpoint before applying proposal let mut initial_state = HashMap::new(); initial_state.insert("model_size".to_string(), vec![1000.0]); initial_state.insert("accuracy".to_string(), vec![0.95]); let checkpoint = manager .create_checkpoint_for_proposal(&proposal, initial_state) .await; assert!(checkpoint.is_ok()); let checkpoint = checkpoint.unwrap(); manager.save_checkpoint(&checkpoint).await.unwrap(); // Simulate proposal application result let mut new_state = HashMap::new(); new_state.insert("model_size".to_string(), vec![250.0]); // Smaller model new_state.insert("accuracy".to_string(), vec![0.85]); // Lower accuracy let validation_result = manager .validate_proposal_result(&proposal, &new_state) .await; assert!(validation_result.is_ok()); }