Skip to main content

oxigdal_workflow/engine/
executor.rs

1//! Workflow execution engine.
2
3use 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/// Task executor trait - implement this to define custom task execution logic.
17#[async_trait]
18pub trait TaskExecutor: Send + Sync {
19    /// Execute a task.
20    async fn execute(&self, task: &TaskNode, context: &ExecutionContext) -> Result<TaskOutput>;
21}
22
23/// Execution context provided to task executors.
24#[derive(Debug, Clone)]
25pub struct ExecutionContext {
26    /// Workflow execution ID.
27    pub execution_id: String,
28    /// Task ID.
29    pub task_id: String,
30    /// Shared workflow state.
31    pub state: Arc<RwLock<WorkflowState>>,
32    /// Input data from previous tasks.
33    pub inputs: std::collections::HashMap<String, serde_json::Value>,
34}
35
36/// Task execution output.
37#[derive(Debug, Clone)]
38pub struct TaskOutput {
39    /// Output data.
40    pub data: Option<serde_json::Value>,
41    /// Execution logs.
42    pub logs: Vec<String>,
43}
44
45/// Workflow executor configuration.
46#[derive(Debug, Clone)]
47pub struct ExecutorConfig {
48    /// Maximum concurrent tasks.
49    pub max_concurrent_tasks: usize,
50    /// Enable state persistence.
51    pub enable_persistence: bool,
52    /// State directory.
53    pub state_dir: String,
54    /// Resource pool.
55    pub resource_pool: ResourcePool,
56    /// Retry on failure.
57    pub retry_on_failure: bool,
58    /// Stop on first failure.
59    pub stop_on_failure: bool,
60    /// Checkpoint interval (save checkpoint every N tasks).
61    pub checkpoint_interval: usize,
62    /// Enable checkpoint-based recovery.
63    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, // Save after each task by default
79            enable_checkpointing: true,
80        }
81    }
82}
83
84/// Workflow executor.
85pub struct WorkflowExecutor<E: TaskExecutor> {
86    /// Configuration.
87    config: ExecutorConfig,
88    /// Task executor implementation.
89    task_executor: Arc<E>,
90    /// State persistence.
91    persistence: Option<StatePersistence>,
92    /// Semaphore bounding the number of tasks executing concurrently within a level.
93    semaphore: Arc<Semaphore>,
94    /// Checkpoint sequence counter.
95    checkpoint_sequence: AtomicU64,
96    /// Tasks completed since last checkpoint.
97    tasks_since_checkpoint: AtomicU64,
98}
99
100impl<E: TaskExecutor> WorkflowExecutor<E> {
101    /// Create a new workflow executor.
102    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    /// Save a checkpoint if conditions are met.
121    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    /// Force save a checkpoint immediately.
150    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    /// Execute a workflow.
174    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        // Validate DAG
186        dag.validate()?;
187
188        // Create initial workflow state
189        let mut state = WorkflowState::new(workflow_id.clone(), execution_id.clone(), workflow_id);
190
191        // Initialize task states
192        for task in dag.tasks() {
193            state.init_task(task.id.clone());
194        }
195
196        state.start();
197
198        // Save initial state
199        if let Some(ref persistence) = self.persistence {
200            persistence.save(&state).await?;
201        }
202
203        // Save initial checkpoint with DAG
204        self.save_checkpoint_now(&state, &dag).await?;
205
206        let state_arc = Arc::new(RwLock::new(state));
207
208        // Create execution plan
209        let execution_plan = create_execution_plan(&dag)?;
210
211        info!(
212            "Execution plan created with {} levels",
213            execution_plan.len()
214        );
215
216        // Execute tasks level by level
217        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            // Save checkpoint after each level
223            {
224                let state_guard = state_arc.read().await;
225                self.maybe_save_checkpoint(&state_guard, &dag).await?;
226            }
227
228            // Check for failures
229            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                    // Save final checkpoint on failure
253                    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        // Complete workflow
267        let mut state_guard = state_arc.write().await;
268
269        // Check if all tasks completed successfully
270        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        // Save final state
282        if let Some(ref persistence) = self.persistence {
283            persistence.save(&state_guard).await?;
284        }
285
286        // Save final checkpoint
287        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    /// Execute a level of tasks concurrently.
302    ///
303    /// All tasks in a level are independent (they share no dependency edges), so
304    /// they are driven concurrently as a set of futures, each acquiring a
305    /// permit from the executor semaphore first so that at most
306    /// `max_concurrent_tasks` run at any instant. This turns the wall-clock cost
307    /// of a level from the sum of its task durations into roughly the slowest
308    /// task (subject to the concurrency cap), which matters for I/O-bound work.
309    ///
310    /// Futures borrow `&self`, `dag`, and `state`, so they are polled on the
311    /// current task rather than `tokio::spawn`ed — this avoids `'static`/`Send`
312    /// bounds on the caller-supplied `TaskExecutor` while still interleaving all
313    /// tasks at their `.await` points.
314    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                // Bound concurrency. `acquire_owned` only errors if the semaphore
328                // is closed, which never happens here; fall back to running
329                // without a permit rather than panicking.
330                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    /// Execute a single task.
353    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        // Check dependencies
368        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        // Mark task as started
376        {
377            let mut state_guard = state.write().await;
378            state_guard.start_task(task_id)?;
379        }
380
381        // Execute with retry
382        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            // Gather inputs from dependencies
402            let inputs = self.gather_inputs(task_id, dag, state).await?;
403
404            // Create execution context
405            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            // Execute task with timeout
416            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                    // Task succeeded
422                    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        // All attempts failed
446        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    /// Check if all task dependencies are met.
455    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    /// Gather inputs from task dependencies.
478    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    /// Resume a workflow from a saved checkpoint.
500    ///
501    /// This method reconstructs the DAG from the checkpoint and continues
502    /// execution from where it left off, handling:
503    /// - Completed tasks (skipped)
504    /// - Interrupted tasks (reset to pending for retry)
505    /// - Failed tasks (depending on configuration)
506    /// - Pending tasks (executed normally)
507    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        // Try to load checkpoint first (includes DAG)
514        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        // Prepare checkpoint for resumption (reset interrupted tasks)
531        checkpoint.prepare_for_resume()?;
532
533        // Update checkpoint sequence for this executor
534        self.checkpoint_sequence
535            .store(checkpoint.sequence + 1, Ordering::SeqCst);
536
537        // Resume execution with recovered DAG and state
538        self.resume_from_checkpoint(checkpoint).await
539    }
540
541    /// Resume execution from a checkpoint.
542    async fn resume_from_checkpoint(
543        &self,
544        checkpoint: WorkflowCheckpoint,
545    ) -> Result<WorkflowState> {
546        let dag = checkpoint.dag.clone();
547
548        // Log recovery information (before moving state)
549        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        // Ensure workflow is in running state
557        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        // Save state update
570        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        // Create execution plan from DAG
577        let execution_plan = create_execution_plan(&dag)?;
578
579        info!("Resuming execution with {} levels", execution_plan.len());
580
581        // Execute tasks level by level, skipping completed ones
582        for (level_idx, level) in execution_plan.iter().enumerate() {
583            // Filter out already completed or skipped tasks
584            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            // Save checkpoint after each level
617            {
618                let state_guard = state_arc.read().await;
619                self.maybe_save_checkpoint(&state_guard, &dag).await?;
620            }
621
622            // Check for failures
623            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                    // Save final checkpoint on failure
647                    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        // Complete workflow
661        let mut state_guard = state_arc.write().await;
662
663        // Check if all tasks completed successfully
664        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        // Save final state
676        if let Some(ref persistence) = self.persistence {
677            persistence.save(&state_guard).await?;
678        }
679
680        // Save final checkpoint
681        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    /// Resume a workflow from a specific checkpoint sequence.
696    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        // Prepare checkpoint for resumption
720        checkpoint.prepare_for_resume()?;
721
722        // Update checkpoint sequence
723        self.checkpoint_sequence
724            .store(sequence + 1, Ordering::SeqCst);
725
726        self.resume_from_checkpoint(checkpoint).await
727    }
728
729    /// Get recovery information for an execution.
730    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    /// List available checkpoints for an execution.
753    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    /// Clean up old checkpoints, keeping only the latest N.
763    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/// Information about workflow recovery state.
797#[derive(Debug, Clone)]
798pub struct RecoveryInfo {
799    /// Execution ID.
800    pub execution_id: String,
801    /// Latest checkpoint sequence number.
802    pub checkpoint_sequence: u64,
803    /// When the checkpoint was created.
804    pub checkpoint_created_at: chrono::DateTime<chrono::Utc>,
805    /// Current workflow status.
806    pub workflow_status: WorkflowStatus,
807    /// Tasks that completed successfully.
808    pub completed_tasks: Vec<String>,
809    /// Tasks that are pending execution.
810    pub pending_tasks: Vec<String>,
811    /// Tasks that were interrupted (running when checkpoint saved).
812    pub interrupted_tasks: Vec<String>,
813    /// Tasks that failed.
814    pub failed_tasks: Vec<String>,
815    /// Tasks that were skipped.
816    pub skipped_tasks: Vec<String>,
817    /// Whether the workflow can be resumed.
818    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    /// Task executor that records the peak number of tasks running at once and
874    /// sleeps a fixed duration, used to observe real concurrency.
875    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(&current),
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        // Serial execution would take ~n * delay_ms; concurrent execution should
935        // finish in well under half that.
936        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(&current),
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}