512 lines
17 KiB
Rust
512 lines
17 KiB
Rust
#![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<PathBuf> {
|
|
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<usize> = 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::<usize>()
|
|
.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);
|
|
}
|