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>
266 lines
7.0 KiB
Rust
266 lines
7.0 KiB
Rust
//! End-to-End Tests for `tli train status` Command
|
|
//!
|
|
//! Tests training job status queries
|
|
|
|
use fxt::proto::ml_training::{
|
|
ml_training_service_server::MlTrainingServiceServer, TrainingStatus,
|
|
ml_training_service_client::MlTrainingServiceClient, GetTrainingJobDetailsRequest,
|
|
};
|
|
use tonic::transport::{Server, Channel};
|
|
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: Query by job ID
|
|
#[tokio::test]
|
|
async fn test_query_by_job_id() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"status_test_1",
|
|
"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);
|
|
|
|
let request = GetTrainingJobDetailsRequest {
|
|
job_id: "status_test_1".to_string(),
|
|
};
|
|
|
|
let response = client.get_training_job_details(request).await;
|
|
assert!(response.is_ok(), "Should successfully query job by ID");
|
|
}
|
|
|
|
/// TEST 2: Query by batch ID
|
|
#[tokio::test]
|
|
async fn test_query_by_batch_id() {
|
|
let state = MockTrainingState::default();
|
|
|
|
// Create batch jobs
|
|
for i in 0..4 {
|
|
let job = test_fixtures::create_test_job(
|
|
&format!("batch_20251022_140000_{}", i),
|
|
"TFT",
|
|
TrainingStatus::Running,
|
|
);
|
|
state.add_job(job);
|
|
}
|
|
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
let jobs = state.list_jobs();
|
|
let batch_jobs: Vec<_> = jobs
|
|
.iter()
|
|
.filter(|j| j.job_id.contains("batch_20251022_140000"))
|
|
.collect();
|
|
|
|
assert_eq!(batch_jobs.len(), 4, "Should find all batch jobs");
|
|
}
|
|
|
|
/// TEST 3: Query by asset symbol
|
|
#[tokio::test]
|
|
async fn test_query_by_asset() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"train_tft_es_fut_20251022",
|
|
"TFT",
|
|
TrainingStatus::Completed,
|
|
);
|
|
state.add_job(job);
|
|
|
|
let _addr = start_mock_server(state.clone()).await;
|
|
|
|
let jobs = state.list_jobs();
|
|
let es_jobs: Vec<_> = jobs
|
|
.iter()
|
|
.filter(|j| j.job_id.contains("es_fut"))
|
|
.collect();
|
|
|
|
assert_eq!(es_jobs.len(), 1);
|
|
}
|
|
|
|
/// TEST 4: Non-existent job → 404 error
|
|
#[tokio::test]
|
|
async fn test_nonexistent_job_404() {
|
|
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 = GetTrainingJobDetailsRequest {
|
|
job_id: "nonexistent_job_123".to_string(),
|
|
};
|
|
|
|
let response = client.get_training_job_details(request).await;
|
|
assert!(response.is_err(), "Should return 404 for non-existent job");
|
|
assert!(response.unwrap_err().message().contains("not found"));
|
|
}
|
|
|
|
/// TEST 5: Status output formatting validation
|
|
#[tokio::test]
|
|
async fn test_status_output_formatting() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let job = test_fixtures::create_test_job_with_times(
|
|
"format_test",
|
|
"DQN",
|
|
TrainingStatus::Completed,
|
|
1729602000,
|
|
1729602060,
|
|
1729602464,
|
|
);
|
|
state.add_job(job.clone());
|
|
|
|
let _addr = start_mock_server(state).await;
|
|
|
|
// Verify job has all required fields for formatting
|
|
assert!(!job.job_id.is_empty());
|
|
assert!(!job.model_type.is_empty());
|
|
assert!(job.created_at > 0);
|
|
assert!(job.started_at > 0);
|
|
assert!(job.completed_at > 0);
|
|
assert_eq!(job.status, TrainingStatus::Completed as i32);
|
|
}
|
|
|
|
/// TEST 6: Multiple status queries (batch)
|
|
#[tokio::test]
|
|
async fn test_multiple_status_queries() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let jobs = test_fixtures::create_test_job_batch();
|
|
for job in &jobs {
|
|
state.add_job(job.clone());
|
|
}
|
|
|
|
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);
|
|
|
|
// Query each job
|
|
for job in &jobs {
|
|
let request = GetTrainingJobDetailsRequest {
|
|
job_id: job.job_id.clone(),
|
|
};
|
|
|
|
let response = client.get_training_job_details(request).await;
|
|
assert!(response.is_ok(), "Should query job {}", job.job_id);
|
|
}
|
|
}
|
|
|
|
/// TEST 7: Status query with metrics
|
|
#[tokio::test]
|
|
async fn test_status_with_metrics() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"metrics_test",
|
|
"PPO",
|
|
TrainingStatus::Running,
|
|
);
|
|
state.add_job(job.clone());
|
|
|
|
let _addr = start_mock_server(state).await;
|
|
|
|
// Verify metrics fields exist
|
|
assert!(job.final_loss >= 0.0);
|
|
assert!(job.best_validation_score >= 0.0);
|
|
}
|
|
|
|
/// TEST 8: Pending job status
|
|
#[tokio::test]
|
|
async fn test_pending_job_status() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"pending_test",
|
|
"MAMBA_2",
|
|
TrainingStatus::Pending,
|
|
);
|
|
state.add_job(job.clone());
|
|
|
|
let _addr = start_mock_server(state).await;
|
|
|
|
assert_eq!(job.status, TrainingStatus::Pending as i32);
|
|
assert_eq!(job.started_at, 0);
|
|
assert_eq!(job.completed_at, 0);
|
|
}
|
|
|
|
/// TEST 9: Failed job status with error
|
|
#[tokio::test]
|
|
async fn test_failed_job_status() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"failed_test",
|
|
"TFT",
|
|
TrainingStatus::Failed,
|
|
);
|
|
state.add_job(job.clone());
|
|
|
|
let _addr = start_mock_server(state).await;
|
|
|
|
assert_eq!(job.status, TrainingStatus::Failed as i32);
|
|
}
|
|
|
|
/// TEST 10: Stopped job status
|
|
#[tokio::test]
|
|
async fn test_stopped_job_status() {
|
|
let state = MockTrainingState::default();
|
|
|
|
let job = test_fixtures::create_test_job(
|
|
"stopped_test",
|
|
"DQN",
|
|
TrainingStatus::Stopped,
|
|
);
|
|
state.add_job(job.clone());
|
|
|
|
let _addr = start_mock_server(state).await;
|
|
|
|
assert_eq!(job.status, TrainingStatus::Stopped as i32);
|
|
}
|