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>
236 lines
7.0 KiB
Rust
236 lines
7.0 KiB
Rust
//! End-to-End Tests for `tli train start` Command
|
|
//!
|
|
//! Tests training job creation and validation
|
|
|
|
use fxt::proto::ml_training::{
|
|
ml_training_service_server::MlTrainingServiceServer, TrainingStatus,
|
|
};
|
|
use tonic::transport::Server;
|
|
use std::net::SocketAddr;
|
|
use tokio::time::Duration;
|
|
|
|
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: Start single asset, single model training
|
|
#[tokio::test]
|
|
async fn test_start_single_asset_single_model() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
// Create a simple training request
|
|
let job_id = state.next_job_id();
|
|
assert_eq!(job_id, "train_job_1");
|
|
|
|
// Verify job ID is unique
|
|
let next_id = state.next_job_id();
|
|
assert_ne!(job_id, next_id);
|
|
}
|
|
|
|
/// TEST 2: Start multiple assets, all models (batch training)
|
|
#[tokio::test]
|
|
async fn test_start_multiple_assets_all_models() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
// Create 4 training jobs (4 models)
|
|
let models = vec!["TFT", "DQN", "PPO", "MAMBA_2"];
|
|
for model in &models {
|
|
let job_id = state.next_job_id();
|
|
let job = test_fixtures::create_test_job(&job_id, model, TrainingStatus::Pending);
|
|
state.add_job(job);
|
|
}
|
|
|
|
let jobs = state.list_jobs();
|
|
assert_eq!(jobs.len(), 4, "Should have 4 training jobs");
|
|
|
|
// Verify all model types present
|
|
let model_types: Vec<_> = jobs.iter().map(|j| j.model_type.as_str()).collect();
|
|
assert!(model_types.contains(&"TFT"));
|
|
assert!(model_types.contains(&"DQN"));
|
|
assert!(model_types.contains(&"PPO"));
|
|
assert!(model_types.contains(&"MAMBA_2"));
|
|
}
|
|
|
|
/// TEST 3: Invalid asset symbol - error handling
|
|
#[tokio::test]
|
|
async fn test_invalid_asset_symbol() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
// In a real scenario, this would validate the asset symbol
|
|
// For mock, we just verify error handling structure exists
|
|
let job_id = state.next_job_id();
|
|
assert!(job_id.starts_with("train_job_"));
|
|
}
|
|
|
|
/// TEST 4: File not found - clear error
|
|
#[tokio::test]
|
|
async fn test_file_not_found_error() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
// Create a job with non-existent file path (would fail in real scenario)
|
|
let job = test_fixtures::create_test_job(
|
|
"train_job_1",
|
|
"TFT",
|
|
TrainingStatus::Failed,
|
|
);
|
|
state.add_job(job.clone());
|
|
|
|
let retrieved = state.get_job("train_job_1");
|
|
assert!(retrieved.is_some());
|
|
assert_eq!(retrieved.unwrap().status, TrainingStatus::Failed as i32);
|
|
}
|
|
|
|
/// TEST 5: Duplicate job detection
|
|
#[tokio::test]
|
|
async fn test_duplicate_job_detection() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
// Create first job
|
|
let job1 = test_fixtures::create_test_job(
|
|
"train_tft_es_20251022_140000",
|
|
"TFT",
|
|
TrainingStatus::Running,
|
|
);
|
|
state.add_job(job1);
|
|
|
|
// Try to create same job (would be detected in real implementation)
|
|
let job2 = test_fixtures::create_test_job(
|
|
"train_tft_es_20251022_140000",
|
|
"TFT",
|
|
TrainingStatus::Running,
|
|
);
|
|
state.add_job(job2);
|
|
|
|
// The mock overwrites, but real implementation would reject
|
|
let jobs = state.list_jobs();
|
|
assert_eq!(jobs.len(), 1, "Duplicate job should be handled");
|
|
}
|
|
|
|
/// TEST 6: GPU availability check
|
|
#[tokio::test]
|
|
async fn test_gpu_availability() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
// Create job with GPU enabled
|
|
let job = test_fixtures::create_test_job(
|
|
"train_gpu_1",
|
|
"MAMBA_2",
|
|
TrainingStatus::Pending,
|
|
);
|
|
state.add_job(job);
|
|
|
|
let retrieved = state.get_job("train_gpu_1").unwrap();
|
|
assert_eq!(retrieved.model_type, "MAMBA_2");
|
|
}
|
|
|
|
/// TEST 7: Job creation timestamp validation
|
|
#[tokio::test]
|
|
async fn test_job_creation_timestamp() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
let before = chrono::Utc::now().timestamp();
|
|
let job = test_fixtures::create_test_job(
|
|
"train_timestamp_test",
|
|
"TFT",
|
|
TrainingStatus::Pending,
|
|
);
|
|
let after = chrono::Utc::now().timestamp();
|
|
|
|
state.add_job(job.clone());
|
|
|
|
assert!(job.created_at >= before);
|
|
assert!(job.created_at <= after + 1); // Allow 1 second tolerance
|
|
}
|
|
|
|
/// TEST 8: Job description and tags
|
|
#[tokio::test]
|
|
async fn test_job_description_and_tags() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"train_with_tags",
|
|
"DQN",
|
|
TrainingStatus::Pending,
|
|
);
|
|
|
|
assert!(!job.description.is_empty());
|
|
assert!(!job.tags.is_empty());
|
|
assert_eq!(job.tags.get("env"), Some(&"test".to_string()));
|
|
}
|
|
|
|
/// TEST 9: Multiple concurrent job submissions
|
|
#[tokio::test]
|
|
async fn test_concurrent_job_submissions() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
// Submit 10 jobs concurrently
|
|
let mut handles = vec![];
|
|
for i in 0..10 {
|
|
let state_clone = state.clone();
|
|
let handle = tokio::spawn(async move {
|
|
let job_id = format!("concurrent_job_{}", i);
|
|
let job = test_fixtures::create_test_job(
|
|
&job_id,
|
|
"TFT",
|
|
TrainingStatus::Pending,
|
|
);
|
|
state_clone.add_job(job);
|
|
});
|
|
handles.push(handle);
|
|
}
|
|
|
|
// Wait for all submissions
|
|
for handle in handles {
|
|
handle.await.unwrap();
|
|
}
|
|
|
|
let jobs = state.list_jobs();
|
|
assert_eq!(jobs.len(), 10, "Should have 10 concurrent jobs");
|
|
}
|
|
|
|
/// TEST 10: Job status initialization
|
|
#[tokio::test]
|
|
async fn test_job_status_initialization() {
|
|
let state = MockTrainingState::default();
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"train_init_test",
|
|
"PPO",
|
|
TrainingStatus::Pending,
|
|
);
|
|
state.add_job(job.clone());
|
|
|
|
assert_eq!(job.status, TrainingStatus::Pending as i32);
|
|
assert_eq!(job.started_at, 0);
|
|
assert_eq!(job.completed_at, 0);
|
|
}
|