1use std::collections::HashMap;
33use std::mem::replace;
34use std::sync::{Arc, Mutex, MutexGuard};
35use std::time::Duration;
36
37use rust_decimal::Decimal;
38use serde::{Deserialize, Serialize};
39use serde_json::{Value, json};
40
41use ironflow_store::entities::{RunFilter, RunStatus, StepKind, StepStatus};
42use ironflow_store::store::Store;
43
44use crate::config::StepConfig;
45use crate::error::EngineError;
46use crate::executor::{StepArtifacts, StepOutput};
47
48pub const DEFAULT_PLAN_MAX_DEPTH: u32 = 3;
50
51pub const MAX_PLANNED_STEPS: usize = 1000;
56
57pub const DEFAULT_ESTIMATE_SAMPLE_RUNS: u32 = 20;
59
60#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
79#[serde(tag = "state", rename_all = "snake_case")]
80pub enum ConditionResult {
81 Evaluated {
83 expression: String,
86 value: bool,
88 },
89 Skipped {
92 reason: String,
94 },
95 Unevaluable {
97 expression: String,
99 reason: String,
101 },
102}
103
104#[derive(Debug, Clone, Serialize, Deserialize)]
125pub struct PlannedStep {
126 pub name: String,
128 pub kind: StepKind,
130 pub workflow: String,
132 pub depth: u32,
134 pub depends_on: Vec<String>,
140 pub condition: Option<ConditionResult>,
146 pub parallel_group: Option<String>,
149 #[serde(with = "opt_duration_ms")]
151 pub estimated_duration: Option<Duration>,
152}
153
154#[derive(Debug, Clone, Serialize, Deserialize)]
172pub struct ExecutionPlan {
173 pub workflow: String,
175 pub steps: Vec<PlannedStep>,
177 #[serde(with = "opt_duration_ms")]
179 pub estimated_duration: Option<Duration>,
180 pub max_depth: u32,
182 pub truncated: bool,
184 pub incomplete_reason: Option<String>,
186}
187
188#[derive(Debug, Clone)]
199pub struct PlanOptions {
200 pub max_depth: u32,
202 pub estimate_durations: bool,
204 pub sample_runs: u32,
206}
207
208impl Default for PlanOptions {
209 fn default() -> Self {
210 Self {
211 max_depth: DEFAULT_PLAN_MAX_DEPTH,
212 estimate_durations: true,
213 sample_runs: DEFAULT_ESTIMATE_SAMPLE_RUNS,
214 }
215 }
216}
217
218mod opt_duration_ms {
220 use std::time::Duration;
221
222 use serde::{Deserialize, Deserializer, Serialize, Serializer};
223
224 pub(super) fn serialize<S: Serializer>(
226 value: &Option<Duration>,
227 serializer: S,
228 ) -> Result<S::Ok, S::Error> {
229 let millis = value.map(|d| d.as_millis() as u64);
230 millis.serialize(serializer)
231 }
232
233 pub(super) fn deserialize<'de, D: Deserializer<'de>>(
235 deserializer: D,
236 ) -> Result<Option<Duration>, D::Error> {
237 let millis = Option::<u64>::deserialize(deserializer)?;
238 Ok(millis.map(Duration::from_millis))
239 }
240}
241
242pub(crate) type SharedPlanRecorder = Arc<Mutex<PlanRecorder>>;
244
245pub(crate) struct PlanRecorder {
247 payload: Value,
248 workflow: String,
249 steps: Vec<PlannedStep>,
250 last_names: Vec<String>,
251 pending_condition: Option<ConditionResult>,
252 depth: u32,
253 max_depth: u32,
254 parallel_groups: u32,
255 estimates: HashMap<String, Duration>,
256 truncated: bool,
257 incomplete_reason: Option<String>,
258}
259
260impl PlanRecorder {
261 pub(crate) fn new(
263 workflow: String,
264 payload: Value,
265 max_depth: u32,
266 estimates: HashMap<String, Duration>,
267 ) -> Self {
268 Self {
269 payload,
270 workflow,
271 steps: Vec::new(),
272 last_names: Vec::new(),
273 pending_condition: None,
274 depth: 0,
275 max_depth,
276 parallel_groups: 0,
277 estimates,
278 truncated: false,
279 incomplete_reason: None,
280 }
281 }
282
283 pub(crate) fn payload(&self) -> Value {
285 self.payload.clone()
286 }
287
288 pub(crate) fn swap_payload(&mut self, next: Value) -> Value {
293 replace(&mut self.payload, next)
294 }
295
296 pub(crate) fn set_condition(&mut self, condition: ConditionResult) {
298 self.pending_condition = Some(condition);
299 }
300
301 pub(crate) fn estimate_for(&self, name: &str) -> Option<Duration> {
303 self.estimates.get(name).copied()
304 }
305
306 pub(crate) fn seed_estimate(&mut self, name: &str, duration: Duration) {
311 self.estimates.entry(name.to_string()).or_insert(duration);
312 }
313
314 pub(crate) fn record(
319 &mut self,
320 name: &str,
321 kind: StepKind,
322 workflow: &str,
323 parallel_group: Option<String>,
324 ) -> bool {
325 if self.steps.len() >= MAX_PLANNED_STEPS {
326 self.truncated = true;
327 if self.incomplete_reason.is_none() {
328 self.incomplete_reason = Some(format!("step cap of {MAX_PLANNED_STEPS} reached"));
329 }
330 return false;
331 }
332
333 let estimated_duration = self.estimates.get(name).copied();
334 self.steps.push(PlannedStep {
335 name: name.to_string(),
336 kind,
337 workflow: workflow.to_string(),
338 depth: self.depth,
339 depends_on: self.last_names.clone(),
340 condition: self.pending_condition.take(),
341 parallel_group,
342 estimated_duration,
343 });
344 true
345 }
346
347 pub(crate) fn set_last(&mut self, names: Vec<String>) {
349 self.last_names = names;
350 }
351
352 pub(crate) fn next_group(&mut self) -> String {
354 self.parallel_groups += 1;
355 format!("parallel-{}", self.parallel_groups)
356 }
357
358 pub(crate) fn enter_workflow(&mut self) -> bool {
360 if self.depth + 1 > self.max_depth {
361 self.truncated = true;
362 if self.incomplete_reason.is_none() {
363 self.incomplete_reason = Some(format!(
364 "sub-workflow expansion stopped at depth {}",
365 self.max_depth
366 ));
367 }
368 return false;
369 }
370 self.depth += 1;
371 true
372 }
373
374 pub(crate) fn leave_workflow(&mut self) {
376 self.depth = self.depth.saturating_sub(1);
377 }
378
379 pub(crate) fn fail(&mut self, reason: String) {
381 self.truncated = true;
382 if self.incomplete_reason.is_none() {
383 self.incomplete_reason = Some(reason);
384 }
385 }
386
387 pub(crate) fn snapshot(&self) -> ExecutionPlan {
389 ExecutionPlan {
390 workflow: self.workflow.clone(),
391 estimated_duration: total_estimate(&self.steps),
392 steps: self.steps.clone(),
393 max_depth: self.max_depth,
394 truncated: self.truncated,
395 incomplete_reason: self.incomplete_reason.clone(),
396 }
397 }
398
399 pub(crate) fn into_plan(self) -> ExecutionPlan {
401 ExecutionPlan {
402 workflow: self.workflow,
403 estimated_duration: total_estimate(&self.steps),
404 steps: self.steps,
405 max_depth: self.max_depth,
406 truncated: self.truncated,
407 incomplete_reason: self.incomplete_reason,
408 }
409 }
410}
411
412pub(crate) fn lock_plan(plan: &SharedPlanRecorder) -> MutexGuard<'_, PlanRecorder> {
417 plan.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
418}
419
420fn total_estimate(steps: &[PlannedStep]) -> Option<Duration> {
426 let mut total = Duration::ZERO;
427 let mut seen_any = false;
428 let mut current_group: Option<&str> = None;
429 let mut group_max = Duration::ZERO;
430
431 for step in steps {
432 match step.parallel_group.as_deref() {
433 Some(group) if current_group == Some(group) => {
434 if let Some(estimate) = step.estimated_duration {
435 seen_any = true;
436 group_max = group_max.max(estimate);
437 }
438 }
439 Some(group) => {
440 if current_group.is_some() {
441 total += group_max;
442 }
443 current_group = Some(group);
444 group_max = step.estimated_duration.unwrap_or(Duration::ZERO);
445 if step.estimated_duration.is_some() {
446 seen_any = true;
447 }
448 }
449 None => {
450 if current_group.is_some() {
451 total += group_max;
452 current_group = None;
453 group_max = Duration::ZERO;
454 }
455 if let Some(estimate) = step.estimated_duration {
456 seen_any = true;
457 total += estimate;
458 }
459 }
460 }
461 }
462
463 if current_group.is_some() {
464 total += group_max;
465 }
466
467 seen_any.then_some(total)
468}
469
470pub(crate) fn planned_output(config: &StepConfig, estimate: Option<Duration>) -> StepOutput {
475 let output = match config {
476 StepConfig::Shell(_) => json!({"stdout": "", "stderr": "", "exit_code": 0}),
477 StepConfig::Http(_) => json!({"status": 200, "headers": {}, "body": ""}),
478 StepConfig::Agent(_) => json!({}),
479 StepConfig::Workflow(c) => json!({
480 "run_id": Value::Null,
481 "workflow_name": c.workflow_name,
482 "status": "completed",
483 "cost_usd": 0,
484 "duration_ms": 0,
485 }),
486 StepConfig::Approval(_) | StepConfig::Decision(_) | StepConfig::Delay(_) => Value::Null,
487 };
488
489 StepOutput {
490 output,
491 duration_ms: estimate.map(|d| d.as_millis() as u64).unwrap_or(0),
492 cost_usd: Decimal::ZERO,
493 input_tokens: None,
494 cache_read_input_tokens: None,
495 cache_creation_input_tokens: None,
496 output_tokens: None,
497 model: None,
498 debug_messages: None,
499 artifacts: StepArtifacts::default(),
500 account_id: None,
501 }
502}
503
504pub(crate) fn planned_custom_output(estimate: Option<Duration>) -> StepOutput {
506 StepOutput {
507 output: json!({}),
508 duration_ms: estimate.map(|d| d.as_millis() as u64).unwrap_or(0),
509 cost_usd: Decimal::ZERO,
510 input_tokens: None,
511 cache_read_input_tokens: None,
512 cache_creation_input_tokens: None,
513 output_tokens: None,
514 model: None,
515 debug_messages: None,
516 artifacts: StepArtifacts::default(),
517 account_id: None,
518 }
519}
520
521pub async fn estimate_durations(
547 store: &Arc<dyn Store>,
548 workflow_name: &str,
549 sample_runs: u32,
550) -> Result<HashMap<String, Duration>, EngineError> {
551 let page = store
552 .list_runs(
553 RunFilter {
554 workflow_name: Some(workflow_name.to_string()),
555 status: Some(RunStatus::Completed),
556 has_steps: Some(true),
557 ..RunFilter::default()
558 },
559 1,
560 sample_runs.clamp(1, 100),
561 )
562 .await?;
563
564 let mut totals: HashMap<String, (u64, u64)> = HashMap::new();
565 for run in &page.items {
566 for step in store.list_steps(run.id).await? {
567 if step.status.state != StepStatus::Completed {
568 continue;
569 }
570 let entry = totals.entry(step.name).or_insert((0, 0));
571 entry.0 += step.duration_ms;
572 entry.1 += 1;
573 }
574 }
575
576 Ok(totals
577 .into_iter()
578 .filter(|(_, (_, count))| *count > 0)
579 .map(|(name, (sum, count))| (name, Duration::from_millis(sum / count)))
580 .collect())
581}
582
583impl ExecutionPlan {
584 pub fn estimated_duration_ms(&self) -> Option<u64> {
604 self.estimated_duration.map(|d| d.as_millis() as u64)
605 }
606}
607
608impl PlannedStep {
609 pub fn estimated_duration_ms(&self) -> Option<u64> {
632 self.estimated_duration.map(|d| d.as_millis() as u64)
633 }
634}
635
636#[cfg(test)]
637mod tests {
638 use std::collections::HashMap;
639 use std::thread::spawn;
640 use std::time::Duration;
641
642 use serde_json::{from_value, json, to_value};
643
644 use crate::config::{HttpConfig, ShellConfig};
645
646 use super::*;
647
648 fn step(name: &str, group: Option<&str>, estimate: Option<u64>) -> PlannedStep {
649 PlannedStep {
650 name: name.to_string(),
651 kind: StepKind::Shell,
652 workflow: "wf".to_string(),
653 depth: 0,
654 depends_on: Vec::new(),
655 condition: None,
656 parallel_group: group.map(str::to_string),
657 estimated_duration: estimate.map(Duration::from_millis),
658 }
659 }
660
661 #[test]
662 fn total_estimate_sums_sequential_steps() {
663 let steps = vec![
664 step("a", None, Some(100)),
665 step("b", None, Some(250)),
666 step("c", None, Some(50)),
667 ];
668 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(400)));
669 }
670
671 #[test]
672 fn total_estimate_counts_a_parallel_group_once_at_its_slowest() {
673 let steps = vec![
674 step("build", None, Some(100)),
675 step("t1", Some("parallel-1"), Some(300)),
676 step("t2", Some("parallel-1"), Some(700)),
677 step("t3", Some("parallel-1"), Some(200)),
678 step("deploy", None, Some(100)),
679 ];
680 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(900)));
681 }
682
683 #[test]
684 fn total_estimate_handles_a_trailing_parallel_group() {
685 let steps = vec![
686 step("build", None, Some(100)),
687 step("t1", Some("parallel-1"), Some(300)),
688 step("t2", Some("parallel-1"), Some(700)),
689 ];
690 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(800)));
691 }
692
693 #[test]
694 fn total_estimate_is_none_without_any_estimate() {
695 let steps = vec![step("a", None, None), step("b", None, None)];
696 assert_eq!(total_estimate(&steps), None);
697 }
698
699 #[test]
700 fn planned_step_duration_round_trips_as_milliseconds() {
701 let original = step("a", None, Some(1234));
702 let value = to_value(&original).expect("serialize");
703 assert_eq!(value["estimated_duration"], 1234);
704
705 let back: PlannedStep = from_value(value).expect("deserialize");
706 assert_eq!(back.estimated_duration, Some(Duration::from_millis(1234)));
707 }
708
709 #[test]
710 fn planned_step_duration_round_trips_when_absent() {
711 let original = step("a", None, None);
712 let value = to_value(&original).expect("serialize");
713 assert!(value["estimated_duration"].is_null());
714
715 let back: PlannedStep = from_value(value).expect("deserialize");
716 assert_eq!(back.estimated_duration, None);
717 }
718
719 #[test]
720 fn planned_shell_and_http_outputs_look_successful() {
721 let shell = planned_output(&StepConfig::Shell(ShellConfig::new("echo hi")), None);
722 assert!(shell.is_success());
723
724 let http = planned_output(
725 &StepConfig::Http(HttpConfig::get("https://example.com")),
726 None,
727 );
728 assert!(http.is_success());
729 }
730
731 #[test]
732 fn planned_output_carries_the_estimate_as_its_duration() {
733 let output = planned_output(
734 &StepConfig::Shell(ShellConfig::new("echo hi")),
735 Some(Duration::from_millis(900)),
736 );
737 assert_eq!(output.duration_ms, 900);
738 assert_eq!(output.cost_usd, Decimal::ZERO);
739 }
740
741 #[test]
742 fn planned_custom_output_is_an_empty_object() {
743 let output = planned_custom_output(None);
744 assert_eq!(output.output, json!({}));
745 assert_eq!(output.duration_ms, 0);
746 }
747
748 #[test]
749 fn condition_result_serializes_its_state_tag() {
750 let evaluated = to_value(ConditionResult::Evaluated {
751 expression: "env == prod".to_string(),
752 value: false,
753 })
754 .expect("serialize");
755 assert_eq!(evaluated["state"], "evaluated");
756 assert_eq!(evaluated["value"], false);
757
758 let skipped = to_value(ConditionResult::Skipped {
759 reason: "not prod".to_string(),
760 })
761 .expect("serialize");
762 assert_eq!(skipped["state"], "skipped");
763
764 let unevaluable = to_value(ConditionResult::Unevaluable {
765 expression: "build succeeded".to_string(),
766 reason: "depends on a step output".to_string(),
767 })
768 .expect("serialize");
769 assert_eq!(unevaluable["state"], "unevaluable");
770 }
771
772 #[test]
773 fn recorder_records_dependencies_and_conditions() {
774 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
775 assert!(recorder.record("a", StepKind::Shell, "wf", None));
776 recorder.set_last(vec!["a".to_string()]);
777 recorder.set_condition(ConditionResult::Skipped {
778 reason: "nope".to_string(),
779 });
780 assert!(recorder.record("b", StepKind::Shell, "wf", None));
781
782 let plan = recorder.into_plan();
783 assert_eq!(plan.steps.len(), 2);
784 assert_eq!(plan.steps[1].depends_on, vec!["a".to_string()]);
785 assert!(matches!(
786 plan.steps[1].condition,
787 Some(ConditionResult::Skipped { .. })
788 ));
789 assert!(plan.steps[0].condition.is_none());
790 }
791
792 #[test]
793 fn recorder_stops_at_the_step_cap() {
794 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
795 for index in 0..MAX_PLANNED_STEPS {
796 assert!(recorder.record(&format!("s{index}"), StepKind::Shell, "wf", None));
797 }
798 assert!(!recorder.record("overflow", StepKind::Shell, "wf", None));
799
800 let plan = recorder.into_plan();
801 assert_eq!(plan.steps.len(), MAX_PLANNED_STEPS);
802 assert!(plan.truncated);
803 assert!(
804 plan.incomplete_reason
805 .expect("a reason")
806 .contains("step cap")
807 );
808 }
809
810 #[test]
811 fn recorder_refuses_to_expand_past_the_depth_limit() {
812 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 1, HashMap::new());
813 assert!(recorder.enter_workflow());
814 assert!(!recorder.enter_workflow());
815
816 let plan = recorder.snapshot();
817 assert!(plan.truncated);
818 assert!(
819 plan.incomplete_reason
820 .expect("a reason")
821 .contains("depth 1")
822 );
823 }
824
825 #[test]
826 fn recorder_allocates_successive_parallel_group_names() {
827 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
828 assert_eq!(recorder.next_group(), "parallel-1");
829 assert_eq!(recorder.next_group(), "parallel-2");
830 }
831
832 #[test]
833 fn recorder_swaps_and_restores_the_payload() {
834 let mut recorder =
835 PlanRecorder::new("wf".to_string(), json!({"env": "prod"}), 3, HashMap::new());
836 let previous = recorder.swap_payload(json!({"env": "dev"}));
837 assert_eq!(previous, json!({"env": "prod"}));
838 assert_eq!(recorder.payload(), json!({"env": "dev"}));
839 recorder.swap_payload(previous);
840 assert_eq!(recorder.payload(), json!({"env": "prod"}));
841 }
842
843 #[test]
844 fn recorder_keeps_the_first_failure_reason() {
845 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
846 recorder.fail("first".to_string());
847 recorder.fail("second".to_string());
848 let plan = recorder.into_plan();
849 assert_eq!(plan.incomplete_reason.as_deref(), Some("first"));
850 assert!(plan.truncated);
851 }
852
853 #[test]
854 fn recorder_uses_the_history_estimate_for_a_known_step() {
855 let estimates = HashMap::from([("build".to_string(), Duration::from_millis(400))]);
856 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, estimates);
857 assert_eq!(
858 recorder.estimate_for("build"),
859 Some(Duration::from_millis(400))
860 );
861 assert_eq!(recorder.estimate_for("unknown"), None);
862 recorder.record("build", StepKind::Shell, "wf", None);
863 let plan = recorder.into_plan();
864 assert_eq!(plan.estimated_duration, Some(Duration::from_millis(400)));
865 }
866
867 #[test]
868 fn lock_plan_recovers_from_poisoning() {
869 let shared: SharedPlanRecorder = Arc::new(Mutex::new(PlanRecorder::new(
870 "wf".to_string(),
871 json!({}),
872 3,
873 HashMap::new(),
874 )));
875 let poisoner = Arc::clone(&shared);
876 let _ = spawn(move || {
877 let _guard = poisoner.lock().expect("lock");
878 panic!("poison the mutex");
879 })
880 .join();
881
882 assert_eq!(lock_plan(&shared).payload(), json!({}));
883 }
884}