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 }
501}
502
503pub(crate) fn planned_custom_output(estimate: Option<Duration>) -> StepOutput {
505 StepOutput {
506 output: json!({}),
507 duration_ms: estimate.map(|d| d.as_millis() as u64).unwrap_or(0),
508 cost_usd: Decimal::ZERO,
509 input_tokens: None,
510 cache_read_input_tokens: None,
511 cache_creation_input_tokens: None,
512 output_tokens: None,
513 model: None,
514 debug_messages: None,
515 artifacts: StepArtifacts::default(),
516 }
517}
518
519pub async fn estimate_durations(
545 store: &Arc<dyn Store>,
546 workflow_name: &str,
547 sample_runs: u32,
548) -> Result<HashMap<String, Duration>, EngineError> {
549 let page = store
550 .list_runs(
551 RunFilter {
552 workflow_name: Some(workflow_name.to_string()),
553 status: Some(RunStatus::Completed),
554 has_steps: Some(true),
555 ..RunFilter::default()
556 },
557 1,
558 sample_runs.clamp(1, 100),
559 )
560 .await?;
561
562 let mut totals: HashMap<String, (u64, u64)> = HashMap::new();
563 for run in &page.items {
564 for step in store.list_steps(run.id).await? {
565 if step.status.state != StepStatus::Completed {
566 continue;
567 }
568 let entry = totals.entry(step.name).or_insert((0, 0));
569 entry.0 += step.duration_ms;
570 entry.1 += 1;
571 }
572 }
573
574 Ok(totals
575 .into_iter()
576 .filter(|(_, (_, count))| *count > 0)
577 .map(|(name, (sum, count))| (name, Duration::from_millis(sum / count)))
578 .collect())
579}
580
581impl ExecutionPlan {
582 pub fn estimated_duration_ms(&self) -> Option<u64> {
602 self.estimated_duration.map(|d| d.as_millis() as u64)
603 }
604}
605
606impl PlannedStep {
607 pub fn estimated_duration_ms(&self) -> Option<u64> {
630 self.estimated_duration.map(|d| d.as_millis() as u64)
631 }
632}
633
634#[cfg(test)]
635mod tests {
636 use std::collections::HashMap;
637 use std::thread::spawn;
638 use std::time::Duration;
639
640 use serde_json::{from_value, json, to_value};
641
642 use crate::config::{HttpConfig, ShellConfig};
643
644 use super::*;
645
646 fn step(name: &str, group: Option<&str>, estimate: Option<u64>) -> PlannedStep {
647 PlannedStep {
648 name: name.to_string(),
649 kind: StepKind::Shell,
650 workflow: "wf".to_string(),
651 depth: 0,
652 depends_on: Vec::new(),
653 condition: None,
654 parallel_group: group.map(str::to_string),
655 estimated_duration: estimate.map(Duration::from_millis),
656 }
657 }
658
659 #[test]
660 fn total_estimate_sums_sequential_steps() {
661 let steps = vec![
662 step("a", None, Some(100)),
663 step("b", None, Some(250)),
664 step("c", None, Some(50)),
665 ];
666 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(400)));
667 }
668
669 #[test]
670 fn total_estimate_counts_a_parallel_group_once_at_its_slowest() {
671 let steps = vec![
672 step("build", None, Some(100)),
673 step("t1", Some("parallel-1"), Some(300)),
674 step("t2", Some("parallel-1"), Some(700)),
675 step("t3", Some("parallel-1"), Some(200)),
676 step("deploy", None, Some(100)),
677 ];
678 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(900)));
679 }
680
681 #[test]
682 fn total_estimate_handles_a_trailing_parallel_group() {
683 let steps = vec![
684 step("build", None, Some(100)),
685 step("t1", Some("parallel-1"), Some(300)),
686 step("t2", Some("parallel-1"), Some(700)),
687 ];
688 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(800)));
689 }
690
691 #[test]
692 fn total_estimate_is_none_without_any_estimate() {
693 let steps = vec![step("a", None, None), step("b", None, None)];
694 assert_eq!(total_estimate(&steps), None);
695 }
696
697 #[test]
698 fn planned_step_duration_round_trips_as_milliseconds() {
699 let original = step("a", None, Some(1234));
700 let value = to_value(&original).expect("serialize");
701 assert_eq!(value["estimated_duration"], 1234);
702
703 let back: PlannedStep = from_value(value).expect("deserialize");
704 assert_eq!(back.estimated_duration, Some(Duration::from_millis(1234)));
705 }
706
707 #[test]
708 fn planned_step_duration_round_trips_when_absent() {
709 let original = step("a", None, None);
710 let value = to_value(&original).expect("serialize");
711 assert!(value["estimated_duration"].is_null());
712
713 let back: PlannedStep = from_value(value).expect("deserialize");
714 assert_eq!(back.estimated_duration, None);
715 }
716
717 #[test]
718 fn planned_shell_and_http_outputs_look_successful() {
719 let shell = planned_output(&StepConfig::Shell(ShellConfig::new("echo hi")), None);
720 assert!(shell.is_success());
721
722 let http = planned_output(
723 &StepConfig::Http(HttpConfig::get("https://example.com")),
724 None,
725 );
726 assert!(http.is_success());
727 }
728
729 #[test]
730 fn planned_output_carries_the_estimate_as_its_duration() {
731 let output = planned_output(
732 &StepConfig::Shell(ShellConfig::new("echo hi")),
733 Some(Duration::from_millis(900)),
734 );
735 assert_eq!(output.duration_ms, 900);
736 assert_eq!(output.cost_usd, Decimal::ZERO);
737 }
738
739 #[test]
740 fn planned_custom_output_is_an_empty_object() {
741 let output = planned_custom_output(None);
742 assert_eq!(output.output, json!({}));
743 assert_eq!(output.duration_ms, 0);
744 }
745
746 #[test]
747 fn condition_result_serializes_its_state_tag() {
748 let evaluated = to_value(ConditionResult::Evaluated {
749 expression: "env == prod".to_string(),
750 value: false,
751 })
752 .expect("serialize");
753 assert_eq!(evaluated["state"], "evaluated");
754 assert_eq!(evaluated["value"], false);
755
756 let skipped = to_value(ConditionResult::Skipped {
757 reason: "not prod".to_string(),
758 })
759 .expect("serialize");
760 assert_eq!(skipped["state"], "skipped");
761
762 let unevaluable = to_value(ConditionResult::Unevaluable {
763 expression: "build succeeded".to_string(),
764 reason: "depends on a step output".to_string(),
765 })
766 .expect("serialize");
767 assert_eq!(unevaluable["state"], "unevaluable");
768 }
769
770 #[test]
771 fn recorder_records_dependencies_and_conditions() {
772 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
773 assert!(recorder.record("a", StepKind::Shell, "wf", None));
774 recorder.set_last(vec!["a".to_string()]);
775 recorder.set_condition(ConditionResult::Skipped {
776 reason: "nope".to_string(),
777 });
778 assert!(recorder.record("b", StepKind::Shell, "wf", None));
779
780 let plan = recorder.into_plan();
781 assert_eq!(plan.steps.len(), 2);
782 assert_eq!(plan.steps[1].depends_on, vec!["a".to_string()]);
783 assert!(matches!(
784 plan.steps[1].condition,
785 Some(ConditionResult::Skipped { .. })
786 ));
787 assert!(plan.steps[0].condition.is_none());
788 }
789
790 #[test]
791 fn recorder_stops_at_the_step_cap() {
792 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
793 for index in 0..MAX_PLANNED_STEPS {
794 assert!(recorder.record(&format!("s{index}"), StepKind::Shell, "wf", None));
795 }
796 assert!(!recorder.record("overflow", StepKind::Shell, "wf", None));
797
798 let plan = recorder.into_plan();
799 assert_eq!(plan.steps.len(), MAX_PLANNED_STEPS);
800 assert!(plan.truncated);
801 assert!(
802 plan.incomplete_reason
803 .expect("a reason")
804 .contains("step cap")
805 );
806 }
807
808 #[test]
809 fn recorder_refuses_to_expand_past_the_depth_limit() {
810 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 1, HashMap::new());
811 assert!(recorder.enter_workflow());
812 assert!(!recorder.enter_workflow());
813
814 let plan = recorder.snapshot();
815 assert!(plan.truncated);
816 assert!(
817 plan.incomplete_reason
818 .expect("a reason")
819 .contains("depth 1")
820 );
821 }
822
823 #[test]
824 fn recorder_allocates_successive_parallel_group_names() {
825 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
826 assert_eq!(recorder.next_group(), "parallel-1");
827 assert_eq!(recorder.next_group(), "parallel-2");
828 }
829
830 #[test]
831 fn recorder_swaps_and_restores_the_payload() {
832 let mut recorder =
833 PlanRecorder::new("wf".to_string(), json!({"env": "prod"}), 3, HashMap::new());
834 let previous = recorder.swap_payload(json!({"env": "dev"}));
835 assert_eq!(previous, json!({"env": "prod"}));
836 assert_eq!(recorder.payload(), json!({"env": "dev"}));
837 recorder.swap_payload(previous);
838 assert_eq!(recorder.payload(), json!({"env": "prod"}));
839 }
840
841 #[test]
842 fn recorder_keeps_the_first_failure_reason() {
843 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
844 recorder.fail("first".to_string());
845 recorder.fail("second".to_string());
846 let plan = recorder.into_plan();
847 assert_eq!(plan.incomplete_reason.as_deref(), Some("first"));
848 assert!(plan.truncated);
849 }
850
851 #[test]
852 fn recorder_uses_the_history_estimate_for_a_known_step() {
853 let estimates = HashMap::from([("build".to_string(), Duration::from_millis(400))]);
854 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, estimates);
855 assert_eq!(
856 recorder.estimate_for("build"),
857 Some(Duration::from_millis(400))
858 );
859 assert_eq!(recorder.estimate_for("unknown"), None);
860 recorder.record("build", StepKind::Shell, "wf", None);
861 let plan = recorder.into_plan();
862 assert_eq!(plan.estimated_duration, Some(Duration::from_millis(400)));
863 }
864
865 #[test]
866 fn lock_plan_recovers_from_poisoning() {
867 let shared: SharedPlanRecorder = Arc::new(Mutex::new(PlanRecorder::new(
868 "wf".to_string(),
869 json!({}),
870 3,
871 HashMap::new(),
872 )));
873 let poisoner = Arc::clone(&shared);
874 let _ = spawn(move || {
875 let _guard = poisoner.lock().expect("lock");
876 panic!("poison the mutex");
877 })
878 .join();
879
880 assert_eq!(lock_plan(&shared).payload(), json!({}));
881 }
882}