use crate::traits::{EvaluationResult, QualityScore};
use crate::EvaluationError;
use scirs2_core::random::prelude::*;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::{broadcast, mpsc, RwLock};
use tokio::time::{timeout, Duration};
use uuid::Uuid;
use voirs_sdk::AudioBuffer;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DistributedConfig {
pub max_workers: usize,
pub min_workers: usize,
pub task_timeout_seconds: u64,
pub heartbeat_interval_seconds: u64,
pub max_retries: u32,
pub load_balancing: LoadBalancingStrategy,
pub fault_tolerance: FaultToleranceConfig,
pub auto_scaling: AutoScalingConfig,
pub edge_computing: EdgeComputingConfig,
pub cluster_config: ClusterConfig,
}
impl Default for DistributedConfig {
fn default() -> Self {
Self {
max_workers: 10,
min_workers: 1,
task_timeout_seconds: 300,
heartbeat_interval_seconds: 30,
max_retries: 3,
load_balancing: LoadBalancingStrategy::RoundRobin,
fault_tolerance: FaultToleranceConfig::default(),
auto_scaling: AutoScalingConfig::default(),
edge_computing: EdgeComputingConfig::default(),
cluster_config: ClusterConfig::default(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AutoScalingConfig {
pub enabled: bool,
pub target_cpu_utilization: f32,
pub target_memory_utilization: f32,
pub target_queue_depth: usize,
pub scale_up_threshold: u32,
pub scale_down_threshold: u32,
pub cooldown_seconds: u64,
pub scale_up_increment: usize,
pub scale_down_decrement: usize,
}
impl Default for AutoScalingConfig {
fn default() -> Self {
Self {
enabled: true,
target_cpu_utilization: 70.0,
target_memory_utilization: 80.0,
target_queue_depth: 20,
scale_up_threshold: 3,
scale_down_threshold: 5,
cooldown_seconds: 300,
scale_up_increment: 2,
scale_down_decrement: 1,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EdgeComputingConfig {
pub enabled: bool,
pub bandwidth_aware: bool,
pub edge_caching: bool,
pub max_bandwidth_mbps: f32,
pub compress_transfers: bool,
pub edge_to_edge_enabled: bool,
pub edge_latency_threshold_ms: u64,
}
impl Default for EdgeComputingConfig {
fn default() -> Self {
Self {
enabled: false,
bandwidth_aware: true,
edge_caching: true,
max_bandwidth_mbps: 100.0,
compress_transfers: true,
edge_to_edge_enabled: false,
edge_latency_threshold_ms: 50,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClusterConfig {
pub enabled: bool,
pub cluster_name: String,
pub consensus_enabled: bool,
pub consensus_algorithm: ConsensusAlgorithm,
pub replication_enabled: bool,
pub replication_factor: usize,
pub partition_tolerance: bool,
pub discovery_method: ClusterDiscoveryMethod,
}
impl Default for ClusterConfig {
fn default() -> Self {
Self {
enabled: false,
cluster_name: "voirs-eval-cluster".to_string(),
consensus_enabled: false,
consensus_algorithm: ConsensusAlgorithm::Raft,
replication_enabled: false,
replication_factor: 3,
partition_tolerance: true,
discovery_method: ClusterDiscoveryMethod::Static,
}
}
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub enum ConsensusAlgorithm {
Raft,
Paxos,
LeaderElection,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub enum ClusterDiscoveryMethod {
Static,
Dns,
Kubernetes,
Consul,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub enum LoadBalancingStrategy {
RoundRobin,
LeastLoaded,
Random,
Weighted,
LatencyAware,
BandwidthAware,
Adaptive,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FaultToleranceConfig {
pub auto_redistribute: bool,
pub health_monitoring: bool,
pub max_node_failures: u32,
pub result_validation: bool,
}
impl Default for FaultToleranceConfig {
fn default() -> Self {
Self {
auto_redistribute: true,
health_monitoring: true,
max_node_failures: 3,
result_validation: true,
}
}
}
pub type TaskId = Uuid;
pub type WorkerId = Uuid;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EvaluationTask {
pub id: TaskId,
pub task_type: TaskType,
pub audio_data: Vec<u8>, pub reference_data: Option<Vec<u8>>,
pub parameters: TaskParameters,
pub priority: TaskPriority,
pub max_execution_time: Duration,
pub retry_count: u32,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum TaskType {
QualityMetrics,
PronunciationAssessment,
ComparativeAnalysis,
PerceptualEvaluation,
StatisticalAnalysis,
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskParameters {
pub metrics: Vec<String>,
pub language: Option<String>,
pub sample_rate: Option<u32>,
pub channels: Option<u16>,
pub custom_params: HashMap<String, String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum TaskPriority {
Low,
Normal,
High,
Critical,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskResult {
pub task_id: TaskId,
pub worker_id: WorkerId,
pub result: Result<EvaluationOutput, String>,
pub execution_time: Duration,
pub resource_usage: ResourceUsage,
pub completed_at: std::time::SystemTime,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EvaluationOutput {
pub quality_scores: HashMap<String, f32>,
pub metrics: HashMap<String, serde_json::Value>,
pub metadata: HashMap<String, String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ResourceUsage {
pub cpu_usage: f32,
pub memory_usage: f32,
pub disk_io: f32,
pub network_io: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerInfo {
pub id: WorkerId,
pub name: String,
pub capabilities: WorkerCapabilities,
pub status: WorkerStatus,
pub current_load: f32,
pub last_heartbeat: std::time::SystemTime,
pub performance_metrics: PerformanceMetrics,
pub network_metrics: NetworkMetrics,
pub location: Option<WorkerLocation>,
pub is_edge_node: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerCapabilities {
pub max_concurrent_tasks: usize,
pub supported_task_types: Vec<TaskType>,
pub available_memory: f32,
pub cpu_cores: usize,
pub specialized_hardware: Vec<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum WorkerStatus {
Online,
Busy,
Offline,
Failed,
Draining,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NetworkMetrics {
pub latency_ms: f32,
pub bandwidth_mbps: f32,
pub packet_loss: f32,
pub jitter_ms: f32,
pub total_data_transferred_mb: f32,
}
impl Default for NetworkMetrics {
fn default() -> Self {
Self {
latency_ms: 10.0,
bandwidth_mbps: 100.0,
packet_loss: 0.0,
jitter_ms: 1.0,
total_data_transferred_mb: 0.0,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerLocation {
pub region: String,
pub zone: Option<String>,
pub latitude: Option<f64>,
pub longitude: Option<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PerformanceMetrics {
pub tasks_completed: u64,
pub tasks_failed: u64,
pub avg_execution_time: Duration,
pub success_rate: f32,
pub throughput: f32,
}
pub struct DistributedEvaluator {
config: DistributedConfig,
workers: Arc<RwLock<HashMap<WorkerId, WorkerInfo>>>,
task_queue: Arc<RwLock<Vec<EvaluationTask>>>,
running_tasks: Arc<RwLock<HashMap<TaskId, (WorkerId, std::time::SystemTime)>>>,
completed_tasks: Arc<RwLock<HashMap<TaskId, TaskResult>>>,
task_sender: mpsc::UnboundedSender<EvaluationTask>,
result_receiver: Arc<RwLock<mpsc::UnboundedReceiver<TaskResult>>>,
shutdown_sender: broadcast::Sender<()>,
stats: Arc<RwLock<SystemStatistics>>,
scaling_state: Arc<RwLock<AutoScalingState>>,
cluster_state: Arc<RwLock<ClusterState>>,
}
#[derive(Debug, Clone, Default)]
pub struct AutoScalingState {
pub last_scaling_action: Option<std::time::SystemTime>,
pub scale_up_counter: u32,
pub scale_down_counter: u32,
pub current_system_load: f32,
pub scaling_recommendation: ScalingRecommendation,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ScalingRecommendation {
ScaleUp,
ScaleDown,
#[default]
NoAction,
}
#[derive(Debug, Clone, Default)]
pub struct ClusterState {
pub is_leader: bool,
pub leader_id: Option<WorkerId>,
pub members: Vec<WorkerId>,
pub last_election: Option<std::time::SystemTime>,
pub term: u64,
}
#[derive(Debug, Clone, Default)]
pub struct SystemStatistics {
pub tasks_submitted: u64,
pub tasks_completed: u64,
pub tasks_failed: u64,
pub total_execution_time: Duration,
pub avg_completion_time: Duration,
pub system_throughput: f32,
pub active_workers: usize,
pub failed_workers: usize,
}
impl DistributedEvaluator {
pub fn new(config: DistributedConfig) -> Self {
let (task_sender, task_receiver) = mpsc::unbounded_channel();
let (result_sender, result_receiver) = mpsc::unbounded_channel();
let (shutdown_sender, _) = broadcast::channel(1);
let evaluator = Self {
config: config.clone(),
workers: Arc::new(RwLock::new(HashMap::new())),
task_queue: Arc::new(RwLock::new(Vec::new())),
running_tasks: Arc::new(RwLock::new(HashMap::new())),
completed_tasks: Arc::new(RwLock::new(HashMap::new())),
task_sender,
result_receiver: Arc::new(RwLock::new(result_receiver)),
shutdown_sender: shutdown_sender.clone(),
stats: Arc::new(RwLock::new(SystemStatistics::default())),
scaling_state: Arc::new(RwLock::new(AutoScalingState::default())),
cluster_state: Arc::new(RwLock::new(ClusterState::default())),
};
evaluator.start_background_tasks(task_receiver, result_sender);
if config.auto_scaling.enabled {
evaluator.start_auto_scaling_monitor();
}
if config.cluster_config.enabled {
evaluator.start_cluster_management();
}
evaluator
}
fn start_background_tasks(
&self,
mut task_receiver: mpsc::UnboundedReceiver<EvaluationTask>,
result_sender: mpsc::UnboundedSender<TaskResult>,
) {
let workers = Arc::clone(&self.workers);
let running_tasks = Arc::clone(&self.running_tasks);
let completed_tasks = Arc::clone(&self.completed_tasks);
let config = self.config.clone();
let stats = Arc::clone(&self.stats);
let mut shutdown_receiver = self.shutdown_sender.subscribe();
tokio::spawn(async move {
loop {
tokio::select! {
Some(task) = task_receiver.recv() => {
Self::distribute_task(
task,
Arc::clone(&workers),
Arc::clone(&running_tasks),
result_sender.clone(),
config.clone(),
).await;
}
_ = shutdown_receiver.recv() => {
break;
}
}
}
});
let workers_monitor = Arc::clone(&self.workers);
let stats_monitor = Arc::clone(&self.stats);
let config_monitor = self.config.clone();
let mut shutdown_monitor = self.shutdown_sender.subscribe();
tokio::spawn(async move {
let mut interval = tokio::time::interval(Duration::from_secs(
config_monitor.heartbeat_interval_seconds,
));
loop {
tokio::select! {
_ = interval.tick() => {
Self::monitor_worker_health(
Arc::clone(&workers_monitor),
Arc::clone(&stats_monitor),
).await;
}
_ = shutdown_monitor.recv() => {
break;
}
}
}
});
}
pub async fn register_worker(&self, worker_info: WorkerInfo) -> EvaluationResult<()> {
let mut workers = self.workers.write().await;
workers.insert(worker_info.id, worker_info);
let mut stats = self.stats.write().await;
stats.active_workers = workers.len();
Ok(())
}
pub async fn unregister_worker(&self, worker_id: WorkerId) -> EvaluationResult<()> {
let mut workers = self.workers.write().await;
workers.remove(&worker_id);
let mut stats = self.stats.write().await;
stats.active_workers = workers.len();
Ok(())
}
pub async fn submit_task(&self, task: EvaluationTask) -> EvaluationResult<TaskId> {
let task_id = task.id;
let mut queue = self.task_queue.write().await;
queue.push(task.clone());
queue.sort_by_key(|b| std::cmp::Reverse(b.priority));
self.task_sender
.send(task)
.map_err(|e| EvaluationError::QualityEvaluationError {
message: format!("Failed to submit task: {}", e),
source: None,
})?;
let mut stats = self.stats.write().await;
stats.tasks_submitted += 1;
Ok(task_id)
}
pub async fn get_result(&self, task_id: TaskId) -> EvaluationResult<Option<TaskResult>> {
let completed_tasks = self.completed_tasks.read().await;
Ok(completed_tasks.get(&task_id).cloned())
}
pub async fn wait_for_completion(&self, task_id: TaskId) -> EvaluationResult<TaskResult> {
let timeout_duration = Duration::from_secs(self.config.task_timeout_seconds);
timeout(timeout_duration, async {
loop {
if let Some(result) = self.get_result(task_id).await? {
return Ok(result);
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
})
.await
.map_err(|_| EvaluationError::QualityEvaluationError {
message: "Task completion timeout".to_string(),
source: None,
})?
}
async fn distribute_task(
task: EvaluationTask,
workers: Arc<RwLock<HashMap<WorkerId, WorkerInfo>>>,
running_tasks: Arc<RwLock<HashMap<TaskId, (WorkerId, std::time::SystemTime)>>>,
result_sender: mpsc::UnboundedSender<TaskResult>,
config: DistributedConfig,
) {
let selected_worker = Self::select_worker(&workers, &task, config.load_balancing).await;
if let Some(worker_id) = selected_worker {
let mut running = running_tasks.write().await;
running.insert(task.id, (worker_id, std::time::SystemTime::now()));
let task_clone = task.clone();
tokio::spawn(async move {
let result = Self::execute_task_on_worker(task_clone, worker_id).await;
let _ = result_sender.send(result);
});
}
}
async fn select_worker(
workers: &Arc<RwLock<HashMap<WorkerId, WorkerInfo>>>,
task: &EvaluationTask,
strategy: LoadBalancingStrategy,
) -> Option<WorkerId> {
let workers_read = workers.read().await;
let available_workers: Vec<_> = workers_read
.values()
.filter(|w| w.status == WorkerStatus::Online)
.filter(|w| {
w.capabilities
.supported_task_types
.contains(&task.task_type)
})
.collect();
if available_workers.is_empty() {
return None;
}
match strategy {
LoadBalancingStrategy::RoundRobin => {
available_workers.first().map(|w| w.id)
}
LoadBalancingStrategy::LeastLoaded => {
available_workers
.iter()
.min_by(|a, b| {
a.current_load
.partial_cmp(&b.current_load)
.unwrap_or(std::cmp::Ordering::Equal)
})
.map(|w| w.id)
}
LoadBalancingStrategy::Random => {
if available_workers.is_empty() {
None
} else {
let mut rng = scirs2_core::random::Random::seed(0);
let idx = rng.random_range(0..available_workers.len());
Some(available_workers[idx].id)
}
}
LoadBalancingStrategy::Weighted => {
available_workers
.iter()
.max_by_key(|w| w.capabilities.max_concurrent_tasks)
.map(|w| w.id)
}
LoadBalancingStrategy::LatencyAware => {
available_workers
.iter()
.min_by(|a, b| {
a.network_metrics
.latency_ms
.partial_cmp(&b.network_metrics.latency_ms)
.expect("value should be present")
})
.map(|w| w.id)
}
LoadBalancingStrategy::BandwidthAware => {
available_workers
.iter()
.max_by(|a, b| {
a.network_metrics
.bandwidth_mbps
.partial_cmp(&b.network_metrics.bandwidth_mbps)
.expect("value should be present")
})
.map(|w| w.id)
}
LoadBalancingStrategy::Adaptive => {
available_workers
.iter()
.min_by(|a, b| {
let score_a = Self::calculate_worker_score(a);
let score_b = Self::calculate_worker_score(b);
score_a
.partial_cmp(&score_b)
.unwrap_or(std::cmp::Ordering::Equal)
})
.map(|w| w.id)
}
}
}
fn calculate_worker_score(worker: &WorkerInfo) -> f32 {
let load_weight = 0.4;
let latency_weight = 0.3;
let bandwidth_weight = 0.2;
let success_rate_weight = 0.1;
let load_score = worker.current_load * 100.0;
let latency_score = worker.network_metrics.latency_ms;
let bandwidth_score = 100.0 - worker.network_metrics.bandwidth_mbps.min(100.0);
let success_rate_score = (1.0 - worker.performance_metrics.success_rate) * 100.0;
load_weight * load_score
+ latency_weight * latency_score
+ bandwidth_weight * bandwidth_score
+ success_rate_weight * success_rate_score
}
async fn execute_task_on_worker(task: EvaluationTask, worker_id: WorkerId) -> TaskResult {
use scirs2_core::random::{rngs::StdRng, Rng, SeedableRng};
let seed = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
let mut rng = Random::seed(seed);
let start_time = std::time::SystemTime::now();
let execution_time = Duration::from_millis(rng.random::<u64>() % 5000 + 1000);
tokio::time::sleep(execution_time).await;
let result = if rng.random::<f32>() > 0.1 {
Ok(EvaluationOutput {
quality_scores: {
let mut scores = HashMap::new();
scores.insert("pesq".to_string(), rng.random::<f32>() * 5.0);
scores.insert("stoi".to_string(), rng.random::<f32>());
scores
},
metrics: HashMap::new(),
metadata: HashMap::new(),
})
} else {
Err("Simulated task failure".to_string())
};
TaskResult {
task_id: task.id,
worker_id,
result,
execution_time,
resource_usage: ResourceUsage {
cpu_usage: rng.random::<f32>() * 100.0,
memory_usage: rng.random::<f32>() * 1024.0,
disk_io: rng.random::<f32>() * 100.0,
network_io: rng.random::<f32>() * 50.0,
},
completed_at: start_time,
}
}
async fn monitor_worker_health(
workers: Arc<RwLock<HashMap<WorkerId, WorkerInfo>>>,
stats: Arc<RwLock<SystemStatistics>>,
) {
let mut workers_write = workers.write().await;
let mut failed_count = 0;
for worker in workers_write.values_mut() {
let now = std::time::SystemTime::now();
let time_since_heartbeat = now
.duration_since(worker.last_heartbeat)
.unwrap_or_default();
if time_since_heartbeat > Duration::from_secs(60) {
worker.status = WorkerStatus::Failed;
failed_count += 1;
}
}
let mut stats_write = stats.write().await;
stats_write.failed_workers = failed_count;
stats_write.active_workers = workers_write.len() - failed_count;
}
pub async fn get_statistics(&self) -> SystemStatistics {
self.stats.read().await.clone()
}
pub async fn get_workers(&self) -> Vec<WorkerInfo> {
let workers = self.workers.read().await;
workers.values().cloned().collect()
}
fn start_auto_scaling_monitor(&self) {
let workers = Arc::clone(&self.workers);
let task_queue = Arc::clone(&self.task_queue);
let scaling_state = Arc::clone(&self.scaling_state);
let config = self.config.clone();
let mut shutdown_receiver = self.shutdown_sender.subscribe();
tokio::spawn(async move {
let mut interval = tokio::time::interval(Duration::from_secs(30));
loop {
tokio::select! {
_ = interval.tick() => {
Self::evaluate_scaling_needs(
Arc::clone(&workers),
Arc::clone(&task_queue),
Arc::clone(&scaling_state),
&config.auto_scaling,
config.min_workers,
config.max_workers,
).await;
}
_ = shutdown_receiver.recv() => {
break;
}
}
}
});
}
async fn evaluate_scaling_needs(
workers: Arc<RwLock<HashMap<WorkerId, WorkerInfo>>>,
task_queue: Arc<RwLock<Vec<EvaluationTask>>>,
scaling_state: Arc<RwLock<AutoScalingState>>,
config: &AutoScalingConfig,
min_workers: usize,
max_workers: usize,
) {
let workers_read = workers.read().await;
let queue_read = task_queue.read().await;
let mut state = scaling_state.write().await;
let active_workers = workers_read
.values()
.filter(|w| w.status == WorkerStatus::Online)
.count();
let queue_depth = queue_read.len();
let total_cpu: f32 = workers_read.values().map(|w| w.current_load).sum();
let avg_cpu = if active_workers > 0 {
total_cpu / active_workers as f32 * 100.0
} else {
0.0
};
let total_memory: f32 = workers_read
.values()
.map(|w| {
(w.capabilities.available_memory - w.performance_metrics.throughput)
/ w.capabilities.available_memory
* 100.0
})
.sum();
let avg_memory = if active_workers > 0 {
total_memory / active_workers as f32
} else {
0.0
};
state.current_system_load = (avg_cpu + avg_memory) / 2.0;
let in_cooldown = if let Some(last_action) = state.last_scaling_action {
let elapsed = std::time::SystemTime::now()
.duration_since(last_action)
.unwrap_or_default();
elapsed < Duration::from_secs(config.cooldown_seconds)
} else {
false
};
if in_cooldown {
return;
}
let should_scale_up = active_workers < max_workers
&& (avg_cpu > config.target_cpu_utilization
|| avg_memory > config.target_memory_utilization
|| queue_depth > config.target_queue_depth);
let should_scale_down = active_workers > min_workers
&& avg_cpu < config.target_cpu_utilization * 0.5
&& avg_memory < config.target_memory_utilization * 0.5
&& queue_depth < config.target_queue_depth / 2;
if should_scale_up {
state.scale_up_counter += 1;
state.scale_down_counter = 0;
if state.scale_up_counter >= config.scale_up_threshold {
state.scaling_recommendation = ScalingRecommendation::ScaleUp;
state.scale_up_counter = 0;
state.last_scaling_action = Some(std::time::SystemTime::now());
}
} else if should_scale_down {
state.scale_down_counter += 1;
state.scale_up_counter = 0;
if state.scale_down_counter >= config.scale_down_threshold {
state.scaling_recommendation = ScalingRecommendation::ScaleDown;
state.scale_down_counter = 0;
state.last_scaling_action = Some(std::time::SystemTime::now());
}
} else {
state.scale_up_counter = 0;
state.scale_down_counter = 0;
state.scaling_recommendation = ScalingRecommendation::NoAction;
}
}
pub async fn get_scaling_recommendation(&self) -> ScalingRecommendation {
let state = self.scaling_state.read().await;
state.scaling_recommendation
}
pub async fn get_scaling_state(&self) -> AutoScalingState {
self.scaling_state.read().await.clone()
}
fn start_cluster_management(&self) {
let cluster_state = Arc::clone(&self.cluster_state);
let config = self.config.cluster_config.clone();
let mut shutdown_receiver = self.shutdown_sender.subscribe();
tokio::spawn(async move {
let mut interval = tokio::time::interval(Duration::from_secs(10));
loop {
tokio::select! {
_ = interval.tick() => {
Self::perform_cluster_maintenance(
Arc::clone(&cluster_state),
&config,
).await;
}
_ = shutdown_receiver.recv() => {
break;
}
}
}
});
}
async fn perform_cluster_maintenance(
cluster_state: Arc<RwLock<ClusterState>>,
config: &ClusterConfig,
) {
if !config.consensus_enabled {
return;
}
let mut state = cluster_state.write().await;
if state.leader_id.is_none() {
state.term += 1;
state.is_leader = true; state.leader_id = Some(Uuid::new_v4());
state.last_election = Some(std::time::SystemTime::now());
}
}
pub async fn get_cluster_state(&self) -> ClusterState {
self.cluster_state.read().await.clone()
}
pub async fn is_cluster_leader(&self) -> bool {
let state = self.cluster_state.read().await;
state.is_leader
}
pub async fn shutdown(&self) -> EvaluationResult<()> {
let _ = self.shutdown_sender.send(());
Ok(())
}
}
pub fn create_evaluation_task(
task_type: TaskType,
audio: &AudioBuffer,
reference: Option<&AudioBuffer>,
parameters: TaskParameters,
) -> EvaluationTask {
EvaluationTask {
id: Uuid::new_v4(),
task_type,
audio_data: serialize_audio_buffer(audio),
reference_data: reference.map(serialize_audio_buffer),
parameters,
priority: TaskPriority::Normal,
max_execution_time: Duration::from_secs(300),
retry_count: 0,
}
}
fn serialize_audio_buffer(audio: &AudioBuffer) -> Vec<u8> {
oxicode::serde::encode_to_vec(audio, oxicode::config::standard()).unwrap_or_default()
}
pub fn deserialize_audio_buffer(data: &[u8]) -> Option<AudioBuffer> {
oxicode::serde::decode_from_slice(data, oxicode::config::standard())
.ok()
.map(|(v, _)| v)
}
#[cfg(test)]
mod tests {
use super::*;
use voirs_sdk::AudioBuffer;
#[test]
fn test_distributed_config_default() {
let config = DistributedConfig::default();
assert_eq!(config.max_workers, 10);
assert_eq!(config.task_timeout_seconds, 300);
assert!(matches!(
config.load_balancing,
LoadBalancingStrategy::RoundRobin
));
}
#[test]
fn test_task_creation() {
let samples = vec![0.1, 0.2, -0.1, -0.2];
let audio = AudioBuffer::new(samples, 16000, 1);
let parameters = TaskParameters {
metrics: vec!["pesq".to_string(), "stoi".to_string()],
language: Some("en".to_string()),
sample_rate: Some(16000),
channels: Some(1),
custom_params: HashMap::new(),
};
let task = create_evaluation_task(TaskType::QualityMetrics, &audio, None, parameters);
assert!(matches!(task.task_type, TaskType::QualityMetrics));
assert_eq!(task.priority, TaskPriority::Normal);
assert!(!task.audio_data.is_empty());
}
#[tokio::test]
async fn test_distributed_evaluator_creation() {
let config = DistributedConfig::default();
let evaluator = DistributedEvaluator::new(config);
let stats = evaluator.get_statistics().await;
assert_eq!(stats.active_workers, 0);
assert_eq!(stats.tasks_submitted, 0);
}
#[tokio::test]
async fn test_worker_registration() {
let config = DistributedConfig::default();
let evaluator = DistributedEvaluator::new(config);
let worker_info = WorkerInfo {
id: Uuid::new_v4(),
name: "test-worker".to_string(),
capabilities: WorkerCapabilities {
max_concurrent_tasks: 4,
supported_task_types: vec![TaskType::QualityMetrics],
available_memory: 1024.0,
cpu_cores: 4,
specialized_hardware: vec![],
},
status: WorkerStatus::Online,
current_load: 0.0,
last_heartbeat: std::time::SystemTime::now(),
performance_metrics: PerformanceMetrics {
tasks_completed: 0,
tasks_failed: 0,
avg_execution_time: Duration::from_secs(0),
success_rate: 0.0,
throughput: 0.0,
},
network_metrics: NetworkMetrics::default(),
location: None,
is_edge_node: false,
};
evaluator.register_worker(worker_info).await.unwrap();
let workers = evaluator.get_workers().await;
assert_eq!(workers.len(), 1);
assert_eq!(workers[0].name, "test-worker");
}
#[tokio::test]
async fn test_task_submission() {
let config = DistributedConfig::default();
let evaluator = DistributedEvaluator::new(config);
let samples = vec![0.1, 0.2, -0.1, -0.2];
let audio = AudioBuffer::new(samples, 16000, 1);
let parameters = TaskParameters {
metrics: vec!["pesq".to_string()],
language: Some("en".to_string()),
sample_rate: Some(16000),
channels: Some(1),
custom_params: HashMap::new(),
};
let task = create_evaluation_task(TaskType::QualityMetrics, &audio, None, parameters);
let task_id = evaluator.submit_task(task).await.unwrap();
let stats = evaluator.get_statistics().await;
assert_eq!(stats.tasks_submitted, 1);
let result = evaluator.get_result(task_id).await.unwrap();
assert!(result.is_none());
}
#[test]
fn test_task_priority_ordering() {
let high = TaskPriority::High;
let low = TaskPriority::Low;
let critical = TaskPriority::Critical;
assert!(critical > high);
assert!(high > low);
}
#[test]
fn test_audio_serialization() {
let samples = vec![0.1, 0.2, -0.1, -0.2];
let audio = AudioBuffer::new(samples.clone(), 16000, 1);
let serialized = serialize_audio_buffer(&audio);
assert!(!serialized.is_empty());
let deserialized = deserialize_audio_buffer(&serialized);
assert!(deserialized.is_some());
let restored_audio = deserialized.unwrap();
assert_eq!(restored_audio.sample_rate(), 16000);
assert_eq!(restored_audio.channels(), 1);
}
#[tokio::test]
async fn test_auto_scaling_config() {
let config = AutoScalingConfig::default();
assert!(config.enabled);
assert_eq!(config.target_cpu_utilization, 70.0);
assert_eq!(config.target_memory_utilization, 80.0);
assert_eq!(config.scale_up_threshold, 3);
assert_eq!(config.scale_down_threshold, 5);
}
#[tokio::test]
async fn test_auto_scaling_state() {
let config = DistributedConfig {
auto_scaling: AutoScalingConfig {
enabled: true,
..Default::default()
},
..Default::default()
};
let evaluator = DistributedEvaluator::new(config);
let state = evaluator.get_scaling_state().await;
assert_eq!(state.scale_up_counter, 0);
assert_eq!(state.scale_down_counter, 0);
assert_eq!(
state.scaling_recommendation,
ScalingRecommendation::NoAction
);
}
#[tokio::test]
async fn test_cluster_state() {
let config = DistributedConfig {
cluster_config: ClusterConfig {
enabled: true,
consensus_enabled: true,
..Default::default()
},
..Default::default()
};
let evaluator = DistributedEvaluator::new(config);
tokio::time::sleep(Duration::from_millis(100)).await;
let state = evaluator.get_cluster_state().await;
assert!(state.term >= 1);
assert!(state.leader_id.is_some());
}
#[tokio::test]
async fn test_edge_computing_config() {
let config = EdgeComputingConfig::default();
assert!(!config.enabled); assert!(config.bandwidth_aware);
assert!(config.edge_caching);
assert_eq!(config.max_bandwidth_mbps, 100.0);
}
#[tokio::test]
async fn test_network_metrics() {
let metrics = NetworkMetrics::default();
assert_eq!(metrics.latency_ms, 10.0);
assert_eq!(metrics.bandwidth_mbps, 100.0);
assert_eq!(metrics.packet_loss, 0.0);
}
#[tokio::test]
async fn test_worker_location() {
let location = WorkerLocation {
region: "us-east-1".to_string(),
zone: Some("us-east-1a".to_string()),
latitude: Some(40.7128),
longitude: Some(-74.0060),
};
assert_eq!(location.region, "us-east-1");
assert_eq!(location.zone, Some("us-east-1a".to_string()));
}
#[tokio::test]
async fn test_latency_aware_load_balancing() {
let config = DistributedConfig {
load_balancing: LoadBalancingStrategy::LatencyAware,
..Default::default()
};
let evaluator = DistributedEvaluator::new(config);
let worker1 = WorkerInfo {
id: Uuid::new_v4(),
name: "high-latency".to_string(),
capabilities: WorkerCapabilities {
max_concurrent_tasks: 4,
supported_task_types: vec![TaskType::QualityMetrics],
available_memory: 1024.0,
cpu_cores: 4,
specialized_hardware: vec![],
},
status: WorkerStatus::Online,
current_load: 0.5,
last_heartbeat: std::time::SystemTime::now(),
performance_metrics: PerformanceMetrics {
tasks_completed: 10,
tasks_failed: 0,
avg_execution_time: Duration::from_secs(2),
success_rate: 1.0,
throughput: 5.0,
},
network_metrics: NetworkMetrics {
latency_ms: 50.0,
bandwidth_mbps: 100.0,
packet_loss: 0.0,
jitter_ms: 1.0,
total_data_transferred_mb: 100.0,
},
location: None,
is_edge_node: false,
};
let worker2 = WorkerInfo {
id: Uuid::new_v4(),
name: "low-latency".to_string(),
capabilities: WorkerCapabilities {
max_concurrent_tasks: 4,
supported_task_types: vec![TaskType::QualityMetrics],
available_memory: 1024.0,
cpu_cores: 4,
specialized_hardware: vec![],
},
status: WorkerStatus::Online,
current_load: 0.5,
last_heartbeat: std::time::SystemTime::now(),
performance_metrics: PerformanceMetrics {
tasks_completed: 10,
tasks_failed: 0,
avg_execution_time: Duration::from_secs(2),
success_rate: 1.0,
throughput: 5.0,
},
network_metrics: NetworkMetrics {
latency_ms: 10.0, bandwidth_mbps: 100.0,
packet_loss: 0.0,
jitter_ms: 1.0,
total_data_transferred_mb: 100.0,
},
location: None,
is_edge_node: false,
};
evaluator.register_worker(worker1).await.unwrap();
evaluator.register_worker(worker2).await.unwrap();
let workers = evaluator.get_workers().await;
assert_eq!(workers.len(), 2);
}
#[tokio::test]
async fn test_bandwidth_aware_load_balancing() {
let config = DistributedConfig {
load_balancing: LoadBalancingStrategy::BandwidthAware,
..Default::default()
};
let evaluator = DistributedEvaluator::new(config);
let worker = WorkerInfo {
id: Uuid::new_v4(),
name: "high-bandwidth".to_string(),
capabilities: WorkerCapabilities {
max_concurrent_tasks: 4,
supported_task_types: vec![TaskType::QualityMetrics],
available_memory: 1024.0,
cpu_cores: 4,
specialized_hardware: vec![],
},
status: WorkerStatus::Online,
current_load: 0.5,
last_heartbeat: std::time::SystemTime::now(),
performance_metrics: PerformanceMetrics {
tasks_completed: 10,
tasks_failed: 0,
avg_execution_time: Duration::from_secs(2),
success_rate: 1.0,
throughput: 5.0,
},
network_metrics: NetworkMetrics {
latency_ms: 10.0,
bandwidth_mbps: 1000.0, packet_loss: 0.0,
jitter_ms: 1.0,
total_data_transferred_mb: 100.0,
},
location: None,
is_edge_node: false,
};
evaluator.register_worker(worker).await.unwrap();
let workers = evaluator.get_workers().await;
assert_eq!(workers.len(), 1);
assert_eq!(workers[0].network_metrics.bandwidth_mbps, 1000.0);
}
#[tokio::test]
async fn test_adaptive_load_balancing() {
let config = DistributedConfig {
load_balancing: LoadBalancingStrategy::Adaptive,
..Default::default()
};
let evaluator = DistributedEvaluator::new(config);
let worker = WorkerInfo {
id: Uuid::new_v4(),
name: "adaptive-worker".to_string(),
capabilities: WorkerCapabilities {
max_concurrent_tasks: 4,
supported_task_types: vec![TaskType::QualityMetrics],
available_memory: 1024.0,
cpu_cores: 4,
specialized_hardware: vec![],
},
status: WorkerStatus::Online,
current_load: 0.3,
last_heartbeat: std::time::SystemTime::now(),
performance_metrics: PerformanceMetrics {
tasks_completed: 100,
tasks_failed: 2,
avg_execution_time: Duration::from_secs(1),
success_rate: 0.98,
throughput: 10.0,
},
network_metrics: NetworkMetrics {
latency_ms: 15.0,
bandwidth_mbps: 500.0,
packet_loss: 0.1,
jitter_ms: 2.0,
total_data_transferred_mb: 1000.0,
},
location: Some(WorkerLocation {
region: "us-west-2".to_string(),
zone: Some("us-west-2a".to_string()),
latitude: Some(45.5231),
longitude: Some(-122.6765),
}),
is_edge_node: true,
};
evaluator.register_worker(worker).await.unwrap();
let workers = evaluator.get_workers().await;
assert_eq!(workers.len(), 1);
assert!(workers[0].is_edge_node);
}
#[test]
fn test_scaling_recommendation() {
assert_eq!(
ScalingRecommendation::default(),
ScalingRecommendation::NoAction
);
}
#[test]
fn test_consensus_algorithm() {
let config = ClusterConfig::default();
assert!(matches!(
config.consensus_algorithm,
ConsensusAlgorithm::Raft
));
}
#[test]
fn test_cluster_discovery_method() {
let config = ClusterConfig::default();
assert!(matches!(
config.discovery_method,
ClusterDiscoveryMethod::Static
));
}
}