Files
foxhunt/bin/fxt/tests/e2e/train_watch_e2e_test.rs
jgrusewski 9c3d741a08 refactor: restructure repo — crates/, bin/, testing/ layout
Move 17 library crates into crates/, CLI binary into bin/fxt,
consolidate 10 test crates into testing/, split config crate
from deployment config files.

Root directory reduced from 38+ to ~17 directories.
All Cargo.toml paths and build.rs proto refs updated.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-25 11:56:00 +01:00

471 lines
12 KiB
Rust

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