//! Advanced Job Spawner Tests //! //! Tests cover: //! - Transaction rollback under concurrent load //! - Deadlock prevention validation //! - Batch creation with partial failures //! - PostgreSQL connection pool exhaustion //! - State consistency under high load #![allow( dead_code, unused_variables, unused_imports, clippy::unwrap_used, clippy::expect_used, clippy::indexing_slicing )] use ml_training_service::job_spawner::{Asset, JobSpawner, ModelType}; use sqlx::PgPool; use std::path::PathBuf; use std::sync::Arc; use std::time::Instant; use tokio::time::Duration; /// Helper to run migrations async fn run_migrations(pool: &PgPool) { sqlx::migrate!("../../migrations") .run(pool) .await .expect("Failed to run migrations"); } // ============================================================================ // CONCURRENT LOAD TESTS // ============================================================================ #[sqlx::test] async fn test_concurrent_batch_creation_100_batches(pool: PgPool) { run_migrations(&pool).await; let spawner = Arc::new(JobSpawner::new(pool.clone())); let mut handles = vec![]; // Spawn 100 concurrent batch creation tasks for i in 0..100 { let spawner_clone = Arc::clone(&spawner); let handle = tokio::spawn(async move { let assets = vec![Asset { symbol: format!("ASSET{}", i), data_file: PathBuf::from(format!("/test/ASSET{}.parquet", i)), }]; let models = vec![ModelType::DQN]; spawner_clone.spawn_batch(assets, models).await }); handles.push(handle); } // Wait for all tasks let mut success_count = 0; let mut failure_count = 0; for handle in handles { match handle.await.unwrap() { Ok(_) => success_count += 1, Err(_) => failure_count += 1, } } println!("Concurrent batches: {} success, {} failure", success_count, failure_count); // All should succeed (no deadlocks) assert_eq!(success_count, 100); assert_eq!(failure_count, 0); // Verify all batches in database let count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(count, 100); } #[sqlx::test] async fn test_concurrent_same_assets_different_models(pool: PgPool) { run_migrations(&pool).await; let spawner = Arc::new(JobSpawner::new(pool.clone())); let mut handles = vec![]; let assets = vec![Asset { symbol: "ES.FUT".to_string(), data_file: PathBuf::from("/test/ES.parquet"), }]; // 4 concurrent tasks, each training different model let models = [ModelType::DQN, ModelType::PPO, ModelType::MAMBA, ModelType::TFT]; for model in models { let spawner_clone = Arc::clone(&spawner); let assets_clone = assets.clone(); let handle = tokio::spawn(async move { spawner_clone .spawn_batch(assets_clone, vec![model]) .await }); handles.push(handle); } // All should succeed for handle in handles { let result = handle.await.unwrap(); assert!(result.is_ok()); } // Should have 4 batches, 4 child jobs total let batch_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); let job_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM child_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(batch_count, 4); assert_eq!(job_count, 4); } #[sqlx::test] async fn test_high_frequency_batch_creation(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); // Create 50 batches as fast as possible let start = Instant::now(); for i in 0..50 { let assets = vec![Asset { symbol: format!("ASSET{}", i), data_file: PathBuf::from(format!("/test/ASSET{}.parquet", i)), }]; let models = vec![ModelType::DQN]; spawner .spawn_batch(assets, models) .await .expect("Batch creation failed"); } let elapsed = start.elapsed(); println!("50 batches created in {:?} ({:.2}ms per batch)", elapsed, elapsed.as_millis() as f64 / 50.0); // Target: <10ms average per batch let avg_ms = elapsed.as_millis() as f64 / 50.0; assert!(avg_ms < 100.0, "Batch creation too slow: {:.2}ms", avg_ms); // Verify count let count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(count, 50); } // ============================================================================ // TRANSACTION ROLLBACK TESTS // ============================================================================ #[sqlx::test] async fn test_rollback_on_empty_assets(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); let initial_batch_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); // Try to create batch with empty assets let result = spawner.spawn_batch(vec![], vec![ModelType::DQN]).await; assert!(result.is_err()); // Verify no batch created (rollback successful) let final_batch_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(initial_batch_count, final_batch_count); } #[sqlx::test] async fn test_rollback_on_empty_models(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); let assets = vec![Asset { symbol: "ES.FUT".to_string(), data_file: PathBuf::from("/test/ES.parquet"), }]; let initial_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); // Try with empty models let result = spawner.spawn_batch(assets, vec![]).await; assert!(result.is_err()); let final_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(initial_count, final_count); } #[sqlx::test] async fn test_rollback_consistency_child_jobs(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); let initial_batch: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); let initial_jobs: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM child_jobs") .fetch_one(&pool) .await .unwrap(); // Create invalid batch let result = spawner.spawn_batch(vec![], vec![ModelType::DQN]).await; assert!(result.is_err()); // Both tables should remain unchanged let final_batch: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); let final_jobs: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM child_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(initial_batch, final_batch); assert_eq!(initial_jobs, final_jobs); } // ============================================================================ // DEADLOCK PREVENTION // ============================================================================ #[sqlx::test] async fn test_no_deadlock_with_concurrent_batch_queries(pool: PgPool) { run_migrations(&pool).await; let spawner = Arc::new(JobSpawner::new(pool.clone())); // Create a batch first let assets = vec![Asset { symbol: "ES.FUT".to_string(), data_file: PathBuf::from("/test/ES.parquet"), }]; let batch = spawner .spawn_batch(assets, vec![ModelType::DQN, ModelType::PPO]) .await .unwrap(); let mut handles = vec![]; // Concurrent operations on same batch for i in 0..20 { let spawner_clone = Arc::clone(&spawner); let batch_id = batch.batch_id; let handle = tokio::spawn(async move { if i % 3 == 0 { // Get batch status spawner_clone.get_batch_status(batch_id).await } else if i % 3 == 1 { // Get next pending job spawner_clone.get_next_pending_job().await.map(|_| None) } else { // Get batch status again spawner_clone.get_batch_status(batch_id).await } }); handles.push(handle); } // All should complete without deadlock let timeout = Duration::from_secs(10); let results = tokio::time::timeout(timeout, async { let mut results = vec![]; for handle in handles { results.push(handle.await); } results }) .await; assert!(results.is_ok(), "Deadlock detected!"); } #[sqlx::test] async fn test_concurrent_get_next_pending_no_starvation(pool: PgPool) { run_migrations(&pool).await; let spawner = Arc::new(JobSpawner::new(pool.clone())); // Create 10 batches with jobs for i in 0..10 { let assets = vec![Asset { symbol: format!("ASSET{}", i), data_file: PathBuf::from(format!("/test/ASSET{}.parquet", i)), }]; spawner .spawn_batch(assets, vec![ModelType::DQN]) .await .unwrap(); } let mut handles = vec![]; let jobs_found = Arc::new(tokio::sync::Mutex::new(vec![])); // 10 workers concurrently getting next job for _ in 0..10 { let spawner_clone = Arc::clone(&spawner); let jobs_clone = Arc::clone(&jobs_found); let handle = tokio::spawn(async move { if let Ok(Some(job)) = spawner_clone.get_next_pending_job().await { let mut jobs = jobs_clone.lock().await; jobs.push(job.id); } }); handles.push(handle); } for handle in handles { handle.await.unwrap(); } let jobs = jobs_found.lock().await; // All 10 jobs should be claimed without duplicates assert_eq!(jobs.len(), 10); let unique: std::collections::HashSet<_> = jobs.iter().collect(); assert_eq!(unique.len(), 10, "Duplicate jobs claimed!"); } // ============================================================================ // PARTIAL FAILURE HANDLING // ============================================================================ #[sqlx::test] async fn test_large_batch_with_validation(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); // Create large batch: 10 assets × 4 models = 40 jobs let assets: Vec = (0..10) .map(|i| Asset { symbol: format!("ASSET{}", i), data_file: PathBuf::from(format!("/test/ASSET{}.parquet", i)), }) .collect(); let models = vec![ModelType::DQN, ModelType::PPO, ModelType::MAMBA, ModelType::TFT]; let batch = spawner.spawn_batch(assets, models).await.unwrap(); // Verify all 40 jobs created let job_count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM child_jobs WHERE batch_id = $1" ) .bind(batch.batch_id) .fetch_one(&pool) .await .unwrap(); assert_eq!(job_count, 40); // Verify batch metadata let (total, pending): (i32, i32) = sqlx::query_as( "SELECT total_jobs, pending_jobs FROM batch_jobs WHERE id = $1" ) .bind(batch.batch_id) .fetch_one(&pool) .await .unwrap(); assert_eq!(total, 40); assert_eq!(pending, 40); } // ============================================================================ // CONNECTION POOL STRESS // ============================================================================ #[sqlx::test] async fn test_spawn_batches_rapid_fire(pool: PgPool) { run_migrations(&pool).await; let spawner = Arc::new(JobSpawner::new(pool.clone())); // Rapidly create 30 batches without waiting let mut handles = vec![]; for i in 0..30 { let spawner_clone = Arc::clone(&spawner); let handle = tokio::spawn(async move { let assets = vec![Asset { symbol: format!("RAPID{}", i), data_file: PathBuf::from(format!("/test/RAPID{}.parquet", i)), }]; spawner_clone .spawn_batch(assets, vec![ModelType::DQN]) .await }); handles.push(handle); } // All should succeed for (i, handle) in handles.into_iter().enumerate() { let result = handle.await.unwrap(); assert!(result.is_ok(), "Batch {} failed", i); } let count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(count, 30); } // ============================================================================ // EDGE CASES // ============================================================================ #[sqlx::test] async fn test_maximum_batch_size(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); // 100 assets × 4 models = 400 jobs let assets: Vec = (0..100) .map(|i| Asset { symbol: format!("MAX{}", i), data_file: PathBuf::from(format!("/test/MAX{}.parquet", i)), }) .collect(); let models = vec![ModelType::DQN, ModelType::PPO, ModelType::MAMBA, ModelType::TFT]; let start = Instant::now(); let batch = spawner.spawn_batch(assets, models).await.unwrap(); let elapsed = start.elapsed(); println!("400-job batch created in {:?}", elapsed); // Verify let count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM child_jobs WHERE batch_id = $1" ) .bind(batch.batch_id) .fetch_one(&pool) .await .unwrap(); assert_eq!(count, 400); // Should complete in reasonable time (<1s) assert!(elapsed.as_secs() < 5, "Large batch too slow: {:?}", elapsed); } #[sqlx::test] async fn test_unicode_asset_symbols(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); let assets = vec![Asset { symbol: "测试符号".to_string(), // Chinese characters data_file: PathBuf::from("/test/unicode.parquet"), }]; let result = spawner .spawn_batch(assets, vec![ModelType::DQN]) .await; // Should handle gracefully (may succeed or fail depending on DB encoding) let _ = result; } #[sqlx::test] async fn test_very_long_file_paths(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); let long_path = "/very/long/path/".to_string() + &"subdir/".repeat(50) + "file.parquet"; let assets = vec![Asset { symbol: "ES.FUT".to_string(), data_file: PathBuf::from(long_path), }]; let result = spawner .spawn_batch(assets, vec![ModelType::DQN]) .await; assert!(result.is_ok()); } #[sqlx::test] async fn test_special_characters_in_paths(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); let assets = vec![Asset { symbol: "ES.FUT".to_string(), data_file: PathBuf::from("/test/data (2024)/ES FUT [180d].parquet"), }]; let result = spawner .spawn_batch(assets, vec![ModelType::DQN]) .await; assert!(result.is_ok()); } #[sqlx::test] async fn test_idempotency_check(pool: PgPool) { run_migrations(&pool).await; let spawner = JobSpawner::new(pool.clone()); let assets = vec![Asset { symbol: "ES.FUT".to_string(), data_file: PathBuf::from("/test/ES.parquet"), }]; // Create same batch twice let batch1 = spawner .spawn_batch(assets.clone(), vec![ModelType::DQN]) .await .unwrap(); let batch2 = spawner .spawn_batch(assets, vec![ModelType::DQN]) .await .unwrap(); // Should create separate batches (no built-in idempotency) assert_ne!(batch1.batch_id, batch2.batch_id); let count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM batch_jobs") .fetch_one(&pool) .await .unwrap(); assert_eq!(count, 2); }