Files
foxhunt/services/ml_training_service/tests/grpc_error_handling.rs
jgrusewski 9ffdb03e89 🚀 Wave 134: Zero Compilation Errors - 65 Agents, 194 Fixes, 530+ Tests
## Summary
- **Total Agents**: 65 (24 coverage + 41 error fixes)
- **Compilation Errors**: 194 → 0 
- **New Tests**: 530+ tests (~17,500 lines)
- **Success Rate**: 100%

## Phase 1: Test Coverage Expansion (Waves 1-3)
- Wave 1-3: 24 agents deployed
- Created comprehensive test suites across all modules
- Added 530+ tests for baseline, advanced, and integration coverage

## Phase 2: Error Elimination (Waves 4-14)
- Wave 4 (12 agents): Fixed 162 errors (Enum Display, tower util, borrow checker)
- Wave 7 (1 agent): Fixed 52 ML proto errors (DataSource, Hyperparameters)
- Wave 8 (1 agent): Fixed 33 Trading proto errors (SubmitOrderRequest)
- Wave 12 (4 agents): Fixed 13 ComplianceRequirements field errors
- Wave 13 (3 agents): Fixed 16 data crate test errors
- Wave 14 (2 agents): Fixed final 2 data lib errors

## Infrastructure Improvements
- Added MinIO Docker service for S3 E2E testing
- Created S3Config::for_minio_testing() helper
- Added storage test_helpers module
- Fixed proto field mappings across all services
- Added tower "util" feature for ServiceExt

## Key Error Patterns Fixed
- Proto field name changes (120+ instances)
- Enum Display trait usage (31 instances)
- Borrow checker errors (20+ instances)
- Missing methods/features (40+ instances)
- Struct field additions (Order, ComplianceRequirements)

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-11 17:06:02 +02:00

