#![allow( unused_variables, unused_imports, clippy::unwrap_used, clippy::expect_used, clippy::indexing_slicing, clippy::while_let_loop, clippy::collapsible_match )] //! gRPC Streaming Progress Tests //! //! This module tests real-time streaming progress updates for ML training jobs. //! Tests cover single job streaming, batch multiplexing, delta updates, client disconnects, //! and error handling. use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; use chrono::Utc; use tokio::sync::{broadcast, mpsc, RwLock}; use tokio::time::timeout; use uuid::Uuid; use ml_training_service::orchestrator::{ JobStatus, ResourceUsage, TrainingJob, TrainingStatusUpdate, }; // ============================================================================ // Test Helpers // ============================================================================ /// Mock orchestrator for testing streaming struct MockOrchestrator { jobs: Arc>>, status_broadcasters: Arc>>>, } impl MockOrchestrator { fn new() -> Self { Self { jobs: Arc::new(RwLock::new(HashMap::new())), status_broadcasters: Arc::new(RwLock::new(HashMap::new())), } } async fn add_job(&self, job: TrainingJob) -> broadcast::Sender { let job_id = job.id; let (tx, _rx) = broadcast::channel(100); self.jobs.write().await.insert(job_id, job); self.status_broadcasters.write().await.insert(job_id, tx.clone()); tx } async fn get_job(&self, job_id: Uuid) -> Option { self.jobs.read().await.get(&job_id).cloned() } async fn subscribe( &self, job_id: Uuid, ) -> Option> { self.status_broadcasters .read() .await .get(&job_id) .map(|tx| tx.subscribe()) } } /// Create a test training job fn create_test_job(model_type: &str, asset: &str) -> TrainingJob { let config = ml::training_pipeline::ProductionTrainingConfig::default(); let mut tags = HashMap::new(); tags.insert("asset".to_string(), asset.to_string()); TrainingJob::new( model_type.to_string(), config, format!("{} training for {}", model_type, asset), tags, ) } /// Create a test status update fn create_status_update( job: &TrainingJob, epoch: u32, total_epochs: u32, progress: f32, ) -> TrainingStatusUpdate { TrainingStatusUpdate { job_id: job.id, status: JobStatus::Running, progress_percentage: progress, current_epoch: epoch, total_epochs, metrics: HashMap::new(), message: format!("Epoch {}/{}", epoch, total_epochs), timestamp: Utc::now(), financial_metrics: None, resource_usage: ResourceUsage { cpu_usage_percent: 50.0, memory_usage_gb: 2.5, gpu_usage_percent: Some(75.0), gpu_memory_usage_gb: Some(1.2), active_workers: 1, }, } } // ============================================================================ // Test Cases // ============================================================================ #[tokio::test] async fn test_stream_single_job_progress() { // ARRANGE: Create mock orchestrator with one job let orchestrator = MockOrchestrator::new(); let job = create_test_job("DQN", "ES.FUT"); let job_id = job.id; let broadcaster = orchestrator.add_job(job.clone()).await; // Create test streaming channel let (tx, mut rx) = mpsc::channel(100); // ACT: Spawn task that simulates streaming logic let job_clone = job.clone(); let orchestrator_clone = orchestrator.subscribe(job_id).await.unwrap(); tokio::spawn(async move { let mut receiver = orchestrator_clone; loop { match receiver.recv().await { Ok(update) => { if tx.send(update).await.is_err() { break; } } Err(_) => break, } } }); // Send progress updates let update1 = create_status_update(&job_clone, 1, 10, 10.0); let update2 = create_status_update(&job_clone, 5, 10, 50.0); let update3 = create_status_update(&job_clone, 10, 10, 100.0); broadcaster.send(update1.clone()).unwrap(); broadcaster.send(update2.clone()).unwrap(); broadcaster.send(update3.clone()).unwrap(); // ASSERT: Receive all updates in order let received1 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(received1.current_epoch, 1); assert_eq!(received1.progress_percentage, 10.0); let received2 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(received2.current_epoch, 5); assert_eq!(received2.progress_percentage, 50.0); let received3 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(received3.current_epoch, 10); assert_eq!(received3.progress_percentage, 100.0); } #[tokio::test] async fn test_stream_batch_progress_weighted() { // ARRANGE: Create batch with 3 jobs let orchestrator = Arc::new(MockOrchestrator::new()); let job1 = create_test_job("DQN", "ES.FUT"); let job2 = create_test_job("PPO", "NQ.FUT"); let job3 = create_test_job("MAMBA_2", "6E.FUT"); let broadcaster1 = orchestrator.add_job(job1.clone()).await; let broadcaster2 = orchestrator.add_job(job2.clone()).await; let broadcaster3 = orchestrator.add_job(job3.clone()).await; // Collect job IDs let jobs = vec![job1.clone(), job2.clone(), job3.clone()]; let job_ids = vec![job1.id, job2.id, job3.id]; // Create aggregation channel let (tx, mut rx) = mpsc::channel(100); // ACT: Spawn task that aggregates batch progress let orchestrator_clone = orchestrator.clone(); tokio::spawn(async move { let mut progress_map: HashMap = HashMap::new(); for job in &jobs { progress_map.insert(job.id, 0.0); } let (update_tx, mut update_rx) = mpsc::channel(100); // Spawn receiver for each job for job_id in job_ids { let update_tx_clone = update_tx.clone(); if let Some(mut receiver) = orchestrator_clone.subscribe(job_id).await { tokio::spawn(async move { while let Ok(update) = receiver.recv().await { let _ = update_tx_clone.send(update).await; } }); } } drop(update_tx); // Aggregate updates from all jobs while let Some(update) = update_rx.recv().await { progress_map.insert(update.job_id, update.progress_percentage); let avg = progress_map.values().sum::() / jobs.len() as f32; if tx.send(avg).await.is_err() { break; } } }); // Wait for receivers to be ready tokio::time::sleep(Duration::from_millis(50)).await; // Send updates to different jobs broadcaster1.send(create_status_update(&job1, 3, 10, 30.0)).unwrap(); let progress1 = timeout(Duration::from_millis(200), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert!((progress1 - 10.0).abs() < 0.01); // (30 + 0 + 0) / 3 broadcaster2.send(create_status_update(&job2, 6, 10, 60.0)).unwrap(); let progress2 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert!((progress2 - 30.0).abs() < 0.01); // (30 + 60 + 0) / 3 broadcaster3.send(create_status_update(&job3, 9, 10, 90.0)).unwrap(); let progress3 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert!((progress3 - 60.0).abs() < 0.01); // (30 + 60 + 90) / 3 } #[tokio::test] async fn test_stream_delta_only() { // ARRANGE: Create orchestrator with one job let orchestrator = MockOrchestrator::new(); let job = create_test_job("DQN", "ES.FUT"); let job_id = job.id; let broadcaster = orchestrator.add_job(job.clone()).await; let (tx, mut rx) = mpsc::channel(100); // ACT: Spawn delta-filtering task let job_clone = job.clone(); let orchestrator_clone = orchestrator.subscribe(job_id).await.unwrap(); tokio::spawn(async move { let mut receiver = orchestrator_clone; let mut last_progress = 0.0f32; loop { match receiver.recv().await { Ok(update) => { // Only send if progress changed if (update.progress_percentage - last_progress).abs() > 0.001 { last_progress = update.progress_percentage; if tx.send(update).await.is_err() { break; } } } Err(_) => break, } } }); // Send same update twice let update1 = create_status_update(&job_clone, 5, 10, 50.0); broadcaster.send(update1.clone()).unwrap(); broadcaster.send(update1.clone()).unwrap(); // Duplicate // Send different update let update2 = create_status_update(&job_clone, 6, 10, 60.0); broadcaster.send(update2.clone()).unwrap(); // ASSERT: Receive only 2 updates (duplicate filtered) let received1 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(received1.progress_percentage, 50.0); let received2 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(received2.progress_percentage, 60.0); // No third update should arrive let result = timeout(Duration::from_millis(200), rx.recv()).await; assert!(result.is_err(), "Should not receive duplicate update"); } #[tokio::test] async fn test_stream_handles_completion() { // ARRANGE: Create orchestrator with one job let orchestrator = MockOrchestrator::new(); let mut job = create_test_job("DQN", "ES.FUT"); let job_id = job.id; let broadcaster = orchestrator.add_job(job.clone()).await; let (tx, mut rx) = mpsc::channel(100); // ACT: Spawn task that closes on completion let job_clone = job.clone(); let orchestrator_clone = orchestrator.subscribe(job_id).await.unwrap(); tokio::spawn(async move { let mut receiver = orchestrator_clone; loop { match receiver.recv().await { Ok(update) => { let should_close = update.status == JobStatus::Completed; if tx.send(update).await.is_err() { break; } if should_close { break; // Close stream on completion } } Err(_) => break, } } }); // Send running update broadcaster.send(create_status_update(&job_clone, 5, 10, 50.0)).unwrap(); // Send completion update job.status = JobStatus::Completed; let mut completion_update = create_status_update(&job_clone, 10, 10, 100.0); completion_update.status = JobStatus::Completed; broadcaster.send(completion_update).expect("INVARIANT: Channel should not be closed"); // ASSERT: Receive both updates, then stream closes let received1 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(received1.progress_percentage, 50.0); let received2 = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(received2.status, JobStatus::Completed); // Stream should close let result = timeout(Duration::from_millis(100), rx.recv()).await; assert!(result.is_ok()); assert!(result.unwrap().is_none(), "Stream should close"); } #[tokio::test] async fn test_stream_handles_client_disconnect() { // ARRANGE: Create orchestrator with one job let orchestrator = MockOrchestrator::new(); let job = create_test_job("DQN", "ES.FUT"); let job_id = job.id; let broadcaster = orchestrator.add_job(job.clone()).await; let (tx, rx) = mpsc::channel(100); let (cleanup_tx, mut cleanup_rx) = mpsc::channel(1); // ACT: Spawn task that detects disconnect let job_clone = job.clone(); let orchestrator_clone = orchestrator.subscribe(job_id).await.unwrap(); tokio::spawn(async move { let mut receiver = orchestrator_clone; loop { match receiver.recv().await { Ok(update) => { if tx.send(update).await.is_err() { // Client disconnected - send cleanup signal let _ = cleanup_tx.send(true).await; break; } } Err(_) => break, } } }); // Immediately drop receiver (simulate client disconnect) drop(rx); // Send update (should trigger disconnect detection) broadcaster.send(create_status_update(&job_clone, 5, 10, 50.0)).unwrap(); // ASSERT: Cleanup signal received let cleanup = timeout(Duration::from_millis(200), cleanup_rx.recv()) .await .expect("Timeout waiting for cleanup") .expect("Cleanup channel closed"); assert!(cleanup, "Cleanup should be triggered on disconnect"); } #[tokio::test] async fn test_stream_nonexistent_job() { // ARRANGE: Create empty orchestrator let orchestrator = MockOrchestrator::new(); let nonexistent_id = Uuid::new_v4(); // ACT: Try to subscribe to nonexistent job let receiver = orchestrator.subscribe(nonexistent_id).await; // ASSERT: Should return None assert!(receiver.is_none(), "Should return None for nonexistent job"); } #[tokio::test] async fn test_stream_multiplexed_jobs() { // ARRANGE: Create batch with 16 jobs (max multiplexing) let orchestrator = Arc::new(MockOrchestrator::new()); let mut jobs = Vec::new(); let mut broadcasters = Vec::new(); for i in 0..16 { let job = create_test_job("DQN", &format!("ASSET_{}", i)); let broadcaster = orchestrator.add_job(job.clone()).await; jobs.push(job); broadcasters.push(broadcaster); } // Create aggregation channel let (tx, mut rx) = mpsc::channel(100); // ACT: Spawn task that handles all 16 jobs let job_ids: Vec = jobs.iter().map(|j| j.id).collect(); let jobs_clone = jobs.clone(); let orchestrator_clone = orchestrator.clone(); tokio::spawn(async move { let mut progress_map: HashMap = HashMap::new(); for job in &jobs_clone { progress_map.insert(job.id, 0.0); } let (update_tx, mut update_rx) = mpsc::channel(100); // Spawn receiver for each job for job_id in job_ids { let update_tx_clone = update_tx.clone(); if let Some(mut receiver) = orchestrator_clone.subscribe(job_id).await { tokio::spawn(async move { while let Ok(update) = receiver.recv().await { let _ = update_tx_clone.send(update).await; } }); } } drop(update_tx); // Aggregate updates from all jobs while let Some(update) = update_rx.recv().await { progress_map.insert(update.job_id, update.progress_percentage); let avg = progress_map.values().sum::() / jobs_clone.len() as f32; if tx.send(avg).await.is_err() { break; } } }); // Wait for all receivers to be ready tokio::time::sleep(Duration::from_millis(50)).await; // Send updates to 4 different jobs broadcasters[0].send(create_status_update(&jobs[0], 1, 10, 10.0)).unwrap(); broadcasters[5].send(create_status_update(&jobs[5], 2, 10, 20.0)).unwrap(); broadcasters[10].send(create_status_update(&jobs[10], 3, 10, 30.0)).unwrap(); broadcasters[15].send(create_status_update(&jobs[15], 4, 10, 40.0)).unwrap(); // Wait for aggregation tokio::time::sleep(Duration::from_millis(150)).await; // ASSERT: Should receive aggregated progress let progress = timeout(Duration::from_millis(200), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); // (10 + 0 + 0 + ... + 0) / 16 = 0.625, then updates... // Final: (10 + 20 + 30 + 40 + 0*12) / 16 = 6.25 assert!(progress > 0.0, "Should receive aggregated progress"); } #[tokio::test] async fn test_stream_initial_state_transmission() { // ARRANGE: Create job already at 50% progress let orchestrator = MockOrchestrator::new(); let mut job = create_test_job("DQN", "ES.FUT"); job.progress_percentage = 50.0; job.current_epoch = 5; job.total_epochs = 10; job.status = JobStatus::Running; let job_id = job.id; let broadcaster = orchestrator.add_job(job.clone()).await; let (tx, mut rx) = mpsc::channel(100); // ACT: Spawn task that sends initial state let initial_state = orchestrator.get_job(job_id).await.unwrap(); let orchestrator_clone = orchestrator.subscribe(job_id).await.unwrap(); tokio::spawn(async move { // Send initial state immediately let initial_update = TrainingStatusUpdate { job_id: initial_state.id, status: initial_state.status.clone(), progress_percentage: initial_state.progress_percentage, current_epoch: initial_state.current_epoch, total_epochs: initial_state.total_epochs, metrics: HashMap::new(), message: "Initial state".to_string(), timestamp: Utc::now(), financial_metrics: None, resource_usage: ResourceUsage { cpu_usage_percent: 0.0, memory_usage_gb: 0.0, gpu_usage_percent: None, gpu_memory_usage_gb: None, active_workers: 0, }, }; let _ = tx.send(initial_update).await; // Then stream live updates let mut receiver = orchestrator_clone; loop { match receiver.recv().await { Ok(update) => { if tx.send(update).await.is_err() { break; } } Err(_) => break, } } }); // ASSERT: First update is initial state let initial = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(initial.progress_percentage, 50.0); assert_eq!(initial.current_epoch, 5); assert_eq!(initial.message, "Initial state"); // Now send live update broadcaster.send(create_status_update(&job, 6, 10, 60.0)).unwrap(); let live_update = timeout(Duration::from_millis(100), rx.recv()) .await .expect("Timeout") .expect("Channel closed"); assert_eq!(live_update.progress_percentage, 60.0); assert_eq!(live_update.current_epoch, 6); } #[tokio::test] async fn test_stream_broadcast_lag_handling() { // ARRANGE: Create orchestrator with small capacity broadcast channel let orchestrator = MockOrchestrator::new(); let job = create_test_job("DQN", "ES.FUT"); let job_id = job.id; // Create small capacity broadcaster let (tx, _rx) = broadcast::channel::(2); // Small capacity orchestrator.status_broadcasters.write().await.insert(job_id, tx.clone()); orchestrator.jobs.write().await.insert(job_id, job.clone()); let (result_tx, mut result_rx) = mpsc::channel(100); // ACT: Spawn receiver task let mut receiver = tx.subscribe(); tokio::spawn(async move { loop { match receiver.recv().await { Ok(update) => { let _ = result_tx.send(Ok(update)).await; } Err(broadcast::error::RecvError::Lagged(n)) => { // Handle lag gracefully let _ = result_tx.send(Err(n)).await; } Err(_) => break, } } }); // Send 5 messages rapidly to overflow capacity for i in 0..5 { tx.send(create_status_update(&job, i, 10, (i * 10) as f32)).unwrap(); } // ASSERT: Should receive lag error let mut received_lag = false; for _ in 0..6 { if let Ok(result) = timeout(Duration::from_millis(100), result_rx.recv()).await { if let Some(Err(_lag_count)) = result { received_lag = true; break; } } } assert!(received_lag, "Should detect broadcast lag"); }