1use std::collections::{HashMap, HashSet};
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(
77 &mut self,
78 steps: Vec<(&str, StepConfig)>,
79 fail_fast: bool,
80 ) -> Result<Vec<ParallelStepResult>, EngineError> {
81 if steps.is_empty() {
82 return Ok(Vec::new());
83 }
84 reject_duplicate_names(&steps)?;
85
86 if let Some(plan) = self.plan().cloned() {
89 self.position += 1;
90 let mut results = Vec::with_capacity(steps.len());
91 let mut names = Vec::with_capacity(steps.len());
92 {
93 let mut recorder = lock_plan(&plan);
94 let group = recorder.next_group();
95 for (name, config) in &steps {
96 let wave = Some(group.clone());
97 if !recorder.record(name, config.kind(), &self.workflow_name, wave) {
98 break;
99 }
100 names.push((*name).to_string());
101 let estimate = recorder.estimate_for(name);
102 let mut output = planned_output(config, estimate);
103 output.artifacts = StepArtifacts::new(name, None, config.declared_outputs());
104 results.push(ParallelStepResult {
105 name: (*name).to_string(),
106 output,
107 step_id: Uuid::now_v7(),
108 });
109 }
110 recorder.set_last(names);
111 }
112 return Ok(results);
113 }
114
115 self.check_guard_timeout()?;
117
118 let wave_position = self.position;
119 self.position += 1;
120
121 let mut slots = self.replay_wave(wave_position, &steps);
127 if slots.iter().all(Option::is_some) {
128 let results: Vec<ParallelStepResult> = slots.into_iter().flatten().collect();
129 self.last_step_ids = results.iter().map(|r| r.step_id).collect();
130 return Ok(results);
131 }
132
133 let wave_budget: Decimal = steps
136 .iter()
137 .zip(&slots)
138 .filter(|(_, slot)| slot.is_none())
139 .filter_map(|((_, config), _)| match config {
140 StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
141 _ => None,
142 })
143 .map(step_budget_usd)
144 .sum();
145 self.check_run_budget(wave_budget)?;
146
147 let now = Utc::now();
148 let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
149 Vec::with_capacity(steps.len());
150 let mut record_slots: Vec<usize> = Vec::with_capacity(steps.len());
152
153 for (slot, (name, config)) in steps.iter().enumerate() {
154 if slots[slot].is_some() {
155 continue;
156 }
157 let kind = config.kind();
158 let trace_id = step_trace_id(self.run_id, name, wave_position);
159 let step = self
160 .store
161 .create_step(NewStep {
162 run_id: self.run_id,
163 trace_id,
164 name: name.to_string(),
165 kind,
166 position: wave_position,
167 input: Some(to_value(config)?),
168 is_error_handler: false,
169 })
170 .await?;
171
172 self.start_step(step.id, now).await?;
173
174 if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
177 self.fail_step(step.id, &err).await;
178 if !config.allow_failure() {
179 return Err(err);
180 }
181 self.has_allowed_failure = true;
182 info!(
183 run_id = %self.run_id,
184 step = %name,
185 error = %err,
186 "parallel step input preparation failed but allow_failure is set, skipping"
187 );
188 continue;
189 }
190
191 let mut config_with_trace = config.clone();
192 let step_trace = self.trace_context.child();
193 match config_with_trace {
194 StepConfig::Agent(ref mut agent_config) => {
195 agent_config.trace_context = Some(step_trace);
196 }
197 StepConfig::Http(ref mut http_config) => {
198 http_config.trace_context = Some(step_trace);
199 }
200 _ => {}
201 }
202 step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
203 record_slots.push(slot);
204 }
205
206 let mut join_set = JoinSet::new();
207 let mut task_index: HashMap<Id, usize> = HashMap::new();
208 let parallel_timeout = self.guard_remaining_timeout();
209 for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
210 let provider = self.provider.clone();
211 let interceptor = self.interceptor.clone();
214 let config = config.clone();
215 let step_log_sender = self
216 .log_sender
217 .as_ref()
218 .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
219 let handle = join_set.spawn(async move {
220 let result = match parallel_timeout {
221 Some(dur) => {
222 match timeout(
223 dur,
224 execute_step_config_intercepted(
225 &config,
226 &provider,
227 step_log_sender,
228 interceptor.as_ref(),
229 ),
230 )
231 .await
232 {
233 Ok(r) => r,
234 Err(_elapsed) => {
235 Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
236 elapsed_secs: 0,
237 max: 0,
238 }))
239 }
240 }
241 }
242 None => {
243 execute_step_config_intercepted(
244 &config,
245 &provider,
246 step_log_sender,
247 interceptor.as_ref(),
248 )
249 .await
250 }
251 };
252 (idx, result)
253 });
254 task_index.insert(handle.id(), idx);
255 }
256
257 let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
259 vec![None; step_records.len()];
260 let mut first_error: Option<EngineError> = None;
261
262 while let Some(join_result) = join_set.join_next().await {
263 let (idx, step_result) = match join_result {
264 Ok(r) => r,
265 Err(e) => {
266 let error_msg = format!("join error: {e}");
267 if let Some(&idx) = task_index.get(&e.id()) {
268 let (step_id, _, step_name, _) = &step_records[idx];
269 let completed_at = Utc::now();
270 error!(
271 run_id = %self.run_id,
272 step = %step_name,
273 error = %error_msg,
274 "parallel step panicked or was cancelled"
275 );
276 if let Err(store_err) = self
277 .store
278 .update_step(
279 *step_id,
280 StepUpdate {
281 status: Some(StepStatus::Failed),
282 error: Some(error_msg.clone()),
283 completed_at: Some(completed_at),
284 ..StepUpdate::default()
285 },
286 )
287 .await
288 {
289 error!(
290 run_id = %self.run_id,
291 step_id = %step_id,
292 error = %store_err,
293 "failed to persist JoinError for step"
294 );
295 }
296 indexed_results[idx] = Some(Err(error_msg.clone()));
297 }
298 if first_error.is_none() {
299 first_error = Some(EngineError::StepConfig(error_msg));
300 }
301 if fail_fast {
302 join_set.abort_all();
303 }
304 continue;
305 }
306 };
307
308 let (step_id, step_trace, step_name, step_config) = &step_records[idx];
309 let completed_at = Utc::now();
310
311 if let Err(err) = self
312 .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
313 .await
314 {
315 self.fail_step(*step_id, &err).await;
316 indexed_results[idx] = Some(Err(err.to_string()));
317 if first_error.is_none() {
318 first_error = Some(err);
319 }
320 if fail_fast {
321 join_set.abort_all();
322 }
323 continue;
324 }
325
326 match step_result {
327 Ok(output) => {
328 self.total_cost_usd += output.cost_usd;
329 self.total_duration_ms += output.duration_ms;
330
331 if matches!(step_config, StepConfig::Agent(_)) {
333 let tokens = output.total_tokens();
334 if tokens > 0
335 && let Err(guard_err) = self.guard_record_tokens(tokens)
336 {
337 if first_error.is_none() {
338 first_error = Some(guard_err);
339 }
340 if fail_fast {
341 join_set.abort_all();
342 }
343 }
344 }
345
346 let debug_messages_json = output.debug_messages_json();
347
348 self.store
349 .update_step(
350 *step_id,
351 StepUpdate {
352 status: Some(StepStatus::Completed),
353 output: Some(output.output.clone()),
354 duration_ms: Some(output.duration_ms),
355 cost_usd: Some(output.cost_usd),
356 input_tokens: output.input_tokens,
357 cache_read_input_tokens: output.cache_read_input_tokens,
358 cache_creation_input_tokens: output.cache_creation_input_tokens,
359 output_tokens: output.output_tokens,
360 completed_at: Some(completed_at),
361 debug_messages: debug_messages_json,
362 ..StepUpdate::default()
363 },
364 )
365 .await?;
366
367 self.step_results.push(StepResult::from_success(
368 *step_trace,
369 step_name,
370 &output,
371 ));
372
373 if let Some(ref bus) = self.event_bus
374 && matches!(step_config, StepConfig::Agent(_))
375 {
376 let tokens = output.total_tokens();
377 bus.publish(
378 self.run_id,
379 WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
380 step_name: step_name.clone(),
381 tokens,
382 cost_usd: output.cost_usd,
383 }),
384 );
385 }
386
387 info!(
388 run_id = %self.run_id,
389 step = %step_name,
390 trace_id = %step_trace,
391 duration_ms = output.duration_ms,
392 "parallel step completed"
393 );
394
395 indexed_results[idx] = Some(Ok(output));
396 }
397 Err(err) => {
398 let err_msg = err.to_string();
399 let debug_messages_json = extract_debug_messages_from_error(&err);
400 let partial = extract_partial_usage_from_error(&err);
401 let raw_response_output = extract_raw_response_from_error(&err);
402
403 if let Some(ref usage) = partial {
404 if let Some(cost) = usage.cost_usd {
405 self.total_cost_usd += cost;
406 }
407 if let Some(dur) = usage.duration_ms {
408 self.total_duration_ms += dur;
409 }
410 }
411
412 if let Err(store_err) = self
413 .store
414 .update_step(
415 *step_id,
416 StepUpdate {
417 status: Some(StepStatus::Failed),
418 error: Some(err_msg.clone()),
419 output: raw_response_output.clone(),
420 completed_at: Some(completed_at),
421 debug_messages: debug_messages_json,
422 duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
423 cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
424 input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
425 cache_read_input_tokens: partial
426 .as_ref()
427 .and_then(|p| p.cache_read_input_tokens),
428 cache_creation_input_tokens: partial
429 .as_ref()
430 .and_then(|p| p.cache_creation_input_tokens),
431 output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
432 ..StepUpdate::default()
433 },
434 )
435 .await
436 {
437 error!(
438 step_id = %step_id,
439 error = %store_err,
440 "failed to persist parallel step failure"
441 );
442 }
443
444 let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
445 let err_cost = partial
446 .as_ref()
447 .and_then(|p| p.cost_usd)
448 .unwrap_or(Decimal::ZERO);
449 self.step_results.push(StepResult::from_failure(
450 *step_trace,
451 step_name,
452 &err_msg,
453 err_duration,
454 err_cost,
455 ));
456
457 if step_config.allow_failure() {
458 self.has_allowed_failure = true;
459 info!(
460 run_id = %self.run_id,
461 step = %step_name,
462 error = %err_msg,
463 "parallel step failed but allow_failure is set, continuing"
464 );
465 indexed_results[idx] = Some(Ok(allowed_failure_output(
466 &err_msg,
467 raw_response_output,
468 partial.as_ref(),
469 )));
470 } else {
471 indexed_results[idx] = Some(Err(err_msg.clone()));
472
473 if first_error.is_none() {
474 first_error = Some(err);
475 }
476
477 if fail_fast {
478 join_set.abort_all();
479 }
480 }
481 }
482 }
483 }
484
485 if let Some(err) = first_error {
486 return Err(err);
487 }
488
489 self.persist_progress().await;
490
491 for (idx, (step_id, _trace_id, name, config)) in step_records.iter().enumerate() {
493 let mut output = match indexed_results[idx].take() {
494 Some(Ok(o)) => o,
495 _ => unreachable!("all steps succeeded if no error returned"),
496 };
497 output.artifacts = StepArtifacts::new(name, Some(*step_id), config.declared_outputs());
498 slots[record_slots[idx]] = Some(ParallelStepResult {
499 name: name.clone(),
500 output,
501 step_id: *step_id,
502 });
503 }
504 let results: Vec<ParallelStepResult> = slots.into_iter().flatten().collect();
505 self.last_step_ids = results.iter().map(|r| r.step_id).collect();
506
507 Ok(results)
508 }
509
510 fn replay_wave(
518 &mut self,
519 position: u32,
520 steps: &[(&str, StepConfig)],
521 ) -> Vec<Option<ParallelStepResult>> {
522 let mut slots = Vec::with_capacity(steps.len());
523 for (name, config) in steps {
524 let Some(step) = self
525 .replay_wave_steps
526 .get(&(position, (*name).to_string()))
527 .filter(|step| step.status.state == StepStatus::Completed)
528 else {
529 slots.push(None);
530 continue;
531 };
532
533 let mut output = StepOutput::from(step);
534 let step_id = step.id;
535 self.total_duration_ms += output.duration_ms;
539 output.artifacts = StepArtifacts::new(name, Some(step_id), config.declared_outputs());
540
541 info!(
542 run_id = %self.run_id,
543 step = %name,
544 position,
545 "step replayed from previous execution"
546 );
547
548 slots.push(Some(ParallelStepResult {
549 name: (*name).to_string(),
550 output,
551 step_id,
552 }));
553 }
554 slots
555 }
556}
557
558fn reject_duplicate_names(steps: &[(&str, StepConfig)]) -> Result<(), EngineError> {
564 let mut seen = HashSet::with_capacity(steps.len());
565 for (name, _) in steps {
566 if !seen.insert(*name) {
567 return Err(EngineError::StepConfig(format!(
568 "parallel wave has two steps named {name:?}; each step of a wave needs its own name"
569 )));
570 }
571 }
572 Ok(())
573}