758 lines
27 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Comprehensive gRPC Error Handling Tests for ML Training Service
//!
//! This test suite validates all gRPC error codes and edge cases for the ML Training Service,
//! focusing on training job management, model validation, and resource constraints.
//!
//! Coverage areas:
//! - InvalidArgument: Invalid model types, bad hyperparameters, missing data sources
//! - NotFound: Non-existent training jobs, missing models
//! - FailedPrecondition: GPU unavailable, invalid job state
//! - ResourceExhausted: Too many concurrent jobs, GPU memory
//! - Internal: Training failures, data loading errors
//! - Aborted: Job cancellation, training interruption
//! - Unavailable: Service temporarily down
//!
//! Total: 13 comprehensive error scenario tests
use anyhow::Result;
use std::collections::HashMap;
use std::time::Duration;
use tonic::{Code, Request};
// Import ML Training Service proto definitions
use ml_training_service::proto::ml_training::{
ml_training_service_client::MlTrainingServiceClient, DataSource, GetTrainingJobDetailsRequest,
Hyperparameters, StartTrainingRequest, StopTrainingRequest,
SubscribeToTrainingStatusRequest,
};
// ============================================================================
// HELPER FUNCTIONS
// ============================================================================
/// Create authenticated ML Training Service client
async fn create_authenticated_client() -> Result<
MlTrainingServiceClient<
tonic::service::interceptor::InterceptedService<
tonic::transport::Channel,
impl Fn(tonic::Request<()>) -> Result<tonic::Request<()>, tonic::Status> + Clone,
>,
>,
> {
let channel = tonic::transport::Channel::from_static("http://localhost:50054")
.connect()
.await?;
// Create valid JWT token for authentication
let token = create_valid_jwt_token()?;
// Create client with interceptor to add auth header
let client =
MlTrainingServiceClient::with_interceptor(channel, move |mut req: Request<()>| {
req.metadata_mut()
.insert("authorization", format!("Bearer {}", token).parse().unwrap());
Ok(req)
});
Ok(client)
}
/// Generate valid JWT token for testing
fn create_valid_jwt_token() -> Result<String> {
use chrono::{Duration, Utc};
use jsonwebtoken::{encode, Algorithm, EncodingKey, Header};
use serde::{Deserialize, Serialize};
#[derive(Serialize, Deserialize)]
struct Claims {
sub: String,
exp: usize,
iat: usize,
iss: String,
aud: String,
roles: Vec<String>,
permissions: Vec<String>,
jti: String,
}
let jwt_secret = std::env::var("JWT_SECRET")
.unwrap_or_else(|_| "OvFLDUbIDak3CSCi5t6zKfsAp65cjTOJ85q9YE+TFY8b361DGg1gSTra2rW6mps3cWrRGQ/NXRA5uftUpMldvOaEHMMgfBs4JjVODDElREdvUFm0EttD1A==".to_string());
let claims = Claims {
sub: "test_ml_user_001".to_string(),
exp: (Utc::now() + Duration::hours(1)).timestamp() as usize,
iat: Utc::now().timestamp() as usize,
iss: "foxhunt-api-gateway".to_string(),
aud: "foxhunt-ml-training".to_string(),
roles: vec!["ml_engineer".to_string()],
permissions: vec!["ml:train".to_string(), "ml:manage".to_string()],
jti: uuid::Uuid::new_v4().to_string(),
};
let token = encode(
&Header::new(Algorithm::HS256),
&claims,
&EncodingKey::from_secret(jwt_secret.as_bytes()),
)?;
Ok(token)
}
// ============================================================================
// INVALID ARGUMENT TESTS (Validation Failures)
// ============================================================================
#[tokio::test]
async fn test_start_training_invalid_model_type_returns_invalid_argument() -> Result<()> {
println!("\n=== Test: Start Training - Invalid Model Type (InvalidArgument) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(StartTrainingRequest {
model_type: "INVALID_MODEL_TYPE".to_string(), // Invalid model type
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/tmp/test_data.parquet".to_string(),
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::DqnParams(
ml_training_service::proto::ml_training::DqnParams {
epochs: 10,
learning_rate: 0.001,
batch_size: 32,
replay_buffer_size: 10000,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay_steps: 10000,
gamma: 0.99,
target_update_frequency: 100,
use_double_dqn: false,
use_dueling: false,
use_prioritized_replay: false,
},
)),
}),
use_gpu: true,
description: "Invalid model test".to_string(),
tags: HashMap::new(),
});
let result = client.start_training(request).await;
assert!(result.is_err(), "Expected error for invalid model type");
let status = result.unwrap_err();
assert_eq!(
status.code(),
Code::InvalidArgument,
"Expected InvalidArgument error code"
);
assert!(
status.message().contains("model") || status.message().contains("type"),
"Error message should mention model type"
);
println!(" ✓ Invalid model type correctly rejected");
Ok(())
}
#[tokio::test]
async fn test_start_training_missing_data_source_returns_invalid_argument() -> Result<()> {
println!("\n=== Test: Start Training - Missing Data Source (InvalidArgument) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(StartTrainingRequest {
model_type: "MAMBA_2".to_string(),
data_source: None, // Invalid: missing data source
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::MambaParams(
ml_training_service::proto::ml_training::MambaParams {
epochs: 10,
learning_rate: 0.001,
batch_size: 32,
state_dim: 256,
hidden_dim: 512,
num_layers: 4,
dt_min: 0.001,
dt_max: 0.1,
use_cuda_kernels: true,
},
)),
}),
use_gpu: true,
description: "Missing data source test".to_string(),
tags: HashMap::new(),
});
let result = client.start_training(request).await;
assert!(result.is_err(), "Expected error for missing data source");
let status = result.unwrap_err();
assert_eq!(status.code(), Code::InvalidArgument);
assert!(
status.message().contains("data") || status.message().contains("source"),
"Error message should mention data source"
);
println!(" ✓ Missing data source correctly rejected");
Ok(())
}
#[tokio::test]
async fn test_start_training_invalid_hyperparameters_returns_invalid_argument() -> Result<()> {
println!("\n=== Test: Start Training - Invalid Hyperparameters (InvalidArgument) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(StartTrainingRequest {
model_type: "DQN".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/tmp/test_data.parquet".to_string(),
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::DqnParams(
ml_training_service::proto::ml_training::DqnParams {
epochs: 0, // Invalid: zero epochs
learning_rate: -0.001, // Invalid: negative learning rate
batch_size: 0, // Invalid: zero batch size
replay_buffer_size: 10000,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay_steps: 10000,
gamma: 0.99,
target_update_frequency: 100,
use_double_dqn: false,
use_dueling: false,
use_prioritized_replay: false,
},
)),
}),
use_gpu: false,
description: "Invalid hyperparameters test".to_string(),
tags: HashMap::new(),
});
let result = client.start_training(request).await;
assert!(result.is_err(), "Expected error for invalid hyperparameters");
let status = result.unwrap_err();
assert_eq!(status.code(), Code::InvalidArgument);
println!(" ✓ Invalid hyperparameters correctly rejected");
Ok(())
}
#[tokio::test]
async fn test_start_training_empty_symbols_returns_invalid_argument() -> Result<()> {
println!("\n=== Test: Start Training - Empty Symbols (InvalidArgument) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(StartTrainingRequest {
model_type: "TFT".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/tmp/test_data.parquet".to_string(),
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::TftParams(
ml_training_service::proto::ml_training::TftParams {
epochs: 10,
learning_rate: 0.001,
batch_size: 32,
hidden_dim: 256,
num_heads: 4,
num_layers: 3,
lookback_window: 100,
forecast_horizon: 10,
dropout_rate: 0.1,
},
)),
}),
use_gpu: false,
description: "Empty symbols test".to_string(),
tags: HashMap::new(),
});
let result = client.start_training(request).await;
assert!(result.is_err(), "Expected error for empty symbols");
let status = result.unwrap_err();
assert_eq!(status.code(), Code::InvalidArgument);
println!(" ✓ Empty symbols list correctly rejected");
Ok(())
}
// ============================================================================
// NOT FOUND TESTS (Non-existent Resources)
// ============================================================================
#[tokio::test]
async fn test_get_training_job_details_nonexistent_job_returns_not_found() -> Result<()> {
println!("\n=== Test: Get Training Job Details - Non-existent Job (NotFound) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(GetTrainingJobDetailsRequest {
job_id: "nonexistent_job_999999".to_string(),
});
let result = client.get_training_job_details(request).await;
assert!(result.is_err(), "Expected error for non-existent job");
let status = result.unwrap_err();
assert_eq!(
status.code(),
Code::NotFound,
"Expected NotFound error code"
);
assert!(
status.message().contains("not found") || status.message().contains("exist"),
"Error message should mention not found"
);
println!(" ✓ Non-existent job correctly returns NotFound");
Ok(())
}
#[tokio::test]
async fn test_stop_training_nonexistent_job_returns_not_found() -> Result<()> {
println!("\n=== Test: Stop Training - Non-existent Job (NotFound) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(StopTrainingRequest {
job_id: "nonexistent_job_888888".to_string(),
reason: "Testing not found".to_string(),
});
let result = client.stop_training(request).await;
assert!(result.is_err(), "Expected error for non-existent job");
let status = result.unwrap_err();
assert_eq!(status.code(), Code::NotFound);
println!(" ✓ Stop non-existent job correctly returns NotFound");
Ok(())
}
#[tokio::test]
async fn test_subscribe_training_status_nonexistent_job_returns_not_found() -> Result<()> {
println!("\n=== Test: Subscribe Training Status - Non-existent Job (NotFound) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(SubscribeToTrainingStatusRequest {
job_id: "nonexistent_job_777777".to_string(),
});
let result = client.subscribe_to_training_status(request).await;
assert!(result.is_err(), "Expected error for non-existent job");
let status = result.unwrap_err();
assert_eq!(status.code(), Code::NotFound);
println!(" ✓ Subscribe to non-existent job correctly returns NotFound");
Ok(())
}
// ============================================================================
// FAILED PRECONDITION TESTS (Invalid State)
// ============================================================================
#[tokio::test]
async fn test_stop_already_completed_job_returns_failed_precondition() -> Result<()> {
println!("\n=== Test: Stop Training - Already Completed (FailedPrecondition) ===");
let mut client = create_authenticated_client().await?;
// Start a very short training job
let start_request = Request::new(StartTrainingRequest {
model_type: "DQN".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/tmp/test_data_tiny.parquet".to_string(),
)),
start_time: 1609459200,
end_time: 1609459201, // 1 second
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::DqnParams(
ml_training_service::proto::ml_training::DqnParams {
epochs: 1, // Single epoch
learning_rate: 0.001,
batch_size: 1,
replay_buffer_size: 10000,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay_steps: 10000,
gamma: 0.99,
target_update_frequency: 100,
use_double_dqn: false,
use_dueling: false,
use_prioritized_replay: false,
},
)),
}),
use_gpu: false,
description: "Quick training test".to_string(),
tags: HashMap::new(),
});
let start_result = client.start_training(start_request).await?;
let job_id = start_result.into_inner().job_id;
// Wait for training to complete
tokio::time::sleep(Duration::from_secs(5)).await;
// Try to stop already completed job
let stop_request = Request::new(StopTrainingRequest {
job_id,
reason: "Testing stop completed".to_string(),
});
let result = client.stop_training(stop_request).await;
if result.is_err() {
let status = result.unwrap_err();
if status.code() == Code::FailedPrecondition {
println!(" ✓ Stop completed job correctly rejected");
} else {
println!(" Got error code: {:?}", status.code());
}
} else {
println!(" Stop succeeded (idempotent operation)");
}
Ok(())
}
#[tokio::test]
#[ignore = "Requires GPU unavailability simulation"]
async fn test_start_training_gpu_unavailable_returns_failed_precondition() -> Result<()> {
println!("\n=== Test: Start Training - GPU Unavailable (FailedPrecondition) ===");
let mut client = create_authenticated_client().await?;
// Request GPU training when GPU is not available
let request = Request::new(StartTrainingRequest {
model_type: "MAMBA_2".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/tmp/test_data.parquet".to_string(),
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::MambaParams(
ml_training_service::proto::ml_training::MambaParams {
epochs: 10,
learning_rate: 0.001,
batch_size: 32,
state_dim: 256,
hidden_dim: 512,
num_layers: 4,
dt_min: 0.001,
dt_max: 0.1,
use_cuda_kernels: true,
},
)),
}),
use_gpu: true, // Request GPU
description: "GPU unavailable test".to_string(),
tags: HashMap::new(),
});
let result = client.start_training(request).await;
if result.is_err() {
let status = result.unwrap_err();
if status.code() == Code::FailedPrecondition {
println!(" ✓ GPU unavailable correctly handled");
} else {
println!(" Request failed with: {:?}", status.code());
}
} else {
println!(" GPU available or fallback to CPU");
}
Ok(())
}
// ============================================================================
// RESOURCE EXHAUSTED TESTS (Too Many Jobs)
// ============================================================================
#[tokio::test]
#[ignore = "Slow test - requires many concurrent jobs"]
async fn test_start_training_too_many_concurrent_jobs_returns_resource_exhausted() -> Result<()> {
println!("\n=== Test: Start Training - Too Many Concurrent Jobs (ResourceExhausted) ===");
let mut client = create_authenticated_client().await?;
// Start many training jobs concurrently
let mut job_ids = vec![];
let mut exhausted = false;
for i in 0..20 {
let request = Request::new(StartTrainingRequest {
model_type: "DQN".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
format!("/tmp/stress_test_{}.parquet", i),
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::DqnParams(
ml_training_service::proto::ml_training::DqnParams {
epochs: 100,
learning_rate: 0.001,
batch_size: 32,
replay_buffer_size: 10000,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay_steps: 10000,
gamma: 0.99,
target_update_frequency: 100,
use_double_dqn: false,
use_dueling: false,
use_prioritized_replay: false,
},
)),
}),
use_gpu: false,
description: format!("Stress test job {}", i),
tags: HashMap::new(),
});
match client.start_training(request).await {
Ok(response) => {
job_ids.push(response.into_inner().job_id);
}
Err(status) => {
if status.code() == Code::ResourceExhausted {
println!(" ✓ Resource exhaustion triggered after {} jobs", i);
exhausted = true;
break;
}
}
}
}
// Clean up all started jobs
for job_id in job_ids {
let _ = client
.stop_training(Request::new(StopTrainingRequest {
job_id,
reason: "Cleanup after test".to_string(),
}))
.await;
}
if !exhausted {
println!(" Resource exhaustion not triggered (high capacity)");
}
Ok(())
}
// ============================================================================
// INTERNAL ERROR TESTS (Training Failures)
// ============================================================================
#[tokio::test]
async fn test_start_training_missing_data_file_returns_internal() -> Result<()> {
println!("\n=== Test: Start Training - Missing Data File (Internal) ===");
let mut client = create_authenticated_client().await?;
let request = Request::new(StartTrainingRequest {
model_type: "PPO".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/nonexistent/path/data.parquet".to_string(), // Non-existent file
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::PpoParams(
ml_training_service::proto::ml_training::PpoParams {
epochs: 10,
learning_rate: 0.001,
batch_size: 32,
clip_ratio: 0.2,
value_loss_coef: 0.5,
entropy_coef: 0.01,
rollout_steps: 2048,
minibatch_size: 64,
gae_lambda: 0.95,
},
)),
}),
use_gpu: false,
description: "Missing data file test".to_string(),
tags: HashMap::new(),
});
let result = client.start_training(request).await;
if result.is_err() {
let status = result.unwrap_err();
// Could be Internal (file I/O error) or InvalidArgument (invalid path)
assert!(
status.code() == Code::Internal || status.code() == Code::InvalidArgument,
"Expected Internal or InvalidArgument for missing file"
);
println!(" ✓ Missing data file correctly handled");
} else {
println!(" Request accepted (may fail during execution)");
}
Ok(())
}
// ============================================================================
// ABORTED TESTS (Job Cancellation)
// ============================================================================
#[tokio::test]
async fn test_stop_running_training_job_succeeds() -> Result<()> {
println!("\n=== Test: Stop Training - Running Job (Aborted) ===");
let mut client = create_authenticated_client().await?;
// Start a long-running training job
let start_request = Request::new(StartTrainingRequest {
model_type: "TFT".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/tmp/test_data.parquet".to_string(),
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::TftParams(
ml_training_service::proto::ml_training::TftParams {
epochs: 1000, // Many epochs
learning_rate: 0.001,
batch_size: 32,
hidden_dim: 256,
num_heads: 4,
num_layers: 3,
lookback_window: 100,
forecast_horizon: 10,
dropout_rate: 0.1,
},
)),
}),
use_gpu: false,
description: "Cancellation test".to_string(),
tags: HashMap::new(),
});
let start_result = client.start_training(start_request).await?;
let job_id = start_result.into_inner().job_id;
// Wait briefly for training to start
tokio::time::sleep(Duration::from_millis(500)).await;
// Stop the running job
let stop_request = Request::new(StopTrainingRequest {
job_id,
reason: "User requested cancellation".to_string(),
});
let result = client.stop_training(stop_request).await;
assert!(result.is_ok(), "Stop training should succeed");
println!(" ✓ Running job successfully stopped");
Ok(())
}
// ============================================================================
// DEADLINE EXCEEDED TESTS (Timeouts)
// ============================================================================
#[tokio::test]
async fn test_start_training_with_short_timeout_may_fail() -> Result<()> {
println!("\n=== Test: Start Training - Short Timeout (DeadlineExceeded) ===");
let channel = tonic::transport::Channel::from_static("http://localhost:50054")
.timeout(Duration::from_micros(1)) // Very short timeout
.connect()
.await?;
let token = create_valid_jwt_token()?;
let mut client =
MlTrainingServiceClient::with_interceptor(channel, move |mut req: Request<()>| {
req.metadata_mut()
.insert("authorization", format!("Bearer {}", token).parse().unwrap());
Ok(req)
});
let request = Request::new(StartTrainingRequest {
model_type: "DQN".to_string(),
data_source: Some(DataSource {
source: Some(ml_training_service::proto::ml_training::data_source::Source::FilePath(
"/tmp/test_data.parquet".to_string(),
)),
start_time: 1609459200,
end_time: 1640995200,
}),
hyperparameters: Some(Hyperparameters {
model_params: Some(ml_training_service::proto::ml_training::hyperparameters::ModelParams::DqnParams(
ml_training_service::proto::ml_training::DqnParams {
epochs: 10,
learning_rate: 0.001,
batch_size: 32,
replay_buffer_size: 10000,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay_steps: 10000,
gamma: 0.99,
target_update_frequency: 100,
use_double_dqn: false,
use_dueling: false,
use_prioritized_replay: false,
},
)),
}),
use_gpu: false,
description: "Timeout test".to_string(),
tags: HashMap::new(),
});
let result = client.start_training(request).await;
if result.is_err() {
let status = result.unwrap_err();
if status.code() == Code::DeadlineExceeded {
println!(" ✓ Request timed out as expected");
} else {
println!(" Request failed with: {:?}", status.code());
}
} else {
println!(" Request completed within deadline");
}
Ok(())
}