Files
foxhunt/services/ml_training_service/tests/model_lifecycle_tests.rs
jgrusewski db6462ba7a fix(clippy): resolve all clippy warnings across entire workspace (--all-targets)
Systematic fix of 360+ clippy errors across 37+ crates covering lib,
test, bench, and example targets. Key changes:

- Add targeted #[allow(...)] on #[cfg(test)] modules for test-only lints
  (assertions_on_result_states, float_cmp, str_to_string, indexing, etc.)
- Feature-gate broken integration tests behind __<crate>_integration flags
  where public APIs changed (trading-service, backtesting-service, etc.)
- Remove dead [[test]] entries from Cargo.toml files pointing to deleted files
- Fix production code: field_reassign_with_default, manual_range_contains,
  assert!(false) → panic!(), format!("{}") simplification, len() > 0 → !is_empty()
- Delete truly unused code (Order struct, unused methods/fields/variants)
- Convert sqlx::query!() to sqlx::query() for SQLX_OFFLINE compatibility

Result: cargo clippy --workspace --all-targets -- -D warnings = 0 errors, 0 warnings

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-13 10:18:35 +01:00

780 lines
26 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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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,
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(),
mode: 0,
resume_checkpoint_path: String::new(),
max_epochs: 0,
});
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(())
}