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