Skip to main content

ironflow_engine/context/steps/
parallel.rs

1//! Parallel step wave for [`WorkflowContext`].
2
3use std::collections::HashMap;
4
5use chrono::Utc;
6use rust_decimal::Decimal;
7use serde_json::to_value;
8use tokio::task::{Id, JoinSet};
9use tokio::time::timeout;
10use tracing::{error, info};
11use uuid::Uuid;
12
13use ironflow_store::models::{NewStep, StepStatus, StepUpdate, step_trace_id};
14
15use crate::budget::step_budget_usd;
16use crate::config::StepConfig;
17use crate::context::WorkflowContext;
18use crate::context::failure::{
19    allowed_failure_output, extract_debug_messages_from_error, extract_partial_usage_from_error,
20    extract_raw_response_from_error,
21};
22use crate::error::EngineError;
23use crate::executor::{
24    ParallelStepResult, StepOutput, StepResult, execute_step_config_intercepted,
25};
26use crate::guard::WorkflowRejection;
27use crate::log_sender::StepLogSender;
28use crate::notify::{WorkflowAgentStepTokensUsedEvent, WorkflowEvent};
29use crate::plan::{lock_plan, planned_output};
30
31impl WorkflowContext {
32    /// Execute multiple steps concurrently (wait-all model).
33    ///
34    /// All steps in the batch execute in parallel via `tokio::JoinSet`.
35    /// Each step is recorded with the same `position` (execution wave).
36    /// Dependencies on previous steps are recorded automatically.
37    ///
38    /// When `fail_fast` is true, remaining steps are aborted on the first
39    /// failure. When false, all steps run to completion and the first
40    /// error is returned.
41    ///
42    /// # Errors
43    ///
44    /// Returns [`EngineError`] if any step fails.
45    ///
46    /// # Examples
47    ///
48    /// ```no_run
49    /// use ironflow_engine::context::WorkflowContext;
50    /// use ironflow_engine::config::{StepConfig, ShellConfig};
51    /// use ironflow_engine::error::EngineError;
52    ///
53    /// # async fn example(ctx: &mut WorkflowContext) -> Result<(), EngineError> {
54    /// let results = ctx.parallel(
55    ///     vec![
56    ///         ("test-unit", StepConfig::Shell(ShellConfig::new("cargo test --lib"))),
57    ///         ("lint", StepConfig::Shell(ShellConfig::new("cargo clippy"))),
58    ///     ],
59    ///     true,
60    /// ).await?;
61    ///
62    /// for r in &results {
63    ///     println!("{}: {:?}", r.name, r.output.output);
64    /// }
65    /// # Ok(())
66    /// # }
67    /// ```
68    pub async fn parallel(
69        &mut self,
70        steps: Vec<(&str, StepConfig)>,
71        fail_fast: bool,
72    ) -> Result<Vec<ParallelStepResult>, EngineError> {
73        if steps.is_empty() {
74            return Ok(Vec::new());
75        }
76
77        // Plan mode: record the whole wave under one parallel group and return
78        // synthetic outputs. No step record is created and nothing runs.
79        if let Some(plan) = self.plan().cloned() {
80            self.position += 1;
81            let mut results = Vec::with_capacity(steps.len());
82            let mut names = Vec::with_capacity(steps.len());
83            {
84                let mut recorder = lock_plan(&plan);
85                let group = recorder.next_group();
86                for (name, config) in &steps {
87                    let wave = Some(group.clone());
88                    if !recorder.record(name, config.kind(), &self.workflow_name, wave) {
89                        break;
90                    }
91                    names.push((*name).to_string());
92                    let estimate = recorder.estimate_for(name);
93                    results.push(ParallelStepResult {
94                        name: (*name).to_string(),
95                        output: planned_output(config, estimate),
96                        step_id: Uuid::now_v7(),
97                    });
98                }
99                recorder.set_last(names);
100            }
101            return Ok(results);
102        }
103
104        // Guard timeout: checked before launching the wave.
105        self.check_guard_timeout()?;
106
107        // Cost cap: the whole wave is charged at once. Refused before any step
108        // record is created, so nothing in the wave starts.
109        let wave_budget: Decimal = steps
110            .iter()
111            .filter_map(|(_, config)| match config {
112                StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
113                _ => None,
114            })
115            .map(step_budget_usd)
116            .sum();
117        self.check_run_budget(wave_budget)?;
118
119        let wave_position = self.position;
120        self.position += 1;
121
122        let now = Utc::now();
123        let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
124            Vec::with_capacity(steps.len());
125
126        for (name, config) in &steps {
127            let kind = config.kind();
128            let trace_id = step_trace_id(self.run_id, name, wave_position);
129            let step = self
130                .store
131                .create_step(NewStep {
132                    run_id: self.run_id,
133                    trace_id,
134                    name: name.to_string(),
135                    kind,
136                    position: wave_position,
137                    input: Some(to_value(config)?),
138                    is_error_handler: false,
139                })
140                .await?;
141
142            self.start_step(step.id, now).await?;
143
144            // Inputs are materialized before any step in the wave starts, so a
145            // missing one fails the wave rather than a half-run command.
146            if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
147                self.fail_step(step.id, &err).await;
148                if !config.allow_failure() {
149                    return Err(err);
150                }
151                self.has_allowed_failure = true;
152                info!(
153                    run_id = %self.run_id,
154                    step = %name,
155                    error = %err,
156                    "parallel step input preparation failed but allow_failure is set, skipping"
157                );
158                continue;
159            }
160
161            let mut config_with_trace = config.clone();
162            let step_trace = self.trace_context.child();
163            match config_with_trace {
164                StepConfig::Agent(ref mut agent_config) => {
165                    agent_config.trace_context = Some(step_trace);
166                }
167                StepConfig::Http(ref mut http_config) => {
168                    http_config.trace_context = Some(step_trace);
169                }
170                _ => {}
171            }
172            step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
173        }
174
175        let mut join_set = JoinSet::new();
176        let mut task_index: HashMap<Id, usize> = HashMap::new();
177        let parallel_timeout = self.guard_remaining_timeout();
178        for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
179            let provider = self.provider.clone();
180            // Each task owns its own handle: `intercept` is synchronous, so no
181            // borrow of the context is held across an await point.
182            let interceptor = self.interceptor.clone();
183            let config = config.clone();
184            let step_log_sender = self
185                .log_sender
186                .as_ref()
187                .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
188            let handle = join_set.spawn(async move {
189                let result = match parallel_timeout {
190                    Some(dur) => {
191                        match timeout(
192                            dur,
193                            execute_step_config_intercepted(
194                                &config,
195                                &provider,
196                                step_log_sender,
197                                interceptor.as_ref(),
198                            ),
199                        )
200                        .await
201                        {
202                            Ok(r) => r,
203                            Err(_elapsed) => {
204                                Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
205                                    elapsed_secs: 0,
206                                    max: 0,
207                                }))
208                            }
209                        }
210                    }
211                    None => {
212                        execute_step_config_intercepted(
213                            &config,
214                            &provider,
215                            step_log_sender,
216                            interceptor.as_ref(),
217                        )
218                        .await
219                    }
220                };
221                (idx, result)
222            });
223            task_index.insert(handle.id(), idx);
224        }
225
226        // JoinSet returns in completion order; indexed_results restores input order.
227        let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
228            vec![None; step_records.len()];
229        let mut first_error: Option<EngineError> = None;
230
231        while let Some(join_result) = join_set.join_next().await {
232            let (idx, step_result) = match join_result {
233                Ok(r) => r,
234                Err(e) => {
235                    let error_msg = format!("join error: {e}");
236                    if let Some(&idx) = task_index.get(&e.id()) {
237                        let (step_id, _, step_name, _) = &step_records[idx];
238                        let completed_at = Utc::now();
239                        error!(
240                            run_id = %self.run_id,
241                            step = %step_name,
242                            error = %error_msg,
243                            "parallel step panicked or was cancelled"
244                        );
245                        if let Err(store_err) = self
246                            .store
247                            .update_step(
248                                *step_id,
249                                StepUpdate {
250                                    status: Some(StepStatus::Failed),
251                                    error: Some(error_msg.clone()),
252                                    completed_at: Some(completed_at),
253                                    ..StepUpdate::default()
254                                },
255                            )
256                            .await
257                        {
258                            error!(
259                                run_id = %self.run_id,
260                                step_id = %step_id,
261                                error = %store_err,
262                                "failed to persist JoinError for step"
263                            );
264                        }
265                        indexed_results[idx] = Some(Err(error_msg.clone()));
266                    }
267                    if first_error.is_none() {
268                        first_error = Some(EngineError::StepConfig(error_msg));
269                    }
270                    if fail_fast {
271                        join_set.abort_all();
272                    }
273                    continue;
274                }
275            };
276
277            let (step_id, step_trace, step_name, step_config) = &step_records[idx];
278            let completed_at = Utc::now();
279
280            if let Err(err) = self
281                .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
282                .await
283            {
284                self.fail_step(*step_id, &err).await;
285                indexed_results[idx] = Some(Err(err.to_string()));
286                if first_error.is_none() {
287                    first_error = Some(err);
288                }
289                if fail_fast {
290                    join_set.abort_all();
291                }
292                continue;
293            }
294
295            match step_result {
296                Ok(output) => {
297                    self.total_cost_usd += output.cost_usd;
298                    self.total_duration_ms += output.duration_ms;
299
300                    // Record token usage in the guard for agent steps.
301                    if matches!(step_config, StepConfig::Agent(_)) {
302                        let tokens = output
303                            .input_tokens
304                            .unwrap_or(0)
305                            .saturating_add(output.output_tokens.unwrap_or(0));
306                        if tokens > 0
307                            && let Err(guard_err) = self.guard_record_tokens(tokens)
308                        {
309                            if first_error.is_none() {
310                                first_error = Some(guard_err);
311                            }
312                            if fail_fast {
313                                join_set.abort_all();
314                            }
315                        }
316                    }
317
318                    let debug_messages_json = output.debug_messages_json();
319
320                    self.store
321                        .update_step(
322                            *step_id,
323                            StepUpdate {
324                                status: Some(StepStatus::Completed),
325                                output: Some(output.output.clone()),
326                                duration_ms: Some(output.duration_ms),
327                                cost_usd: Some(output.cost_usd),
328                                input_tokens: output.input_tokens,
329                                output_tokens: output.output_tokens,
330                                completed_at: Some(completed_at),
331                                debug_messages: debug_messages_json,
332                                ..StepUpdate::default()
333                            },
334                        )
335                        .await?;
336
337                    self.step_results.push(StepResult::from_success(
338                        *step_trace,
339                        step_name,
340                        &output,
341                    ));
342
343                    if let Some(ref bus) = self.event_bus
344                        && matches!(step_config, StepConfig::Agent(_))
345                    {
346                        let tokens = output
347                            .input_tokens
348                            .unwrap_or(0)
349                            .saturating_add(output.output_tokens.unwrap_or(0));
350                        bus.publish(
351                            self.run_id,
352                            WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
353                                step_name: step_name.clone(),
354                                tokens,
355                                cost_usd: output.cost_usd,
356                            }),
357                        );
358                    }
359
360                    info!(
361                        run_id = %self.run_id,
362                        step = %step_name,
363                        trace_id = %step_trace,
364                        duration_ms = output.duration_ms,
365                        "parallel step completed"
366                    );
367
368                    indexed_results[idx] = Some(Ok(output));
369                }
370                Err(err) => {
371                    let err_msg = err.to_string();
372                    let debug_messages_json = extract_debug_messages_from_error(&err);
373                    let partial = extract_partial_usage_from_error(&err);
374                    let raw_response_output = extract_raw_response_from_error(&err);
375
376                    if let Some(ref usage) = partial {
377                        if let Some(cost) = usage.cost_usd {
378                            self.total_cost_usd += cost;
379                        }
380                        if let Some(dur) = usage.duration_ms {
381                            self.total_duration_ms += dur;
382                        }
383                    }
384
385                    if let Err(store_err) = self
386                        .store
387                        .update_step(
388                            *step_id,
389                            StepUpdate {
390                                status: Some(StepStatus::Failed),
391                                error: Some(err_msg.clone()),
392                                output: raw_response_output.clone(),
393                                completed_at: Some(completed_at),
394                                debug_messages: debug_messages_json,
395                                duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
396                                cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
397                                input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
398                                output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
399                                ..StepUpdate::default()
400                            },
401                        )
402                        .await
403                    {
404                        error!(
405                            step_id = %step_id,
406                            error = %store_err,
407                            "failed to persist parallel step failure"
408                        );
409                    }
410
411                    let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
412                    let err_cost = partial
413                        .as_ref()
414                        .and_then(|p| p.cost_usd)
415                        .unwrap_or(Decimal::ZERO);
416                    self.step_results.push(StepResult::from_failure(
417                        *step_trace,
418                        step_name,
419                        &err_msg,
420                        err_duration,
421                        err_cost,
422                    ));
423
424                    if step_config.allow_failure() {
425                        self.has_allowed_failure = true;
426                        info!(
427                            run_id = %self.run_id,
428                            step = %step_name,
429                            error = %err_msg,
430                            "parallel step failed but allow_failure is set, continuing"
431                        );
432                        indexed_results[idx] = Some(Ok(allowed_failure_output(
433                            &err_msg,
434                            raw_response_output,
435                            partial.as_ref(),
436                        )));
437                    } else {
438                        indexed_results[idx] = Some(Err(err_msg.clone()));
439
440                        if first_error.is_none() {
441                            first_error = Some(err);
442                        }
443
444                        if fail_fast {
445                            join_set.abort_all();
446                        }
447                    }
448                }
449            }
450        }
451
452        if let Some(err) = first_error {
453            return Err(err);
454        }
455
456        self.persist_progress().await;
457
458        self.last_step_ids = step_records.iter().map(|(id, _, _, _)| *id).collect();
459
460        // Build results in original order.
461        let results: Vec<ParallelStepResult> = step_records
462            .iter()
463            .enumerate()
464            .map(|(idx, (step_id, _trace_id, name, _))| {
465                let output = match indexed_results[idx].take() {
466                    Some(Ok(o)) => o,
467                    _ => unreachable!("all steps succeeded if no error returned"),
468                };
469                ParallelStepResult {
470                    name: name.clone(),
471                    output,
472                    step_id: *step_id,
473                }
474            })
475            .collect();
476
477        Ok(results)
478    }
479}