#![cfg(feature = "disabled_tests")] use rtx_preprocessing::{ DistributedDataLoader, FaultToleranceConfig, PreprocessingError, RebalancingStrategy, Result, ShardingConfig, ShardingStrategy, WorkerInfo, rtx_stubs::rtx_tensor::{Device, Tensor}, }; use std::fs::{File, create_dir_all}; use std::io::Write; use std::path::{Path, PathBuf}; use std::sync::Arc; use tempfile::TempDir; /// Helper function to create multiple test data files fn create_test_dataset(dir: &Path, file_count: usize, samples_per_file: usize) -> Vec { let mut file_paths = Vec::new(); for file_idx in 0..file_count { let file_name = format!("data_{:03}.bin", file_idx); let file_path = dir.join(&file_name); let mut file = File::create(&file_path).unwrap(); // Write sequential data for testing let start_value = file_idx * samples_per_file; for i in 0..samples_per_file { let value = (start_value + i) as f32; file.write_all(&value.to_le_bytes()).unwrap(); } file.flush().unwrap(); file_paths.push(file_path); } file_paths } #[test] fn test_distributed_data_loader_creation() { let config = ShardingConfig::default(); let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let loader = DistributedDataLoader::new(config.clone(), workers.clone()); assert_eq!(loader.config().num_shards, config.num_shards); assert_eq!(loader.workers().len(), 2); assert!(!loader.is_initialized()); } #[test] fn test_distributed_data_loader_with_round_robin_sharding() { let config = ShardingConfig { num_shards: 4, sharding_strategy: ShardingStrategy::RoundRobin, rebalancing_strategy: RebalancingStrategy::Static, fault_tolerance: FaultToleranceConfig::default(), }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let loader = DistributedDataLoader::new(config, workers); assert_eq!( loader.config().sharding_strategy, ShardingStrategy::RoundRobin ); assert_eq!(loader.config().num_shards, 4); } #[test] fn test_distributed_data_loader_initialization() { let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 8, 256); // 8 files, 256 samples each let config = ShardingConfig { num_shards: 4, sharding_strategy: ShardingStrategy::RoundRobin, rebalancing_strategy: RebalancingStrategy::Static, fault_tolerance: FaultToleranceConfig::default(), }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); // Initialize with data files let result = loader.initialize_shards(&file_paths); assert!(result.is_ok()); assert!(loader.is_initialized()); // Check that shards were created let shard_info = loader.get_shard_info(); assert_eq!(shard_info.len(), 4); // Check that all files were assigned to shards let total_files: usize = shard_info.iter().map(|s| s.file_paths.len()).sum(); assert_eq!(total_files, 8); } #[test] fn test_hash_based_sharding() { let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 12, 100); let config = ShardingConfig { num_shards: 3, sharding_strategy: ShardingStrategy::HashBased, rebalancing_strategy: RebalancingStrategy::Static, fault_tolerance: FaultToleranceConfig::default(), }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), WorkerInfo::new(2, "worker2".to_string(), "127.0.0.1:8002".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); loader.initialize_shards(&file_paths).unwrap(); let shard_info = loader.get_shard_info(); // Hash-based sharding should distribute files somewhat evenly for shard in &shard_info { assert!(shard.file_paths.len() > 0); assert!(shard.file_paths.len() <= 6); // At most half the files in one shard } // Verify same file always goes to same shard (deterministic) let first_file_shard = shard_info .iter() .find(|s| s.file_paths.contains(&file_paths[0])) .unwrap() .shard_id; // Reinitialize and check consistency loader.reset(); loader.initialize_shards(&file_paths).unwrap(); let new_shard_info = loader.get_shard_info(); let new_first_file_shard = new_shard_info .iter() .find(|s| s.file_paths.contains(&file_paths[0])) .unwrap() .shard_id; assert_eq!(first_file_shard, new_first_file_shard); } #[test] fn test_size_aware_sharding() { let temp_dir = TempDir::new().unwrap(); // Create files with different sizes let mut file_paths = Vec::new(); let sizes = vec![100, 200, 50, 300, 150]; // Different sample counts for (idx, &size) in sizes.iter().enumerate() { let file_name = format!("data_{}.bin", idx); let file_path = temp_dir.path().join(&file_name); let mut file = File::create(&file_path).unwrap(); for i in 0..size { let value = i as f32; file.write_all(&value.to_le_bytes()).unwrap(); } file.flush().unwrap(); file_paths.push(file_path); } let config = ShardingConfig { num_shards: 2, sharding_strategy: ShardingStrategy::SizeAware, rebalancing_strategy: RebalancingStrategy::Static, fault_tolerance: FaultToleranceConfig::default(), }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); loader.initialize_shards(&file_paths).unwrap(); let shard_info = loader.get_shard_info(); // Calculate total size for each shard let total_sizes: Vec = shard_info .iter() .map(|shard| { shard .file_paths .iter() .map(|path| { let file_idx = path .file_stem() .unwrap() .to_str() .unwrap() .split('_') .last() .unwrap() .parse::() .unwrap(); sizes[file_idx] * 4 // f32 = 4 bytes }) .sum() }) .collect(); // Size-aware sharding should balance shard sizes let size_diff = (total_sizes[0] as i64 - total_sizes[1] as i64).abs(); let avg_size = (total_sizes[0] + total_sizes[1]) / 2; // Sizes should be within 50% of each other for good balancing assert!(size_diff as usize <= avg_size / 2); } #[test] fn test_shard_assignment_retrieval() { let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 6, 100); let config = ShardingConfig::default(); let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers.clone()); loader.initialize_shards(&file_paths).unwrap(); // Test worker assignment for worker in &workers { let assigned_shards = loader.get_shards_for_worker(worker.worker_id); assert!(!assigned_shards.is_empty()); for shard_id in assigned_shards { let shard = loader.get_shard(shard_id).unwrap(); assert!(shard.assigned_worker_id == worker.worker_id); } } // Test file to shard mapping for file_path in &file_paths { let shard_id = loader.get_shard_for_file(file_path).unwrap(); let shard = loader.get_shard(shard_id).unwrap(); assert!(shard.file_paths.contains(file_path)); } } #[test] fn test_dynamic_rebalancing() { let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 8, 100); let config = ShardingConfig { num_shards: 4, sharding_strategy: ShardingStrategy::RoundRobin, rebalancing_strategy: RebalancingStrategy::Dynamic { imbalance_threshold: 0.3, rebalance_interval_secs: 10, }, fault_tolerance: FaultToleranceConfig::default(), }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); loader.initialize_shards(&file_paths).unwrap(); // Simulate load imbalance by marking one worker as slow loader.report_worker_performance(0, 1000); // Worker 0 slow (1000ms) loader.report_worker_performance(1, 100); // Worker 1 fast (100ms) // Check if rebalancing is triggered let needs_rebalancing = loader.needs_rebalancing(); assert!(needs_rebalancing); // Trigger rebalancing let result = loader.rebalance(); assert!(result.is_ok()); // Verify that shards were redistributed let shard_info = loader.get_shard_info(); let worker1_shards = shard_info .iter() .filter(|s| s.assigned_worker_id == 1) .count(); let worker0_shards = shard_info .iter() .filter(|s| s.assigned_worker_id == 0) .count(); // Worker 1 (faster) should get more shards assert!(worker1_shards >= worker0_shards); } #[test] fn test_fault_tolerance_detection() { let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 4, 100); let config = ShardingConfig { num_shards: 4, sharding_strategy: ShardingStrategy::RoundRobin, rebalancing_strategy: RebalancingStrategy::Static, fault_tolerance: FaultToleranceConfig { enable_health_checks: true, health_check_interval_secs: 5, max_failures: 3, recovery_timeout_secs: 30, }, }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); loader.initialize_shards(&file_paths).unwrap(); // All workers should initially be healthy assert!(loader.is_worker_healthy(0)); assert!(loader.is_worker_healthy(1)); // Simulate worker failure loader.report_worker_failure(0); loader.report_worker_failure(0); loader.report_worker_failure(0); // 3 failures = unhealthy assert!(!loader.is_worker_healthy(0)); assert!(loader.is_worker_healthy(1)); // Check that failed worker's shards are reassigned let failed_worker_shards = loader.get_shards_for_worker(0); assert!(failed_worker_shards.is_empty()); // All shards should now be assigned to healthy worker let healthy_worker_shards = loader.get_shards_for_worker(1); assert_eq!(healthy_worker_shards.len(), 4); } #[test] fn test_fault_tolerance_recovery() { let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 6, 100); let config = ShardingConfig { num_shards: 2, sharding_strategy: ShardingStrategy::RoundRobin, rebalancing_strategy: RebalancingStrategy::Static, fault_tolerance: FaultToleranceConfig { enable_health_checks: true, health_check_interval_secs: 1, max_failures: 2, recovery_timeout_secs: 5, }, }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); loader.initialize_shards(&file_paths).unwrap(); // Simulate worker failure and recovery loader.report_worker_failure(0); loader.report_worker_failure(0); // Worker 0 becomes unhealthy assert!(!loader.is_worker_healthy(0)); // Simulate worker recovery loader.report_worker_recovery(0); assert!(loader.is_worker_healthy(0)); // Verify shards are redistributed after recovery let result = loader.redistribute_after_recovery(0); assert!(result.is_ok()); let worker0_shards = loader.get_shards_for_worker(0); assert!(!worker0_shards.is_empty()); } #[test] fn test_concurrent_access() { use std::sync::Arc; use std::thread; let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 16, 50); let config = ShardingConfig { num_shards: 8, sharding_strategy: ShardingStrategy::RoundRobin, rebalancing_strategy: RebalancingStrategy::Static, fault_tolerance: FaultToleranceConfig::default(), }; let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), WorkerInfo::new(2, "worker2".to_string(), "127.0.0.1:8002".to_string()), WorkerInfo::new(3, "worker3".to_string(), "127.0.0.1:8003".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); loader.initialize_shards(&file_paths).unwrap(); let loader = Arc::new(loader); let mut handles = vec![]; // Spawn multiple threads doing concurrent operations for worker_id in 0..4 { let loader_clone = Arc::clone(&loader); let handle = thread::spawn(move || { // Each thread queries different aspects match worker_id { 0 => { // Thread 0: Query shard assignments let shards = loader_clone.get_shards_for_worker(worker_id); assert!(!shards.is_empty()); } 1 => { // Thread 1: Check worker health let is_healthy = loader_clone.is_worker_healthy(worker_id); assert!(is_healthy); } 2 => { // Thread 2: Get shard info let shard_info = loader_clone.get_shard_info(); assert_eq!(shard_info.len(), 8); } 3 => { // Thread 3: Report performance metrics loader_clone.report_worker_performance(worker_id, 150); } _ => unreachable!(), } }); handles.push(handle); } // Wait for all threads to complete for handle in handles { handle.join().unwrap(); } // Verify system is still consistent let shard_info = loader.get_shard_info(); assert_eq!(shard_info.len(), 8); let total_files: usize = shard_info.iter().map(|s| s.file_paths.len()).sum(); assert_eq!(total_files, 16); } #[test] fn test_distributed_loader_statistics() { let temp_dir = TempDir::new().unwrap(); let file_paths = create_test_dataset(temp_dir.path(), 10, 100); let config = ShardingConfig::default(); let workers = vec![ WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()), WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()), ]; let mut loader = DistributedDataLoader::new(config, workers); loader.initialize_shards(&file_paths).unwrap(); // Get initial statistics let stats = loader.get_statistics(); assert_eq!(stats.total_shards, 2); assert_eq!(stats.total_files, 10); assert_eq!(stats.active_workers, 2); assert_eq!(stats.failed_workers, 0); // Simulate some operations loader.report_worker_performance(0, 200); loader.report_worker_performance(1, 150); // Simulate worker failure loader.report_worker_failure(0); loader.report_worker_failure(0); loader.report_worker_failure(0); let updated_stats = loader.get_statistics(); assert_eq!(updated_stats.active_workers, 1); assert_eq!(updated_stats.failed_workers, 1); // Performance metrics should be recorded assert!(updated_stats.average_worker_performance > 0.0); }