ironflow_engine/context/steps/
parallel.rs1use 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::{ParallelStepResult, StepOutput, StepResult, execute_step_config};
24use crate::guard::WorkflowRejection;
25use crate::log_sender::StepLogSender;
26use crate::notify::{WorkflowAgentStepTokensUsedEvent, WorkflowEvent};
27
28impl WorkflowContext {
29 pub async fn parallel(
66 &mut self,
67 steps: Vec<(&str, StepConfig)>,
68 fail_fast: bool,
69 ) -> Result<Vec<ParallelStepResult>, EngineError> {
70 if steps.is_empty() {
71 return Ok(Vec::new());
72 }
73
74 self.check_guard_timeout()?;
76
77 let wave_budget: Decimal = steps
80 .iter()
81 .filter_map(|(_, config)| match config {
82 StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
83 _ => None,
84 })
85 .map(step_budget_usd)
86 .sum();
87 self.check_run_budget(wave_budget)?;
88
89 let wave_position = self.position;
90 self.position += 1;
91
92 let now = Utc::now();
93 let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
94 Vec::with_capacity(steps.len());
95
96 for (name, config) in &steps {
97 let kind = config.kind();
98 let trace_id = step_trace_id(self.run_id, name, wave_position);
99 let step = self
100 .store
101 .create_step(NewStep {
102 run_id: self.run_id,
103 trace_id,
104 name: name.to_string(),
105 kind,
106 position: wave_position,
107 input: Some(to_value(config)?),
108 is_error_handler: false,
109 })
110 .await?;
111
112 self.start_step(step.id, now).await?;
113
114 if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
117 self.fail_step(step.id, &err).await;
118 if !config.allow_failure() {
119 return Err(err);
120 }
121 self.has_allowed_failure = true;
122 info!(
123 run_id = %self.run_id,
124 step = %name,
125 error = %err,
126 "parallel step input preparation failed but allow_failure is set, skipping"
127 );
128 continue;
129 }
130
131 let mut config_with_trace = config.clone();
132 let step_trace = self.trace_context.child();
133 match config_with_trace {
134 StepConfig::Agent(ref mut agent_config) => {
135 agent_config.trace_context = Some(step_trace);
136 }
137 StepConfig::Http(ref mut http_config) => {
138 http_config.trace_context = Some(step_trace);
139 }
140 _ => {}
141 }
142 step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
143 }
144
145 let mut join_set = JoinSet::new();
146 let mut task_index: HashMap<Id, usize> = HashMap::new();
147 let parallel_timeout = self.guard_remaining_timeout();
148 for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
149 let provider = self.provider.clone();
150 let config = config.clone();
151 let step_log_sender = self
152 .log_sender
153 .as_ref()
154 .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
155 let handle = join_set.spawn(async move {
156 let result = match parallel_timeout {
157 Some(dur) => {
158 match timeout(
159 dur,
160 execute_step_config(&config, &provider, step_log_sender),
161 )
162 .await
163 {
164 Ok(r) => r,
165 Err(_elapsed) => {
166 Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
167 elapsed_secs: 0,
168 max: 0,
169 }))
170 }
171 }
172 }
173 None => execute_step_config(&config, &provider, step_log_sender).await,
174 };
175 (idx, result)
176 });
177 task_index.insert(handle.id(), idx);
178 }
179
180 let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
182 vec![None; step_records.len()];
183 let mut first_error: Option<EngineError> = None;
184
185 while let Some(join_result) = join_set.join_next().await {
186 let (idx, step_result) = match join_result {
187 Ok(r) => r,
188 Err(e) => {
189 let error_msg = format!("join error: {e}");
190 if let Some(&idx) = task_index.get(&e.id()) {
191 let (step_id, _, step_name, _) = &step_records[idx];
192 let completed_at = Utc::now();
193 error!(
194 run_id = %self.run_id,
195 step = %step_name,
196 error = %error_msg,
197 "parallel step panicked or was cancelled"
198 );
199 if let Err(store_err) = self
200 .store
201 .update_step(
202 *step_id,
203 StepUpdate {
204 status: Some(StepStatus::Failed),
205 error: Some(error_msg.clone()),
206 completed_at: Some(completed_at),
207 ..StepUpdate::default()
208 },
209 )
210 .await
211 {
212 error!(
213 run_id = %self.run_id,
214 step_id = %step_id,
215 error = %store_err,
216 "failed to persist JoinError for step"
217 );
218 }
219 indexed_results[idx] = Some(Err(error_msg.clone()));
220 }
221 if first_error.is_none() {
222 first_error = Some(EngineError::StepConfig(error_msg));
223 }
224 if fail_fast {
225 join_set.abort_all();
226 }
227 continue;
228 }
229 };
230
231 let (step_id, step_trace, step_name, step_config) = &step_records[idx];
232 let completed_at = Utc::now();
233
234 if let Err(err) = self
235 .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
236 .await
237 {
238 self.fail_step(*step_id, &err).await;
239 indexed_results[idx] = Some(Err(err.to_string()));
240 if first_error.is_none() {
241 first_error = Some(err);
242 }
243 if fail_fast {
244 join_set.abort_all();
245 }
246 continue;
247 }
248
249 match step_result {
250 Ok(output) => {
251 self.total_cost_usd += output.cost_usd;
252 self.total_duration_ms += output.duration_ms;
253
254 if matches!(step_config, StepConfig::Agent(_)) {
256 let tokens = output
257 .input_tokens
258 .unwrap_or(0)
259 .saturating_add(output.output_tokens.unwrap_or(0));
260 if tokens > 0
261 && let Err(guard_err) = self.guard_record_tokens(tokens)
262 {
263 if first_error.is_none() {
264 first_error = Some(guard_err);
265 }
266 if fail_fast {
267 join_set.abort_all();
268 }
269 }
270 }
271
272 let debug_messages_json = output.debug_messages_json();
273
274 self.store
275 .update_step(
276 *step_id,
277 StepUpdate {
278 status: Some(StepStatus::Completed),
279 output: Some(output.output.clone()),
280 duration_ms: Some(output.duration_ms),
281 cost_usd: Some(output.cost_usd),
282 input_tokens: output.input_tokens,
283 output_tokens: output.output_tokens,
284 completed_at: Some(completed_at),
285 debug_messages: debug_messages_json,
286 ..StepUpdate::default()
287 },
288 )
289 .await?;
290
291 self.step_results.push(StepResult::from_success(
292 *step_trace,
293 step_name,
294 &output,
295 ));
296
297 if let Some(ref bus) = self.event_bus
298 && matches!(step_config, StepConfig::Agent(_))
299 {
300 let tokens = output
301 .input_tokens
302 .unwrap_or(0)
303 .saturating_add(output.output_tokens.unwrap_or(0));
304 bus.publish(
305 self.run_id,
306 WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
307 step_name: step_name.clone(),
308 tokens,
309 cost_usd: output.cost_usd,
310 }),
311 );
312 }
313
314 info!(
315 run_id = %self.run_id,
316 step = %step_name,
317 trace_id = %step_trace,
318 duration_ms = output.duration_ms,
319 "parallel step completed"
320 );
321
322 indexed_results[idx] = Some(Ok(output));
323 }
324 Err(err) => {
325 let err_msg = err.to_string();
326 let debug_messages_json = extract_debug_messages_from_error(&err);
327 let partial = extract_partial_usage_from_error(&err);
328 let raw_response_output = extract_raw_response_from_error(&err);
329
330 if let Some(ref usage) = partial {
331 if let Some(cost) = usage.cost_usd {
332 self.total_cost_usd += cost;
333 }
334 if let Some(dur) = usage.duration_ms {
335 self.total_duration_ms += dur;
336 }
337 }
338
339 if let Err(store_err) = self
340 .store
341 .update_step(
342 *step_id,
343 StepUpdate {
344 status: Some(StepStatus::Failed),
345 error: Some(err_msg.clone()),
346 output: raw_response_output.clone(),
347 completed_at: Some(completed_at),
348 debug_messages: debug_messages_json,
349 duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
350 cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
351 input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
352 output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
353 ..StepUpdate::default()
354 },
355 )
356 .await
357 {
358 error!(
359 step_id = %step_id,
360 error = %store_err,
361 "failed to persist parallel step failure"
362 );
363 }
364
365 let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
366 let err_cost = partial
367 .as_ref()
368 .and_then(|p| p.cost_usd)
369 .unwrap_or(Decimal::ZERO);
370 self.step_results.push(StepResult::from_failure(
371 *step_trace,
372 step_name,
373 &err_msg,
374 err_duration,
375 err_cost,
376 ));
377
378 if step_config.allow_failure() {
379 self.has_allowed_failure = true;
380 info!(
381 run_id = %self.run_id,
382 step = %step_name,
383 error = %err_msg,
384 "parallel step failed but allow_failure is set, continuing"
385 );
386 indexed_results[idx] = Some(Ok(allowed_failure_output(
387 &err_msg,
388 raw_response_output,
389 partial.as_ref(),
390 )));
391 } else {
392 indexed_results[idx] = Some(Err(err_msg.clone()));
393
394 if first_error.is_none() {
395 first_error = Some(err);
396 }
397
398 if fail_fast {
399 join_set.abort_all();
400 }
401 }
402 }
403 }
404 }
405
406 if let Some(err) = first_error {
407 return Err(err);
408 }
409
410 self.persist_progress().await;
411
412 self.last_step_ids = step_records.iter().map(|(id, _, _, _)| *id).collect();
413
414 let results: Vec<ParallelStepResult> = step_records
416 .iter()
417 .enumerate()
418 .map(|(idx, (step_id, _trace_id, name, _))| {
419 let output = match indexed_results[idx].take() {
420 Some(Ok(o)) => o,
421 _ => unreachable!("all steps succeeded if no error returned"),
422 };
423 ParallelStepResult {
424 name: name.clone(),
425 output,
426 step_id: *step_id,
427 }
428 })
429 .collect();
430
431 Ok(results)
432 }
433}