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