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>
471 lines
12 KiB
Rust
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)));
|
|
}
|