1use crate::dag::{ResourcePool, TaskNode, WorkflowDag, create_execution_plan};
4use crate::engine::state::{
5 StatePersistence, TaskStatus, WorkflowCheckpoint, WorkflowState, WorkflowStatus,
6};
7use crate::error::{Result, WorkflowError};
8use async_trait::async_trait;
9use std::sync::Arc;
10use std::sync::atomic::{AtomicU64, Ordering};
11use std::time::Duration;
12use tokio::sync::{RwLock, Semaphore};
13use tokio::time::timeout;
14use tracing::{debug, error, info, warn};
15
16#[async_trait]
18pub trait TaskExecutor: Send + Sync {
19 async fn execute(&self, task: &TaskNode, context: &ExecutionContext) -> Result<TaskOutput>;
21}
22
23#[derive(Debug, Clone)]
25pub struct ExecutionContext {
26 pub execution_id: String,
28 pub task_id: String,
30 pub state: Arc<RwLock<WorkflowState>>,
32 pub inputs: std::collections::HashMap<String, serde_json::Value>,
34}
35
36#[derive(Debug, Clone)]
38pub struct TaskOutput {
39 pub data: Option<serde_json::Value>,
41 pub logs: Vec<String>,
43}
44
45#[derive(Debug, Clone)]
47pub struct ExecutorConfig {
48 pub max_concurrent_tasks: usize,
50 pub enable_persistence: bool,
52 pub state_dir: String,
54 pub resource_pool: ResourcePool,
56 pub retry_on_failure: bool,
58 pub stop_on_failure: bool,
60 pub checkpoint_interval: usize,
62 pub enable_checkpointing: bool,
64}
65
66impl Default for ExecutorConfig {
67 fn default() -> Self {
68 Self {
69 max_concurrent_tasks: 10,
70 enable_persistence: true,
71 state_dir: std::env::temp_dir()
72 .join("oxigdal-workflow")
73 .to_string_lossy()
74 .into_owned(),
75 resource_pool: ResourcePool::default(),
76 retry_on_failure: true,
77 stop_on_failure: false,
78 checkpoint_interval: 1, enable_checkpointing: true,
80 }
81 }
82}
83
84pub struct WorkflowExecutor<E: TaskExecutor> {
86 config: ExecutorConfig,
88 task_executor: Arc<E>,
90 persistence: Option<StatePersistence>,
92 semaphore: Arc<Semaphore>,
94 checkpoint_sequence: AtomicU64,
96 tasks_since_checkpoint: AtomicU64,
98}
99
100impl<E: TaskExecutor> WorkflowExecutor<E> {
101 pub fn new(config: ExecutorConfig, task_executor: E) -> Self {
103 let semaphore = Arc::new(Semaphore::new(config.max_concurrent_tasks));
104 let persistence = if config.enable_persistence {
105 Some(StatePersistence::new(config.state_dir.clone()))
106 } else {
107 None
108 };
109
110 Self {
111 config,
112 task_executor: Arc::new(task_executor),
113 persistence,
114 semaphore,
115 checkpoint_sequence: AtomicU64::new(0),
116 tasks_since_checkpoint: AtomicU64::new(0),
117 }
118 }
119
120 async fn maybe_save_checkpoint(&self, state: &WorkflowState, dag: &WorkflowDag) -> Result<()> {
122 if !self.config.enable_checkpointing {
123 return Ok(());
124 }
125
126 let persistence = match &self.persistence {
127 Some(p) => p,
128 None => return Ok(()),
129 };
130
131 let tasks_completed = self.tasks_since_checkpoint.fetch_add(1, Ordering::SeqCst) + 1;
132
133 if tasks_completed >= self.config.checkpoint_interval as u64 {
134 self.tasks_since_checkpoint.store(0, Ordering::SeqCst);
135 let seq = self.checkpoint_sequence.fetch_add(1, Ordering::SeqCst);
136
137 let checkpoint = WorkflowCheckpoint::new(state.clone(), dag.clone(), seq);
138 persistence.save_checkpoint(&checkpoint).await?;
139
140 debug!(
141 "Saved checkpoint {} for execution {}",
142 seq, state.execution_id
143 );
144 }
145
146 Ok(())
147 }
148
149 async fn save_checkpoint_now(&self, state: &WorkflowState, dag: &WorkflowDag) -> Result<()> {
151 if !self.config.enable_checkpointing {
152 return Ok(());
153 }
154
155 let persistence = match &self.persistence {
156 Some(p) => p,
157 None => return Ok(()),
158 };
159
160 self.tasks_since_checkpoint.store(0, Ordering::SeqCst);
161 let seq = self.checkpoint_sequence.fetch_add(1, Ordering::SeqCst);
162
163 let checkpoint = WorkflowCheckpoint::new(state.clone(), dag.clone(), seq);
164 persistence.save_checkpoint(&checkpoint).await?;
165
166 info!(
167 "Saved checkpoint {} for execution {}",
168 seq, state.execution_id
169 );
170 Ok(())
171 }
172
173 pub async fn execute(
175 &self,
176 workflow_id: String,
177 execution_id: String,
178 dag: WorkflowDag,
179 ) -> Result<WorkflowState> {
180 info!(
181 "Starting workflow execution: workflow_id={}, execution_id={}",
182 workflow_id, execution_id
183 );
184
185 dag.validate()?;
187
188 let mut state = WorkflowState::new(workflow_id.clone(), execution_id.clone(), workflow_id);
190
191 for task in dag.tasks() {
193 state.init_task(task.id.clone());
194 }
195
196 state.start();
197
198 if let Some(ref persistence) = self.persistence {
200 persistence.save(&state).await?;
201 }
202
203 self.save_checkpoint_now(&state, &dag).await?;
205
206 let state_arc = Arc::new(RwLock::new(state));
207
208 let execution_plan = create_execution_plan(&dag)?;
210
211 info!(
212 "Execution plan created with {} levels",
213 execution_plan.len()
214 );
215
216 for (level_idx, level) in execution_plan.iter().enumerate() {
218 info!("Executing level {} with {} tasks", level_idx, level.len());
219
220 let results = self.execute_level(&dag, &state_arc, level).await;
221
222 {
224 let state_guard = state_arc.read().await;
225 self.maybe_save_checkpoint(&state_guard, &dag).await?;
226 }
227
228 let failed_tasks: Vec<_> = results
230 .iter()
231 .filter_map(|(task_id, result)| {
232 if result.is_err() {
233 Some(task_id.clone())
234 } else {
235 None
236 }
237 })
238 .collect();
239
240 if !failed_tasks.is_empty() {
241 error!("Tasks failed: {:?}", failed_tasks);
242
243 if self.config.stop_on_failure {
244 warn!("Stopping workflow execution due to failures");
245 let mut state_guard = state_arc.write().await;
246 state_guard.fail();
247
248 if let Some(ref persistence) = self.persistence {
249 persistence.save(&state_guard).await?;
250 }
251
252 self.save_checkpoint_now(&state_guard, &dag).await?;
254
255 drop(state_guard);
256
257 return Ok(Arc::try_unwrap(state_arc)
258 .map(|rw| rw.into_inner())
259 .unwrap_or_else(|arc| {
260 tokio::task::block_in_place(|| arc.blocking_read().clone())
261 }));
262 }
263 }
264 }
265
266 let mut state_guard = state_arc.write().await;
268
269 let all_completed = state_guard
271 .task_states
272 .values()
273 .all(|ts| ts.status == TaskStatus::Completed || ts.status == TaskStatus::Skipped);
274
275 if all_completed {
276 state_guard.complete();
277 } else {
278 state_guard.fail();
279 }
280
281 if let Some(ref persistence) = self.persistence {
283 persistence.save(&state_guard).await?;
284 }
285
286 self.save_checkpoint_now(&state_guard, &dag).await?;
288
289 info!(
290 "Workflow execution completed: status={:?}",
291 state_guard.status
292 );
293
294 drop(state_guard);
295
296 Ok(Arc::try_unwrap(state_arc)
297 .map(|rw| rw.into_inner())
298 .unwrap_or_else(|arc| tokio::task::block_in_place(|| arc.blocking_read().clone())))
299 }
300
301 async fn execute_level(
315 &self,
316 dag: &WorkflowDag,
317 state: &Arc<RwLock<WorkflowState>>,
318 level: &[String],
319 ) -> Vec<(String, Result<()>)> {
320 use futures::stream::{FuturesUnordered, StreamExt};
321
322 let mut in_flight = FuturesUnordered::new();
323
324 for task_id in level {
325 let semaphore = Arc::clone(&self.semaphore);
326 in_flight.push(async move {
327 let _permit = semaphore.acquire_owned().await.ok();
331 let result = self
332 .execute_task(
333 task_id,
334 dag,
335 state,
336 &*self.task_executor,
337 self.config.retry_on_failure,
338 )
339 .await;
340 (task_id.clone(), result)
341 });
342 }
343
344 let mut results = Vec::with_capacity(level.len());
345 while let Some(item) = in_flight.next().await {
346 results.push(item);
347 }
348
349 results
350 }
351
352 async fn execute_task(
354 &self,
355 task_id: &str,
356 dag: &WorkflowDag,
357 state: &Arc<RwLock<WorkflowState>>,
358 executor: &E,
359 retry_on_failure: bool,
360 ) -> Result<()> {
361 let task = dag
362 .get_task(task_id)
363 .ok_or_else(|| WorkflowError::not_found(format!("Task '{}'", task_id)))?;
364
365 debug!("Executing task: {}", task_id);
366
367 if !self.check_dependencies(task_id, dag, state).await? {
369 warn!("Skipping task {} due to failed dependencies", task_id);
370 let mut state_guard = state.write().await;
371 state_guard.skip_task(task_id)?;
372 return Ok(());
373 }
374
375 {
377 let mut state_guard = state.write().await;
378 state_guard.start_task(task_id)?;
379 }
380
381 let max_attempts = if retry_on_failure {
383 task.retry.max_attempts
384 } else {
385 1
386 };
387
388 let mut last_error = None;
389
390 for attempt in 0..max_attempts {
391 if attempt > 0 {
392 debug!("Retrying task {} (attempt {})", task_id, attempt + 1);
393
394 let delay_ms =
395 task.retry.delay_ms as f64 * task.retry.backoff_multiplier.powi(attempt as i32);
396 let delay_ms = delay_ms.min(task.retry.max_delay_ms as f64) as u64;
397
398 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
399 }
400
401 let inputs = self.gather_inputs(task_id, dag, state).await?;
403
404 let ctx = ExecutionContext {
406 execution_id: {
407 let state_guard = state.read().await;
408 state_guard.execution_id.clone()
409 },
410 task_id: task_id.to_string(),
411 state: Arc::clone(state),
412 inputs,
413 };
414
415 let task_timeout = Duration::from_secs(task.timeout_secs.unwrap_or(300));
417 let execute_result = timeout(task_timeout, executor.execute(task, &ctx)).await;
418
419 match execute_result {
420 Ok(Ok(output)) => {
421 let mut state_guard = state.write().await;
423 state_guard.complete_task(task_id, output.data)?;
424
425 for log in output.logs {
426 state_guard.add_task_log(task_id, log)?;
427 }
428
429 info!("Task {} completed successfully", task_id);
430 return Ok(());
431 }
432 Ok(Err(e)) => {
433 warn!("Task {} failed: {}", task_id, e);
434 last_error = Some(e);
435 }
436 Err(_) => {
437 let timeout_error =
438 WorkflowError::task_timeout(task_id, task_timeout.as_secs());
439 warn!("Task {} timed out", task_id);
440 last_error = Some(timeout_error);
441 }
442 }
443 }
444
445 let error = last_error.unwrap_or_else(|| WorkflowError::execution("Unknown error"));
447 let mut state_guard = state.write().await;
448 state_guard.fail_task(task_id, error.to_string())?;
449
450 error!("Task {} failed after {} attempts", task_id, max_attempts);
451 Err(error)
452 }
453
454 async fn check_dependencies(
456 &self,
457 task_id: &str,
458 dag: &WorkflowDag,
459 state: &Arc<RwLock<WorkflowState>>,
460 ) -> Result<bool> {
461 let dependencies = dag.get_dependencies(task_id);
462 let state_guard = state.read().await;
463
464 for dep_id in dependencies {
465 if let Some(dep_state) = state_guard.get_task_state(&dep_id) {
466 if dep_state.status != TaskStatus::Completed {
467 return Ok(false);
468 }
469 } else {
470 return Ok(false);
471 }
472 }
473
474 Ok(true)
475 }
476
477 async fn gather_inputs(
479 &self,
480 task_id: &str,
481 dag: &WorkflowDag,
482 state: &Arc<RwLock<WorkflowState>>,
483 ) -> Result<std::collections::HashMap<String, serde_json::Value>> {
484 let dependencies = dag.get_dependencies(task_id);
485 let state_guard = state.read().await;
486 let mut inputs = std::collections::HashMap::new();
487
488 for dep_id in dependencies {
489 if let Some(dep_state) = state_guard.get_task_state(&dep_id)
490 && let Some(ref output) = dep_state.output
491 {
492 inputs.insert(dep_id.clone(), output.clone());
493 }
494 }
495
496 Ok(inputs)
497 }
498
499 pub async fn resume(&self, execution_id: String) -> Result<WorkflowState> {
508 let persistence = self
509 .persistence
510 .as_ref()
511 .ok_or_else(|| WorkflowError::state("Persistence is not enabled"))?;
512
513 let mut checkpoint = persistence.load_checkpoint(&execution_id).await.map_err(|e| {
515 WorkflowError::state(format!(
516 "Failed to load checkpoint for recovery: {}. Ensure checkpointing was enabled during execution.",
517 e
518 ))
519 })?;
520
521 if checkpoint.state.is_terminal() {
522 return Err(WorkflowError::state("Cannot resume a terminal workflow"));
523 }
524
525 info!(
526 "Resuming workflow execution: execution_id={}, checkpoint_sequence={}",
527 execution_id, checkpoint.sequence
528 );
529
530 checkpoint.prepare_for_resume()?;
532
533 self.checkpoint_sequence
535 .store(checkpoint.sequence + 1, Ordering::SeqCst);
536
537 self.resume_from_checkpoint(checkpoint).await
539 }
540
541 async fn resume_from_checkpoint(
543 &self,
544 checkpoint: WorkflowCheckpoint,
545 ) -> Result<WorkflowState> {
546 let dag = checkpoint.dag.clone();
547
548 let completed = checkpoint.get_completed_tasks();
550 let pending = checkpoint.get_pending_tasks();
551 let interrupted = checkpoint.get_interrupted_tasks();
552 let failed = checkpoint.get_failed_tasks();
553
554 let mut state = checkpoint.state;
555
556 if state.status != WorkflowStatus::Running {
558 state.status = WorkflowStatus::Running;
559 }
560
561 info!(
562 "Recovery state: {} completed, {} pending, {} interrupted, {} failed",
563 completed.len(),
564 pending.len(),
565 interrupted.len(),
566 failed.len()
567 );
568
569 if let Some(ref persistence) = self.persistence {
571 persistence.save(&state).await?;
572 }
573
574 let state_arc = Arc::new(RwLock::new(state));
575
576 let execution_plan = create_execution_plan(&dag)?;
578
579 info!("Resuming execution with {} levels", execution_plan.len());
580
581 for (level_idx, level) in execution_plan.iter().enumerate() {
583 let tasks_to_execute: Vec<String> = {
585 let state_guard = state_arc.read().await;
586 level
587 .iter()
588 .filter(|task_id| {
589 state_guard
590 .get_task_state(task_id)
591 .map(|ts| {
592 !matches!(ts.status, TaskStatus::Completed | TaskStatus::Skipped)
593 })
594 .unwrap_or(true)
595 })
596 .cloned()
597 .collect()
598 };
599
600 if tasks_to_execute.is_empty() {
601 debug!("Level {} has no tasks to execute, skipping", level_idx);
602 continue;
603 }
604
605 info!(
606 "Resuming level {} with {} tasks (skipping {} completed)",
607 level_idx,
608 tasks_to_execute.len(),
609 level.len() - tasks_to_execute.len()
610 );
611
612 let results = self
613 .execute_level(&dag, &state_arc, &tasks_to_execute)
614 .await;
615
616 {
618 let state_guard = state_arc.read().await;
619 self.maybe_save_checkpoint(&state_guard, &dag).await?;
620 }
621
622 let failed_tasks: Vec<_> = results
624 .iter()
625 .filter_map(|(task_id, result)| {
626 if result.is_err() {
627 Some(task_id.clone())
628 } else {
629 None
630 }
631 })
632 .collect();
633
634 if !failed_tasks.is_empty() {
635 error!("Tasks failed during resume: {:?}", failed_tasks);
636
637 if self.config.stop_on_failure {
638 warn!("Stopping resumed workflow execution due to failures");
639 let mut state_guard = state_arc.write().await;
640 state_guard.fail();
641
642 if let Some(ref persistence) = self.persistence {
643 persistence.save(&state_guard).await?;
644 }
645
646 self.save_checkpoint_now(&state_guard, &dag).await?;
648
649 drop(state_guard);
650
651 return Ok(Arc::try_unwrap(state_arc)
652 .map(|rw| rw.into_inner())
653 .unwrap_or_else(|arc| {
654 tokio::task::block_in_place(|| arc.blocking_read().clone())
655 }));
656 }
657 }
658 }
659
660 let mut state_guard = state_arc.write().await;
662
663 let all_completed = state_guard
665 .task_states
666 .values()
667 .all(|ts| ts.status == TaskStatus::Completed || ts.status == TaskStatus::Skipped);
668
669 if all_completed {
670 state_guard.complete();
671 } else {
672 state_guard.fail();
673 }
674
675 if let Some(ref persistence) = self.persistence {
677 persistence.save(&state_guard).await?;
678 }
679
680 self.save_checkpoint_now(&state_guard, &dag).await?;
682
683 info!(
684 "Resumed workflow execution completed: status={:?}",
685 state_guard.status
686 );
687
688 drop(state_guard);
689
690 Ok(Arc::try_unwrap(state_arc)
691 .map(|rw| rw.into_inner())
692 .unwrap_or_else(|arc| tokio::task::block_in_place(|| arc.blocking_read().clone())))
693 }
694
695 pub async fn resume_from_sequence(
697 &self,
698 execution_id: String,
699 sequence: u64,
700 ) -> Result<WorkflowState> {
701 let persistence = self
702 .persistence
703 .as_ref()
704 .ok_or_else(|| WorkflowError::state("Persistence is not enabled"))?;
705
706 let mut checkpoint = persistence
707 .load_checkpoint_by_sequence(&execution_id, sequence)
708 .await?;
709
710 if checkpoint.state.is_terminal() {
711 return Err(WorkflowError::state("Cannot resume a terminal workflow"));
712 }
713
714 info!(
715 "Resuming workflow from specific checkpoint: execution_id={}, sequence={}",
716 execution_id, sequence
717 );
718
719 checkpoint.prepare_for_resume()?;
721
722 self.checkpoint_sequence
724 .store(sequence + 1, Ordering::SeqCst);
725
726 self.resume_from_checkpoint(checkpoint).await
727 }
728
729 pub async fn get_recovery_info(&self, execution_id: &str) -> Result<RecoveryInfo> {
731 let persistence = self
732 .persistence
733 .as_ref()
734 .ok_or_else(|| WorkflowError::state("Persistence is not enabled"))?;
735
736 let checkpoint = persistence.load_checkpoint(execution_id).await?;
737
738 Ok(RecoveryInfo {
739 execution_id: execution_id.to_string(),
740 checkpoint_sequence: checkpoint.sequence,
741 checkpoint_created_at: checkpoint.created_at,
742 workflow_status: checkpoint.state.status,
743 completed_tasks: checkpoint.get_completed_tasks(),
744 pending_tasks: checkpoint.get_pending_tasks(),
745 interrupted_tasks: checkpoint.get_interrupted_tasks(),
746 failed_tasks: checkpoint.get_failed_tasks(),
747 skipped_tasks: checkpoint.get_skipped_tasks(),
748 can_resume: !checkpoint.state.is_terminal(),
749 })
750 }
751
752 pub async fn list_checkpoints(&self, execution_id: &str) -> Result<Vec<u64>> {
754 let persistence = self
755 .persistence
756 .as_ref()
757 .ok_or_else(|| WorkflowError::state("Persistence is not enabled"))?;
758
759 persistence.list_checkpoints(execution_id).await
760 }
761
762 pub async fn cleanup_checkpoints(
764 &self,
765 execution_id: &str,
766 keep_count: usize,
767 ) -> Result<usize> {
768 let persistence = self
769 .persistence
770 .as_ref()
771 .ok_or_else(|| WorkflowError::state("Persistence is not enabled"))?;
772
773 let checkpoints = persistence.list_checkpoints(execution_id).await?;
774
775 if checkpoints.len() <= keep_count {
776 return Ok(0);
777 }
778
779 let to_delete = checkpoints.len() - keep_count;
780 let mut deleted = 0;
781
782 for seq in checkpoints.iter().take(to_delete) {
783 if persistence
784 .delete_checkpoint(execution_id, *seq)
785 .await
786 .is_ok()
787 {
788 deleted += 1;
789 }
790 }
791
792 Ok(deleted)
793 }
794}
795
796#[derive(Debug, Clone)]
798pub struct RecoveryInfo {
799 pub execution_id: String,
801 pub checkpoint_sequence: u64,
803 pub checkpoint_created_at: chrono::DateTime<chrono::Utc>,
805 pub workflow_status: WorkflowStatus,
807 pub completed_tasks: Vec<String>,
809 pub pending_tasks: Vec<String>,
811 pub interrupted_tasks: Vec<String>,
813 pub failed_tasks: Vec<String>,
815 pub skipped_tasks: Vec<String>,
817 pub can_resume: bool,
819}
820
821#[cfg(test)]
822mod tests {
823 use super::*;
824 use crate::dag::graph::{ResourceRequirements, RetryPolicy};
825 use crate::engine::state::WorkflowStatus;
826 use std::collections::HashMap;
827
828 struct DummyExecutor;
829
830 #[async_trait]
831 impl TaskExecutor for DummyExecutor {
832 async fn execute(
833 &self,
834 _task: &TaskNode,
835 _context: &ExecutionContext,
836 ) -> Result<TaskOutput> {
837 Ok(TaskOutput {
838 data: Some(serde_json::json!({"result": "success"})),
839 logs: vec!["Task executed".to_string()],
840 })
841 }
842 }
843
844 fn create_test_task(id: &str) -> TaskNode {
845 TaskNode {
846 id: id.to_string(),
847 name: id.to_string(),
848 description: None,
849 config: serde_json::json!({}),
850 retry: RetryPolicy::default(),
851 timeout_secs: Some(60),
852 resources: ResourceRequirements::default(),
853 metadata: HashMap::new(),
854 }
855 }
856
857 #[tokio::test]
858 async fn test_simple_workflow() {
859 let mut dag = WorkflowDag::new();
860 dag.add_task(create_test_task("task1")).ok();
861
862 let executor = WorkflowExecutor::new(ExecutorConfig::default(), DummyExecutor);
863
864 let result = executor
865 .execute("wf1".to_string(), "exec1".to_string(), dag)
866 .await;
867
868 assert!(result.is_ok());
869 let state = result.expect("Expected workflow state");
870 assert_eq!(state.status, WorkflowStatus::Completed);
871 }
872
873 struct ConcurrencyProbe {
876 current: Arc<std::sync::atomic::AtomicUsize>,
877 max_seen: Arc<std::sync::atomic::AtomicUsize>,
878 delay_ms: u64,
879 }
880
881 #[async_trait]
882 impl TaskExecutor for ConcurrencyProbe {
883 async fn execute(
884 &self,
885 _task: &TaskNode,
886 _context: &ExecutionContext,
887 ) -> Result<TaskOutput> {
888 let now = self.current.fetch_add(1, Ordering::SeqCst) + 1;
889 self.max_seen.fetch_max(now, Ordering::SeqCst);
890 tokio::time::sleep(Duration::from_millis(self.delay_ms)).await;
891 self.current.fetch_sub(1, Ordering::SeqCst);
892 Ok(TaskOutput {
893 data: None,
894 logs: Vec::new(),
895 })
896 }
897 }
898
899 #[tokio::test]
900 async fn test_execute_level_runs_concurrently() {
901 let n = 6usize;
902 let delay_ms = 150u64;
903
904 let mut dag = WorkflowDag::new();
905 for i in 0..n {
906 dag.add_task(create_test_task(&format!("t{i}"))).ok();
907 }
908
909 let current = Arc::new(std::sync::atomic::AtomicUsize::new(0));
910 let max_seen = Arc::new(std::sync::atomic::AtomicUsize::new(0));
911 let probe = ConcurrencyProbe {
912 current: Arc::clone(¤t),
913 max_seen: Arc::clone(&max_seen),
914 delay_ms,
915 };
916
917 let config = ExecutorConfig {
918 enable_persistence: false,
919 enable_checkpointing: false,
920 max_concurrent_tasks: n,
921 retry_on_failure: false,
922 ..Default::default()
923 };
924 let executor = WorkflowExecutor::new(config, probe);
925
926 let start = std::time::Instant::now();
927 let state = executor
928 .execute("wf".to_string(), "exec".to_string(), dag)
929 .await
930 .expect("workflow should run");
931 let elapsed = start.elapsed();
932
933 assert_eq!(state.status, WorkflowStatus::Completed);
934 assert!(
937 elapsed < Duration::from_millis(delay_ms * n as u64 / 2),
938 "level did not run concurrently: elapsed {:?}",
939 elapsed
940 );
941 assert!(
942 max_seen.load(Ordering::SeqCst) >= 2,
943 "expected overlapping tasks, peak was {}",
944 max_seen.load(Ordering::SeqCst)
945 );
946 }
947
948 #[tokio::test]
949 async fn test_execute_level_respects_max_concurrent() {
950 let n = 6usize;
951
952 let mut dag = WorkflowDag::new();
953 for i in 0..n {
954 dag.add_task(create_test_task(&format!("t{i}"))).ok();
955 }
956
957 let current = Arc::new(std::sync::atomic::AtomicUsize::new(0));
958 let max_seen = Arc::new(std::sync::atomic::AtomicUsize::new(0));
959 let probe = ConcurrencyProbe {
960 current: Arc::clone(¤t),
961 max_seen: Arc::clone(&max_seen),
962 delay_ms: 80,
963 };
964
965 let config = ExecutorConfig {
966 enable_persistence: false,
967 enable_checkpointing: false,
968 max_concurrent_tasks: 2,
969 retry_on_failure: false,
970 ..Default::default()
971 };
972 let executor = WorkflowExecutor::new(config, probe);
973
974 let state = executor
975 .execute("wf".to_string(), "exec".to_string(), dag)
976 .await
977 .expect("workflow should run");
978
979 assert_eq!(state.status, WorkflowStatus::Completed);
980 let peak = max_seen.load(Ordering::SeqCst);
981 assert!(peak >= 1, "at least one task should have run");
982 assert!(
983 peak <= 2,
984 "semaphore cap exceeded: {} tasks ran concurrently",
985 peak
986 );
987 }
988}