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