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>
453 lines
12 KiB
Rust
453 lines
12 KiB
Rust
//! End-to-End Full User Workflow Tests
|
||
//!
|
||
//! Tests complete user scenarios from start to finish
|
||
|
||
use fxt::proto::ml_training::{
|
||
ml_training_service_server::MlTrainingServiceServer, TrainingStatus,
|
||
ml_training_service_client::MlTrainingServiceClient, SubscribeToTrainingStatusRequest,
|
||
StopTrainingRequest, GetTrainingJobDetailsRequest,
|
||
};
|
||
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
|
||
}
|
||
|
||
/// SCENARIO 1: Start → Watch → Complete
|
||
#[tokio::test]
|
||
async fn test_workflow_start_watch_complete() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Start: Create a job
|
||
let job = test_fixtures::create_test_job(
|
||
"workflow_1",
|
||
"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);
|
||
|
||
// Watch: Subscribe to progress
|
||
let request = SubscribeToTrainingStatusRequest {
|
||
job_id: "workflow_1".to_string(),
|
||
};
|
||
|
||
let mut stream = client
|
||
.subscribe_to_training_status(request)
|
||
.await
|
||
.unwrap()
|
||
.into_inner();
|
||
|
||
// Complete: Wait for 100%
|
||
let mut completed = false;
|
||
while let Some(update) = stream.next().await {
|
||
let update = update.unwrap();
|
||
if update.progress_percentage >= 100.0 {
|
||
completed = true;
|
||
break;
|
||
}
|
||
}
|
||
|
||
assert!(completed, "Job should complete");
|
||
|
||
// Verify final status
|
||
let job = state.get_job("workflow_1").unwrap();
|
||
assert_eq!(job.status, TrainingStatus::Completed as i32);
|
||
}
|
||
|
||
/// SCENARIO 2: Start → Stop → Restart
|
||
#[tokio::test]
|
||
async fn test_workflow_start_stop_restart() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Start: Create first job
|
||
let job1 = test_fixtures::create_test_job(
|
||
"workflow_2_first",
|
||
"DQN",
|
||
TrainingStatus::Running,
|
||
);
|
||
state.add_job(job1);
|
||
|
||
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);
|
||
|
||
// Stop: Stop the first job
|
||
let stop_request = StopTrainingRequest {
|
||
job_id: "workflow_2_first".to_string(),
|
||
reason: "User requested restart".to_string(),
|
||
};
|
||
|
||
let stop_response = client.stop_training(stop_request).await;
|
||
assert!(stop_response.is_ok());
|
||
|
||
// Verify stopped
|
||
let job1_stopped = state.get_job("workflow_2_first").unwrap();
|
||
assert_eq!(job1_stopped.status, TrainingStatus::Stopped as i32);
|
||
|
||
// Restart: Create new job
|
||
let job2 = test_fixtures::create_test_job(
|
||
"workflow_2_restart",
|
||
"DQN",
|
||
TrainingStatus::Pending,
|
||
);
|
||
state.add_job(job2);
|
||
|
||
let job2_retrieved = state.get_job("workflow_2_restart").unwrap();
|
||
assert_eq!(job2_retrieved.status, TrainingStatus::Pending as i32);
|
||
}
|
||
|
||
/// SCENARIO 3: Start → Error → Retry
|
||
#[tokio::test]
|
||
async fn test_workflow_start_error_retry() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Start: Create job that will fail
|
||
let job1 = test_fixtures::create_test_job(
|
||
"workflow_3_fail",
|
||
"PPO",
|
||
TrainingStatus::Failed,
|
||
);
|
||
state.add_job(job1);
|
||
|
||
let _addr = start_mock_server(state.clone()).await;
|
||
|
||
// Verify failed
|
||
let job1_failed = state.get_job("workflow_3_fail").unwrap();
|
||
assert_eq!(job1_failed.status, TrainingStatus::Failed as i32);
|
||
|
||
// Retry: Create new job with same parameters
|
||
let job2 = test_fixtures::create_test_job(
|
||
"workflow_3_retry",
|
||
"PPO",
|
||
TrainingStatus::Pending,
|
||
);
|
||
state.add_job(job2);
|
||
|
||
let job2_retrieved = state.get_job("workflow_3_retry").unwrap();
|
||
assert_eq!(job2_retrieved.status, TrainingStatus::Pending as i32);
|
||
}
|
||
|
||
/// SCENARIO 4: Multiple batches concurrently
|
||
#[tokio::test]
|
||
async fn test_workflow_multiple_batches() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Create two batches
|
||
for batch in 0..2 {
|
||
for model in 0..4 {
|
||
let job = test_fixtures::create_test_job(
|
||
&format!("batch_{}_model_{}", batch, model),
|
||
"TFT",
|
||
TrainingStatus::Running,
|
||
);
|
||
state.add_job(job);
|
||
}
|
||
}
|
||
|
||
let _addr = start_mock_server(state.clone()).await;
|
||
|
||
// Verify all jobs created
|
||
let jobs = state.list_jobs();
|
||
assert_eq!(jobs.len(), 8, "Should have 8 jobs (2 batches × 4 models)");
|
||
|
||
// Verify batch organization
|
||
let batch_0: Vec<_> = jobs.iter().filter(|j| j.job_id.contains("batch_0")).collect();
|
||
let batch_1: Vec<_> = jobs.iter().filter(|j| j.job_id.contains("batch_1")).collect();
|
||
|
||
assert_eq!(batch_0.len(), 4);
|
||
assert_eq!(batch_1.len(), 4);
|
||
}
|
||
|
||
/// SCENARIO 5: Long-running batch (resume watch)
|
||
#[tokio::test]
|
||
async fn test_workflow_long_running_resume() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Create a long-running job
|
||
let job = test_fixtures::create_test_job(
|
||
"workflow_5_long",
|
||
"MAMBA_2",
|
||
TrainingStatus::Running,
|
||
);
|
||
state.add_job(job);
|
||
|
||
let addr = start_mock_server(state.clone()).await;
|
||
|
||
let channel1 = Channel::from_shared(format!("http://{}", addr))
|
||
.unwrap()
|
||
.connect()
|
||
.await
|
||
.unwrap();
|
||
|
||
let mut client1 = MlTrainingServiceClient::new(channel1);
|
||
|
||
// Start watching
|
||
let request1 = SubscribeToTrainingStatusRequest {
|
||
job_id: "workflow_5_long".to_string(),
|
||
};
|
||
|
||
let mut stream1 = client1
|
||
.subscribe_to_training_status(request1)
|
||
.await
|
||
.unwrap()
|
||
.into_inner();
|
||
|
||
// Read a few updates
|
||
let mut update_count = 0;
|
||
while let Some(update) = stream1.next().await {
|
||
let _update = update.unwrap();
|
||
update_count += 1;
|
||
if update_count >= 3 {
|
||
break;
|
||
}
|
||
}
|
||
|
||
// Disconnect
|
||
drop(stream1);
|
||
|
||
// Resume watching with new connection
|
||
let channel2 = Channel::from_shared(format!("http://{}", addr))
|
||
.unwrap()
|
||
.connect()
|
||
.await
|
||
.unwrap();
|
||
|
||
let mut client2 = MlTrainingServiceClient::new(channel2);
|
||
|
||
let request2 = SubscribeToTrainingStatusRequest {
|
||
job_id: "workflow_5_long".to_string(),
|
||
};
|
||
|
||
let mut stream2 = client2
|
||
.subscribe_to_training_status(request2)
|
||
.await
|
||
.unwrap()
|
||
.into_inner();
|
||
|
||
// Should be able to resume
|
||
let mut resumed_updates = 0;
|
||
while let Some(update) = stream2.next().await {
|
||
let update = update.unwrap();
|
||
resumed_updates += 1;
|
||
if update.progress_percentage >= 100.0 {
|
||
break;
|
||
}
|
||
}
|
||
|
||
assert!(resumed_updates > 0, "Should receive updates after resume");
|
||
}
|
||
|
||
/// SCENARIO 6: Query status during training
|
||
#[tokio::test]
|
||
async fn test_workflow_status_during_training() {
|
||
let state = MockTrainingState::default();
|
||
|
||
let job = test_fixtures::create_test_job(
|
||
"workflow_6_status",
|
||
"TFT",
|
||
TrainingStatus::Running,
|
||
);
|
||
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);
|
||
|
||
// Query status while running
|
||
let status_request = GetTrainingJobDetailsRequest {
|
||
job_id: "workflow_6_status".to_string(),
|
||
};
|
||
|
||
let status_response = client.get_training_job_details(status_request).await;
|
||
assert!(status_response.is_ok());
|
||
|
||
// Verify job is running
|
||
let job = state.get_job("workflow_6_status").unwrap();
|
||
assert_eq!(job.status, TrainingStatus::Running as i32);
|
||
}
|
||
|
||
/// SCENARIO 7: List jobs during batch execution
|
||
#[tokio::test]
|
||
async fn test_workflow_list_during_batch() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Create batch jobs in various states
|
||
let states = vec![
|
||
TrainingStatus::Pending,
|
||
TrainingStatus::Running,
|
||
TrainingStatus::Running,
|
||
TrainingStatus::Completed,
|
||
];
|
||
|
||
for (i, status) in states.iter().enumerate() {
|
||
let job = test_fixtures::create_test_job(
|
||
&format!("workflow_7_batch_{}", i),
|
||
"DQN",
|
||
*status,
|
||
);
|
||
state.add_job(job);
|
||
}
|
||
|
||
let _addr = start_mock_server(state.clone()).await;
|
||
|
||
// List all jobs
|
||
let all_jobs = state.list_jobs();
|
||
assert_eq!(all_jobs.len(), 4);
|
||
|
||
// Filter running jobs
|
||
let running = test_fixtures::filter_by_status(&all_jobs, TrainingStatus::Running);
|
||
assert_eq!(running.len(), 2);
|
||
|
||
// Filter completed jobs
|
||
let completed = test_fixtures::filter_by_status(&all_jobs, TrainingStatus::Completed);
|
||
assert_eq!(completed.len(), 1);
|
||
}
|
||
|
||
/// SCENARIO 8: Error recovery workflow
|
||
#[tokio::test]
|
||
async fn test_workflow_error_recovery() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Start job
|
||
let job1 = test_fixtures::create_test_job(
|
||
"workflow_8_original",
|
||
"PPO",
|
||
TrainingStatus::Running,
|
||
);
|
||
state.add_job(job1);
|
||
|
||
let _addr = start_mock_server(state.clone()).await;
|
||
|
||
// Simulate error
|
||
state.update_job_status("workflow_8_original", TrainingStatus::Failed);
|
||
|
||
// Verify failed
|
||
let job1_failed = state.get_job("workflow_8_original").unwrap();
|
||
assert_eq!(job1_failed.status, TrainingStatus::Failed as i32);
|
||
|
||
// Recovery: Start new job with adjusted parameters
|
||
let job2 = test_fixtures::create_test_job(
|
||
"workflow_8_recovery",
|
||
"PPO",
|
||
TrainingStatus::Pending,
|
||
);
|
||
state.add_job(job2);
|
||
|
||
// Verify recovery job
|
||
let job2_recovered = state.get_job("workflow_8_recovery").unwrap();
|
||
assert_eq!(job2_recovered.status, TrainingStatus::Pending as i32);
|
||
}
|
||
|
||
/// SCENARIO 9: Complete batch monitoring
|
||
#[tokio::test]
|
||
async fn test_workflow_batch_monitoring() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Create batch
|
||
for i in 0..4 {
|
||
let job = test_fixtures::create_test_job(
|
||
&format!("workflow_9_batch_{}", i),
|
||
"TFT",
|
||
TrainingStatus::Pending,
|
||
);
|
||
state.add_job(job);
|
||
}
|
||
|
||
let _addr = start_mock_server(state.clone()).await;
|
||
|
||
// Monitor: Simulate progress updates
|
||
for i in 0..4 {
|
||
state.update_job_status(&format!("workflow_9_batch_{}", i), TrainingStatus::Running);
|
||
}
|
||
|
||
// Check all running
|
||
let jobs = state.list_jobs();
|
||
for job in &jobs {
|
||
assert_eq!(job.status, TrainingStatus::Running as i32);
|
||
}
|
||
|
||
// Complete batch
|
||
for i in 0..4 {
|
||
state.update_job_status(&format!("workflow_9_batch_{}", i), TrainingStatus::Completed);
|
||
}
|
||
|
||
// Verify all completed
|
||
let jobs = state.list_jobs();
|
||
for job in &jobs {
|
||
assert_eq!(job.status, TrainingStatus::Completed as i32);
|
||
}
|
||
}
|
||
|
||
/// SCENARIO 10: Cleanup after workflow
|
||
#[tokio::test]
|
||
async fn test_workflow_cleanup() {
|
||
let state = MockTrainingState::default();
|
||
|
||
// Create test jobs
|
||
for i in 0..5 {
|
||
let job = test_fixtures::create_test_job(
|
||
&format!("workflow_10_cleanup_{}", i),
|
||
"MAMBA_2",
|
||
TrainingStatus::Completed,
|
||
);
|
||
state.add_job(job);
|
||
}
|
||
|
||
let _addr = start_mock_server(state.clone()).await;
|
||
|
||
// Verify jobs exist
|
||
let jobs = state.list_jobs();
|
||
assert_eq!(jobs.len(), 5);
|
||
|
||
// Cleanup: Clear state
|
||
state.clear();
|
||
|
||
// Verify cleanup
|
||
let jobs_after = state.list_jobs();
|
||
assert_eq!(jobs_after.len(), 0, "Should have no jobs after cleanup");
|
||
}
|