Files
foxhunt/services/ml_training_service/tests/model_lifecycle_tests.rs
jgrusewski 3817b06f19 feat(ml_training_service): add JobSpawner field to MLTrainingServiceImpl
Wire JobSpawner into the gRPC service struct so that training jobs can
be persisted to PostgreSQL before being dispatched to K8s. This is the
first step toward a durable job queue that survives pod restarts.

- Add `job_spawner: Option<Arc<JobSpawner>>` to MLTrainingServiceImpl
- Extend `new()` constructor to accept the spawner parameter
- Add `DatabaseManager::pg_pool()` accessor for cheap PgPool cloning
- Construct JobSpawner in main.rs and pass `Some(job_spawner)` to service
- Add compile-time test verifying the field exists

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-27 22:29:10 +01:00

741 lines
24 KiB
Rust

//! ML Training Service Model Lifecycle Integration Tests
//!
//! Comprehensive tests covering:
//! - Training job submission and validation
//! - Job lifecycle management (start, stop, pause)
//! - Model configuration validation
//! - Hyperparameter validation
//! - Training status updates
//! - Resource management
//! - Concurrent training jobs
//! - Error handling
use anyhow::Result;
use config::{DatabaseConfig, MLConfig};
use ml_training_service::database::DatabaseManager;
use ml_training_service::orchestrator::TrainingOrchestrator;
use ml_training_service::service::{
proto::{
ml_training_service_server::MlTrainingService, DataSource, DqnParams,
GetTrainingJobDetailsRequest, Hyperparameters, ListAvailableModelsRequest,
ListTrainingJobsRequest, MambaParams, StartTrainingRequest, StopTrainingRequest,
TlobParams, TrainingStatus,
},
MLTrainingServiceImpl,
};
use ml_training_service::storage::{ModelStorageManager, StorageConfig};
use std::collections::HashMap;
use std::sync::Arc;
use tonic::Request;
/// Setup test ML training service
async fn setup_ml_training_service() -> Result<MLTrainingServiceImpl> {
let config = MLConfig::default();
// Create test database config with proper field names
let db_config = DatabaseConfig {
url: "postgres://test:test@localhost/test_ml_training".to_string(),
max_connections: 5,
min_connections: 1,
connect_timeout: std::time::Duration::from_secs(30),
query_timeout: std::time::Duration::from_secs(30),
enable_query_logging: false,
application_name: Some("ml_training_test".to_string()),
pool: config::PoolConfig::default(),
transaction: config::TransactionConfig::default(),
};
// Create database manager
let db_manager = Arc::new(DatabaseManager::new(&db_config).await?);
// Create test storage config with proper field names
let storage_config = StorageConfig {
storage_type: "local".to_string(),
local_base_path: Some(std::path::PathBuf::from("/tmp/ml_training_test_models")),
enable_compression: false,
};
// Create storage manager
let storage_manager = Arc::new(ModelStorageManager::new(storage_config).await?);
// Create orchestrator with test dependencies
let orchestrator =
Arc::new(TrainingOrchestrator::new(config.clone(), db_manager, storage_manager).await?);
let tuning_script = std::env::var("TUNER_SCRIPT_PATH")
.unwrap_or_else(|_| "services/ml_training_service/hyperparameter_tuner.py".to_string());
let working_dir = std::env::var("TUNING_WORKING_DIR").unwrap_or_else(|_| ".".to_string());
let tuning_manager = Arc::new(ml_training_service::tuning_manager::TuningManager::new(
tuning_script, working_dir,
));
let promotion_manager = Arc::new(ml_training_service::promotion_manager::PromotionManager::new());
Ok(MLTrainingServiceImpl::new(
orchestrator,
tuning_manager,
config,
promotion_manager,
None, // no K8s dispatcher in tests
None, // no job spawner in tests
))
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_start_training_tlob_transformer() -> Result<()> {
println!("\n=== Test: Start TLOB Transformer Training ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StartTrainingRequest {
model_type: "tlob_transformer".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/orderbook_data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(
ml_training_service::service::proto::hyperparameters::ModelParams::TlobParams(
TlobParams {
epochs: 100,
learning_rate: 0.001,
batch_size: 64,
sequence_length: 50,
hidden_dim: 128,
num_heads: 8,
num_layers: 4,
dropout_rate: 0.1,
use_positional_encoding: true,
},
),
),
}),
use_gpu: true,
description: "test_tlob_training_001".to_string(),
tags: HashMap::new(),
});
let response = service.start_training(request).await?;
let job = response.into_inner();
println!("✓ TLOB training job started: {}", job.job_id);
assert!(!job.job_id.is_empty());
assert_eq!(job.status, TrainingStatus::Pending as i32);
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_start_training_mamba2() -> Result<()> {
println!("\n=== Test: Start MAMBA-2 Training ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StartTrainingRequest {
model_type: "mamba2".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/timeseries_data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(
ml_training_service::service::proto::hyperparameters::ModelParams::MambaParams(
MambaParams {
epochs: 150,
learning_rate: 0.0001,
batch_size: 32,
state_dim: 256,
hidden_dim: 512,
num_layers: 6,
dt_min: 0.001,
dt_max: 0.1,
use_cuda_kernels: true,
},
),
),
}),
use_gpu: true,
description: "test_mamba2_training_001".to_string(),
tags: HashMap::new(),
});
let response = service.start_training(request).await?;
let job = response.into_inner();
println!("✓ MAMBA-2 training job started: {}", job.job_id);
assert!(!job.job_id.is_empty());
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_start_training_dqn() -> Result<()> {
println!("\n=== Test: Start DQN Training ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StartTrainingRequest {
model_type: "dqn".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/rl_environment_data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(
ml_training_service::service::proto::hyperparameters::ModelParams::DqnParams(
DqnParams {
epochs: 200,
learning_rate: 0.0005,
batch_size: 128,
replay_buffer_size: 100000,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay_steps: 10000,
gamma: 0.99,
target_update_frequency: 100,
use_double_dqn: true,
use_dueling: false,
use_prioritized_replay: false,
},
),
),
}),
use_gpu: true,
description: "test_dqn_training_001".to_string(),
tags: HashMap::new(),
});
let response = service.start_training(request).await?;
let job = response.into_inner();
println!("✓ DQN training job started: {}", job.job_id);
assert!(!job.job_id.is_empty());
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_start_training_invalid_model_type() -> Result<()> {
println!("\n=== Test: Reject Invalid Model Type ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StartTrainingRequest {
model_type: "invalid_model_type_xyz".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: "test_invalid_model".to_string(),
tags: HashMap::new(),
});
let result = service.start_training(request).await;
assert!(result.is_err(), "Invalid model type should be rejected");
if let Err(status) = result {
println!("✓ Rejected with: {}", status.message());
assert_eq!(status.code(), tonic::Code::InvalidArgument);
assert!(status.message().contains("model type"));
}
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_start_training_empty_dataset_path() -> Result<()> {
println!("\n=== Test: Reject Empty Dataset Path ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StartTrainingRequest {
model_type: "tlob_transformer".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath("".to_string()),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: "test_empty_dataset".to_string(),
tags: HashMap::new(),
});
let result = service.start_training(request).await;
assert!(result.is_err(), "Empty dataset path should be rejected");
if let Err(status) = result {
println!("✓ Rejected with: {}", status.message());
assert!(status.message().contains("dataset") || status.message().contains("data_source"));
}
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_start_training_invalid_hyperparameters() -> Result<()> {
println!("\n=== Test: Reject Invalid Hyperparameters ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StartTrainingRequest {
model_type: "tlob_transformer".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(
ml_training_service::service::proto::hyperparameters::ModelParams::TlobParams(
TlobParams {
epochs: 0, // Invalid: zero epochs
learning_rate: -0.001, // Invalid: negative learning rate
batch_size: 0, // Invalid: zero batch size
sequence_length: 50,
hidden_dim: 128,
num_heads: 8,
num_layers: 4,
dropout_rate: 0.1,
use_positional_encoding: true,
},
),
),
}),
use_gpu: false,
description: "test_invalid_hyperparams".to_string(),
tags: HashMap::new(),
});
let result = service.start_training(request).await;
assert!(
result.is_err(),
"Invalid hyperparameters should be rejected"
);
if let Err(status) = result {
println!("✓ Rejected with: {}", status.message());
assert_eq!(status.code(), tonic::Code::InvalidArgument);
}
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_stop_training_job() -> Result<()> {
println!("\n=== Test: Stop Training Job ===");
let service = setup_ml_training_service().await?;
// Start a training job first
let start_request = Request::new(StartTrainingRequest {
model_type: "tlob_transformer".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: "test_stop_job".to_string(),
tags: HashMap::new(),
});
let start_response = service.start_training(start_request).await?;
let job_id = start_response.into_inner().job_id;
println!(" Training job started: {}", job_id);
// Stop the job
let stop_request = Request::new(StopTrainingRequest {
job_id: job_id.clone(),
reason: "test_stop".to_string(),
});
let stop_response = service.stop_training(stop_request).await?;
let stop_result = stop_response.into_inner();
println!("✓ Training job stopped: {}", job_id);
assert!(stop_result.success);
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_stop_nonexistent_job() -> Result<()> {
println!("\n=== Test: Stop Nonexistent Training Job ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StopTrainingRequest {
job_id: "nonexistent_job_12345".to_string(),
reason: "test".to_string(),
});
let result = service.stop_training(request).await;
match result {
Ok(response) => {
let stop_result = response.into_inner();
assert!(!stop_result.success, "Stopping nonexistent job should fail");
println!("✓ Stop failed as expected: {}", stop_result.message);
},
Err(status) => {
println!("✓ Rejected with: {}", status.message());
assert_eq!(status.code(), tonic::Code::NotFound);
},
}
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_get_training_job_details() -> Result<()> {
println!("\n=== Test: Get Training Job Details ===");
let service = setup_ml_training_service().await?;
// Start a training job
let start_request = Request::new(StartTrainingRequest {
model_type: "mamba2".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: "test_job_details".to_string(),
tags: HashMap::new(),
});
let start_response = service.start_training(start_request).await?;
let job_id = start_response.into_inner().job_id;
// Get job details
let details_request = Request::new(GetTrainingJobDetailsRequest {
job_id: job_id.clone(),
});
let details_response = service.get_training_job_details(details_request).await?;
let details = details_response.into_inner();
println!("✓ Job details retrieved for: {}", job_id);
if let Some(job_details) = details.job_details {
assert_eq!(job_details.job_id, job_id);
assert_eq!(job_details.description, "test_job_details");
assert_eq!(job_details.model_type, "mamba2");
println!(
" Status: {:?}",
TrainingStatus::try_from(job_details.status).unwrap_or(TrainingStatus::Unknown)
);
} else {
panic!("Expected job_details to be present");
}
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_list_training_jobs() -> Result<()> {
println!("\n=== Test: List Training Jobs ===");
let service = setup_ml_training_service().await?;
// Start a few training jobs
for i in 1..=3 {
let request = Request::new(StartTrainingRequest {
model_type: "tlob_transformer".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: format!("test_list_job_{}", i),
tags: HashMap::new(),
});
let _ = service.start_training(request).await?;
}
// List all jobs
let list_request = Request::new(ListTrainingJobsRequest {
page: 1,
page_size: 10,
status_filter: 0, // UNKNOWN = 0, means no filter
model_type_filter: "".to_string(),
start_time: 0,
end_time: 0,
});
let list_response = service.list_training_jobs(list_request).await?;
let jobs = list_response.into_inner();
println!("✓ Listed {} training jobs", jobs.jobs.len());
assert!(jobs.jobs.len() >= 3);
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_list_available_models() -> Result<()> {
println!("\n=== Test: List Available Models ===");
let service = setup_ml_training_service().await?;
let request = Request::new(ListAvailableModelsRequest {});
let response = service.list_available_models(request).await?;
let models = response.into_inner();
println!("✓ Available models:");
for model in &models.models {
println!(" - {}: {}", model.model_type, model.description);
}
assert!(!models.models.is_empty(), "Should have available models");
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_concurrent_training_jobs() -> Result<()> {
println!("\n=== Test: Concurrent Training Jobs ===");
let service = Arc::new(setup_ml_training_service().await?);
let mut handles = vec![];
// Start 3 training jobs concurrently
for i in 1..=3 {
let svc = service.clone();
let handle = tokio::spawn(async move {
let request = Request::new(StartTrainingRequest {
model_type: "tlob_transformer".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: format!("concurrent_job_{}", i),
tags: HashMap::new(),
});
svc.start_training(request).await
});
handles.push(handle);
}
// Wait for all to complete
let mut success_count = 0;
for handle in handles {
if let Ok(Ok(_)) = handle.await {
success_count += 1;
}
}
println!("{}/3 concurrent training jobs started", success_count);
assert_eq!(success_count, 3, "All concurrent jobs should start");
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_training_job_with_gpu() -> Result<()> {
println!("\n=== Test: Training Job with GPU ===");
let service = setup_ml_training_service().await?;
let request = Request::new(StartTrainingRequest {
model_type: "mamba2".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: true,
description: "test_gpu_training".to_string(),
tags: HashMap::new(),
});
let response = service.start_training(request).await?;
let job = response.into_inner();
println!("✓ Training job with GPU started: {}", job.job_id);
assert!(!job.job_id.is_empty());
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_training_job_with_tags() -> Result<()> {
println!("\n=== Test: Training Job with Tags ===");
let service = setup_ml_training_service().await?;
let mut tags = HashMap::new();
tags.insert("experiment".to_string(), "baseline".to_string());
tags.insert("version".to_string(), "v1.0".to_string());
let request = Request::new(StartTrainingRequest {
model_type: "tlob_transformer".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: "test_tagged_job".to_string(),
tags,
});
let response = service.start_training(request).await?;
let job = response.into_inner();
println!("✓ Training job with tags started: {}", job.job_id);
assert!(!job.job_id.is_empty());
Ok(())
}
#[tokio::test]
#[ignore = "Requires PostgreSQL database and test infrastructure"]
async fn test_training_job_lifecycle() -> Result<()> {
println!("\n=== Test: Complete Training Job Lifecycle ===");
let service = setup_ml_training_service().await?;
// 1. Start training
let start_request = Request::new(StartTrainingRequest {
model_type: "dqn".to_string(),
data_source: Some(DataSource {
source: Some(
ml_training_service::service::proto::data_source::Source::FilePath(
"/data/training/data.parquet".to_string(),
),
),
start_time: 0,
end_time: 0,
}),
hyperparameters: None,
use_gpu: false,
description: "test_lifecycle".to_string(),
tags: HashMap::new(),
});
let start_response = service.start_training(start_request).await?;
let job_id = start_response.into_inner().job_id;
println!(" 1. Training started: {}", job_id);
// 2. Check status
let status_request = Request::new(GetTrainingJobDetailsRequest {
job_id: job_id.clone(),
});
let status_response = service.get_training_job_details(status_request).await?;
let status_details = status_response.into_inner();
if let Some(job_details) = status_details.job_details {
println!(
" 2. Status checked: {:?}",
TrainingStatus::try_from(job_details.status).unwrap_or(TrainingStatus::Unknown)
);
}
// 3. Stop training
let stop_request = Request::new(StopTrainingRequest {
job_id: job_id.clone(),
reason: "test_lifecycle_complete".to_string(),
});
let stop_response = service.stop_training(stop_request).await?;
println!(
" 3. Training stopped: {}",
stop_response.into_inner().success
);
// 4. Verify stopped status
let final_status_request = Request::new(GetTrainingJobDetailsRequest {
job_id: job_id.clone(),
});
let final_status = service
.get_training_job_details(final_status_request)
.await?;
if let Some(job_details) = final_status.into_inner().job_details {
println!(
" 4. Final status: {:?}",
TrainingStatus::try_from(job_details.status).unwrap_or(TrainingStatus::Unknown)
);
}
println!("✓ Complete lifecycle test passed");
Ok(())
}