//! End-to-End Tests for `tli train watch` Command //! //! Tests real-time training progress monitoring use fxt::proto::ml_training::{ ml_training_service_server::MlTrainingServiceServer, TrainingStatus, ml_training_service_client::MlTrainingServiceClient, SubscribeToTrainingStatusRequest, }; use tonic::transport::{Server, Channel}; use std::net::SocketAddr; use tokio::time::Duration; use futures_util::StreamExt; use super::mock_ml_training_service::{MockMlTrainingService, MockTrainingState}; use super::test_fixtures; /// Helper to start mock gRPC server async fn start_mock_server(state: MockTrainingState) -> SocketAddr { let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); let service = MockMlTrainingService::new(state); let listener = tokio::net::TcpListener::bind(addr).await.unwrap(); let local_addr = listener.local_addr().unwrap(); tokio::spawn(async move { Server::builder() .add_service(MlTrainingServiceServer::new(service)) .serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener)) .await .unwrap(); }); tokio::time::sleep(Duration::from_millis(100)).await; local_addr } /// TEST 1: Watch single job progress (0% → 100%) #[tokio::test] async fn test_watch_single_job_progress() { let state = MockTrainingState::default(); // Create a job let job = test_fixtures::create_test_job( "watch_test_1", "TFT", TrainingStatus::Pending, ); state.add_job(job); let addr = start_mock_server(state.clone()).await; // Connect to mock server let channel = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); // Subscribe to training status let request = SubscribeToTrainingStatusRequest { job_id: "watch_test_1".to_string(), }; let mut stream = client .subscribe_to_training_status(request) .await .unwrap() .into_inner(); // Collect all progress updates let mut updates = vec![]; while let Some(update) = stream.next().await { let update = update.unwrap(); let progress = update.progress_percentage; updates.push(update); // Break after 100% to avoid infinite stream if progress >= 100.0 { break; } } assert!(!updates.is_empty(), "Should receive progress updates"); assert_eq!(updates.last().unwrap().progress_percentage, 100.0); assert_eq!(updates.last().unwrap().status, TrainingStatus::Completed as i32); } /// TEST 2: Watch batch progress (weighted calculation) #[tokio::test] async fn test_watch_batch_progress() { let state = MockTrainingState::default(); // Create multiple jobs (batch) for i in 0..4 { let job = test_fixtures::create_test_job( &format!("batch_job_{}", i), "TFT", TrainingStatus::Pending, ); state.add_job(job); } let _addr = start_mock_server(state.clone()).await; // In batch mode, we would watch all jobs and aggregate progress let jobs = state.list_jobs(); assert_eq!(jobs.len(), 4, "Should have 4 batch jobs"); } /// TEST 3: Multiple concurrent watchers (16 streams) #[tokio::test] async fn test_concurrent_watchers() { let state = MockTrainingState::default(); // Create a job let job = test_fixtures::create_test_job( "concurrent_watch", "DQN", TrainingStatus::Pending, ); state.add_job(job); let addr = start_mock_server(state.clone()).await; // Spawn 16 concurrent watchers let mut handles = vec![]; for i in 0..16 { let addr_clone = addr; let handle = tokio::spawn(async move { let channel = Channel::from_shared(format!("http://{}", addr_clone)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); let request = SubscribeToTrainingStatusRequest { job_id: "concurrent_watch".to_string(), }; let mut stream = client .subscribe_to_training_status(request) .await .unwrap() .into_inner(); // Read first update if let Some(update) = stream.next().await { assert!(update.is_ok()); } i }); handles.push(handle); } // Wait for all watchers for handle in handles { let result = handle.await.unwrap(); assert!(result < 16); } } /// TEST 4: Resume watch after disconnect #[tokio::test] async fn test_resume_watch_after_disconnect() { let state = MockTrainingState::default(); let job = test_fixtures::create_test_job( "resume_test", "PPO", TrainingStatus::Running, ); state.add_job(job); let addr = start_mock_server(state.clone()).await; // First connection let channel1 = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client1 = MlTrainingServiceClient::new(channel1); let request1 = SubscribeToTrainingStatusRequest { job_id: "resume_test".to_string(), }; let mut stream1 = client1 .subscribe_to_training_status(request1) .await .unwrap() .into_inner(); // Read one update then disconnect if let Some(update) = stream1.next().await { assert!(update.is_ok()); } drop(stream1); // Reconnect let channel2 = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client2 = MlTrainingServiceClient::new(channel2); let request2 = SubscribeToTrainingStatusRequest { job_id: "resume_test".to_string(), }; let mut stream2 = client2 .subscribe_to_training_status(request2) .await .unwrap() .into_inner(); // Should be able to resume if let Some(update) = stream2.next().await { assert!(update.is_ok()); } } /// TEST 5: Terminal status auto-close #[tokio::test] async fn test_terminal_status_auto_close() { let state = MockTrainingState::default(); let job = test_fixtures::create_test_job( "terminal_test", "TFT", TrainingStatus::Pending, ); state.add_job(job); let addr = start_mock_server(state.clone()).await; let channel = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); let request = SubscribeToTrainingStatusRequest { job_id: "terminal_test".to_string(), }; let mut stream = client .subscribe_to_training_status(request) .await .unwrap() .into_inner(); // Collect all updates until stream ends let mut final_status = None; while let Some(update) = stream.next().await { let update = update.unwrap(); final_status = Some(update.status); if update.progress_percentage >= 100.0 { break; } } assert_eq!(final_status, Some(TrainingStatus::Completed as i32)); } /// TEST 6: Progress percentage validation #[tokio::test] async fn test_progress_percentage_validation() { let state = MockTrainingState::default(); let job = test_fixtures::create_test_job( "progress_test", "MAMBA_2", TrainingStatus::Pending, ); state.add_job(job); let addr = start_mock_server(state.clone()).await; let channel = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); let request = SubscribeToTrainingStatusRequest { job_id: "progress_test".to_string(), }; let mut stream = client .subscribe_to_training_status(request) .await .unwrap() .into_inner(); // Validate all progress values are 0-100 let mut last_progress = 0.0; while let Some(update) = stream.next().await { let update = update.unwrap(); assert!(update.progress_percentage >= 0.0); assert!(update.progress_percentage <= 100.0); assert!(update.progress_percentage >= last_progress, "Progress should be monotonic"); last_progress = update.progress_percentage; if update.progress_percentage >= 100.0 { break; } } } /// TEST 7: Epoch counter validation #[tokio::test] async fn test_epoch_counter() { let state = MockTrainingState::default(); let job = test_fixtures::create_test_job( "epoch_test", "DQN", TrainingStatus::Pending, ); state.add_job(job); let addr = start_mock_server(state.clone()).await; let channel = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); let request = SubscribeToTrainingStatusRequest { job_id: "epoch_test".to_string(), }; let mut stream = client .subscribe_to_training_status(request) .await .unwrap() .into_inner(); let mut max_epoch = 0; while let Some(update) = stream.next().await { let update = update.unwrap(); assert!(update.current_epoch <= update.total_epochs); max_epoch = update.current_epoch; if update.progress_percentage >= 100.0 { break; } } assert!(max_epoch > 0, "Should have progressed through epochs"); } /// TEST 8: Non-existent job watch error #[tokio::test] async fn test_watch_nonexistent_job() { let state = MockTrainingState::default(); let addr = start_mock_server(state).await; let channel = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); let request = SubscribeToTrainingStatusRequest { job_id: "nonexistent_job".to_string(), }; let result = client.subscribe_to_training_status(request).await; assert!(result.is_err(), "Should error on non-existent job"); } /// TEST 9: Stream timeout handling #[tokio::test] async fn test_stream_timeout() { let state = MockTrainingState::default(); let job = test_fixtures::create_test_job( "timeout_test", "PPO", TrainingStatus::Pending, ); state.add_job(job); let addr = start_mock_server(state).await; let channel = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); let request = SubscribeToTrainingStatusRequest { job_id: "timeout_test".to_string(), }; let stream_future = client.subscribe_to_training_status(request); // Should complete within 5 seconds let result = tokio::time::timeout(Duration::from_secs(5), stream_future).await; assert!(result.is_ok(), "Stream should be established within timeout"); } /// TEST 10: Status transition validation #[tokio::test] async fn test_status_transitions() { let state = MockTrainingState::default(); let job = test_fixtures::create_test_job( "transition_test", "TFT", TrainingStatus::Pending, ); state.add_job(job); let addr = start_mock_server(state).await; let channel = Channel::from_shared(format!("http://{}", addr)) .unwrap() .connect() .await .unwrap(); let mut client = MlTrainingServiceClient::new(channel); let request = SubscribeToTrainingStatusRequest { job_id: "transition_test".to_string(), }; let mut stream = client .subscribe_to_training_status(request) .await .unwrap() .into_inner(); let mut statuses = vec![]; while let Some(update) = stream.next().await { let update = update.unwrap(); statuses.push(update.status); if update.progress_percentage >= 100.0 { break; } } // Should see RUNNING → COMPLETED transition assert!(statuses.contains(&(TrainingStatus::Running as i32))); assert_eq!(statuses.last(), Some(&(TrainingStatus::Completed as i32))); }