use std::collections::HashMap;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ParallelTask {
pub id: String,
pub description: String,
pub working_directory: String,
pub depends_on: Vec<String>,
pub max_iterations: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ParallelTaskResult {
pub task_id: String,
pub success: bool,
pub summary: String,
pub iterations: u32,
pub cost: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ParallelPlanStatus {
Pending,
Running {
completed: usize,
total: usize,
},
Completed {
results: Vec<ParallelTaskResult>,
},
Failed {
reason: String,
partial_results: Vec<ParallelTaskResult>,
},
}
#[derive(Debug, Clone)]
pub struct ParallelConfig {
pub max_concurrent: usize,
pub use_mdap: bool,
pub fail_fast: bool,
}
impl Default for ParallelConfig {
fn default() -> Self {
Self {
max_concurrent: 5,
use_mdap: false,
fail_fast: false,
}
}
}
pub struct ParallelCoordinator {
config: ParallelConfig,
tasks: Vec<ParallelTask>,
results: HashMap<String, ParallelTaskResult>,
}
impl ParallelCoordinator {
pub fn new(config: ParallelConfig) -> Self {
Self {
config,
tasks: Vec::new(),
results: HashMap::new(),
}
}
pub fn add_task(&mut self, task: ParallelTask) {
self.tasks.push(task);
}
pub fn ready_tasks(&self) -> Vec<&ParallelTask> {
self.tasks
.iter()
.filter(|t| {
!self.results.contains_key(&t.id)
&& t.depends_on
.iter()
.all(|dep| self.results.get(dep).is_some_and(|r| r.success))
})
.collect()
}
pub fn record_result(&mut self, result: ParallelTaskResult) {
self.results.insert(result.task_id.clone(), result);
}
pub fn is_complete(&self) -> bool {
self.tasks.iter().all(|t| self.results.contains_key(&t.id))
}
pub fn has_failure(&self) -> bool {
self.results.values().any(|r| !r.success)
}
pub fn status(&self) -> ParallelPlanStatus {
if self.results.is_empty() && !self.tasks.is_empty() {
return ParallelPlanStatus::Pending;
}
let completed = self.results.len();
let total = self.tasks.len();
if completed < total {
if self.config.fail_fast && self.has_failure() {
return ParallelPlanStatus::Failed {
reason: "fail-fast: a task failed".to_string(),
partial_results: self.results.values().cloned().collect(),
};
}
return ParallelPlanStatus::Running { completed, total };
}
let results: Vec<ParallelTaskResult> = self.results.values().cloned().collect();
if results.iter().all(|r| r.success) {
ParallelPlanStatus::Completed { results }
} else {
ParallelPlanStatus::Failed {
reason: "one or more tasks failed".to_string(),
partial_results: results,
}
}
}
pub fn stats(&self) -> ParallelStats {
let results: Vec<&ParallelTaskResult> = self.results.values().collect();
ParallelStats {
total_tasks: self.tasks.len(),
completed: results.len(),
succeeded: results.iter().filter(|r| r.success).count(),
failed: results.iter().filter(|r| !r.success).count(),
total_iterations: results.iter().map(|r| r.iterations as u64).sum(),
total_cost: results.iter().map(|r| r.cost).sum(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ParallelStats {
pub total_tasks: usize,
pub completed: usize,
pub succeeded: usize,
pub failed: usize,
pub total_iterations: u64,
pub total_cost: f64,
}
#[cfg(test)]
mod tests {
use super::*;
fn make_task(id: &str, deps: Vec<&str>) -> ParallelTask {
ParallelTask {
id: id.to_string(),
description: format!("Task {id}"),
working_directory: "/tmp".to_string(),
depends_on: deps.into_iter().map(|s| s.to_string()).collect(),
max_iterations: 10,
}
}
fn make_result(task_id: &str, success: bool) -> ParallelTaskResult {
ParallelTaskResult {
task_id: task_id.to_string(),
success,
summary: "done".to_string(),
iterations: 5,
cost: 0.01,
}
}
#[test]
fn new_coordinator_is_empty() {
let coord = ParallelCoordinator::new(ParallelConfig::default());
assert!(coord.is_complete()); assert!(!coord.has_failure());
}
#[test]
fn add_task_and_ready_tasks() {
let mut coord = ParallelCoordinator::new(ParallelConfig::default());
coord.add_task(make_task("a", vec![]));
coord.add_task(make_task("b", vec!["a"]));
let ready = coord.ready_tasks();
assert_eq!(ready.len(), 1);
assert_eq!(ready[0].id, "a");
}
#[test]
fn record_result_unlocks_dependents() {
let mut coord = ParallelCoordinator::new(ParallelConfig::default());
coord.add_task(make_task("a", vec![]));
coord.add_task(make_task("b", vec!["a"]));
coord.record_result(make_result("a", true));
let ready = coord.ready_tasks();
assert_eq!(ready.len(), 1);
assert_eq!(ready[0].id, "b");
}
#[test]
fn failed_dependency_blocks_dependents() {
let mut coord = ParallelCoordinator::new(ParallelConfig::default());
coord.add_task(make_task("a", vec![]));
coord.add_task(make_task("b", vec!["a"]));
coord.record_result(make_result("a", false));
let ready = coord.ready_tasks();
assert!(ready.is_empty());
}
#[test]
fn is_complete_when_all_done() {
let mut coord = ParallelCoordinator::new(ParallelConfig::default());
coord.add_task(make_task("a", vec![]));
coord.add_task(make_task("b", vec![]));
assert!(!coord.is_complete());
coord.record_result(make_result("a", true));
assert!(!coord.is_complete());
coord.record_result(make_result("b", true));
assert!(coord.is_complete());
}
#[test]
fn stats_aggregates_correctly() {
let mut coord = ParallelCoordinator::new(ParallelConfig::default());
coord.add_task(make_task("a", vec![]));
coord.add_task(make_task("b", vec![]));
coord.record_result(make_result("a", true));
coord.record_result(make_result("b", false));
let stats = coord.stats();
assert_eq!(stats.total_tasks, 2);
assert_eq!(stats.completed, 2);
assert_eq!(stats.succeeded, 1);
assert_eq!(stats.failed, 1);
assert_eq!(stats.total_iterations, 10);
assert!((stats.total_cost - 0.02).abs() < f64::EPSILON);
}
#[test]
fn status_transitions() {
let mut coord = ParallelCoordinator::new(ParallelConfig::default());
coord.add_task(make_task("a", vec![]));
assert!(matches!(coord.status(), ParallelPlanStatus::Pending));
coord.record_result(make_result("a", true));
assert!(matches!(
coord.status(),
ParallelPlanStatus::Completed { .. }
));
}
}