1use std::collections::HashMap;
27use std::fmt;
28use std::sync::Arc;
29use std::time::Instant;
30
31use chrono::{DateTime, Utc};
32use futures_util::StreamExt;
33use rust_decimal::Decimal;
34use serde_json::{Value, json};
35use tokio::task::{Id, JoinSet};
36use tracing::{Span, error, info, warn};
37use uuid::Uuid;
38
39use ironflow_core::decision::{DecisionOutput, DecisionProvider};
40use ironflow_core::error::{AgentError, OperationError};
41
42mod decision_impl;
43use ironflow_core::provider::AgentProvider;
44use ironflow_core::trace_context::WorkflowTraceContext;
45use ironflow_store::models::{
46 ArtifactLookup, NewRun, NewStep, NewStepDependency, RunStatus, RunUpdate, Step, StepKind,
47 StepStatus, StepUpdate, TriggerKind, step_trace_id,
48};
49use ironflow_store::store::Store;
50
51use ironflow_artifacts::name::guess_content_type;
52use ironflow_artifacts::stream_from_bytes;
53use ironflow_store::entities::Artifact;
54
55use crate::artifact::{
56 ArtifactSink, ArtifactUpload, StepLocation, collect_outputs, materialize_inputs,
57};
58use crate::budget::step_budget_usd;
59use crate::config::{
60 AgentStepConfig, ApprovalConfig, DecisionConfig, HttpConfig, ShellConfig, StepConfig,
61 WorkflowStepConfig,
62};
63use crate::error::EngineError;
64use crate::executor::{ParallelStepResult, StepOutput, StepResult, execute_step_config};
65use crate::guard::{SharedGuardState, WorkflowGuardConfig, WorkflowRejection};
66use crate::handler::WorkflowHandler;
67use crate::log_sender::{LogSender, StepLogSender};
68use crate::notify::{
69 WorkflowAgentStepTokensUsedEvent, WorkflowApprovalRequiredEvent, WorkflowEvent,
70 WorkflowEventBus, WorkflowStepCompletedEvent, WorkflowStepFailedEvent,
71 WorkflowStepStartedEvent,
72};
73#[cfg(not(feature = "secret-store"))]
74use crate::operation::NoopSecretResolver;
75use crate::operation::{Operation, OperationContext, SecretResolver};
76#[cfg(feature = "secret-store")]
77use ironflow_store::workflow_secrets::ScopedSecretStore;
78
79pub(crate) type HandlerResolver =
81 Arc<dyn Fn(&str) -> Option<Arc<dyn WorkflowHandler>> + Send + Sync>;
82
83pub struct WorkflowContext {
102 run_id: Uuid,
103 workflow_name: String,
104 store: Arc<dyn Store>,
105 provider: Arc<dyn AgentProvider>,
106 decision_provider: Option<Arc<dyn DecisionProvider>>,
110 handler_resolver: Option<HandlerResolver>,
111 position: u32,
112 last_step_ids: Vec<Uuid>,
114 total_cost_usd: Decimal,
116 total_duration_ms: u64,
118 max_cost_usd: Option<Decimal>,
120 inherited_cost_usd: Decimal,
123 replay_steps: HashMap<u32, Step>,
126 granted_approvals: HashMap<u32, u32>,
130 attempt: u32,
132 carried_duration_ms: u64,
135 log_sender: Option<LogSender>,
137 artifact_sink: Option<Arc<dyn ArtifactSink>>,
141 has_allowed_failure: bool,
143 error_handlers: Vec<OnErrorHandler>,
145 guard_state: Option<SharedGuardState>,
147 guard_config: Option<WorkflowGuardConfig>,
149 step_results: Vec<StepResult>,
151 event_bus: Option<WorkflowEventBus>,
153 trace_context: WorkflowTraceContext,
155 operation_ctx: Option<OperationContext>,
157}
158
159struct OnErrorHandler {
161 name: String,
162 config: StepConfig,
163}
164
165impl WorkflowContext {
166 pub fn new(
171 run_id: Uuid,
172 workflow_name: String,
173 store: Arc<dyn Store>,
174 provider: Arc<dyn AgentProvider>,
175 ) -> Self {
176 let trace_context = WorkflowTraceContext::from_workflow_run_id(&run_id.to_string());
177 Self {
178 run_id,
179 workflow_name,
180 store,
181 provider,
182 decision_provider: None,
183 handler_resolver: None,
184 position: 0,
185 last_step_ids: Vec::new(),
186 total_cost_usd: Decimal::ZERO,
187 total_duration_ms: 0,
188 max_cost_usd: None,
189 inherited_cost_usd: Decimal::ZERO,
190 replay_steps: HashMap::new(),
191 granted_approvals: HashMap::new(),
192 attempt: 1,
193 carried_duration_ms: 0,
194 log_sender: None,
195 artifact_sink: None,
196 has_allowed_failure: false,
197 error_handlers: Vec::new(),
198 guard_state: None,
199 guard_config: None,
200 step_results: Vec::new(),
201 event_bus: None,
202 trace_context,
203 operation_ctx: None,
204 }
205 }
206
207 pub(crate) fn with_handler_resolver(
212 run_id: Uuid,
213 workflow_name: String,
214 store: Arc<dyn Store>,
215 provider: Arc<dyn AgentProvider>,
216 resolver: HandlerResolver,
217 ) -> Self {
218 let trace_context = WorkflowTraceContext::from_workflow_run_id(&run_id.to_string());
219 Self {
220 run_id,
221 workflow_name,
222 store,
223 provider,
224 decision_provider: None,
225 handler_resolver: Some(resolver),
226 position: 0,
227 last_step_ids: Vec::new(),
228 total_cost_usd: Decimal::ZERO,
229 total_duration_ms: 0,
230 max_cost_usd: None,
231 inherited_cost_usd: Decimal::ZERO,
232 replay_steps: HashMap::new(),
233 granted_approvals: HashMap::new(),
234 attempt: 1,
235 carried_duration_ms: 0,
236 log_sender: None,
237 artifact_sink: None,
238 has_allowed_failure: false,
239 error_handlers: Vec::new(),
240 guard_state: None,
241 guard_config: None,
242 step_results: Vec::new(),
243 event_bus: None,
244 trace_context,
245 operation_ctx: None,
246 }
247 }
248
249 pub fn set_log_sender(&mut self, sender: LogSender) {
251 self.log_sender = Some(sender);
252 }
253
254 pub fn set_artifact_sink(&mut self, sink: Arc<dyn ArtifactSink>) {
274 self.artifact_sink = Some(sink);
275 }
276
277 pub fn trace_context(&self) -> &WorkflowTraceContext {
283 &self.trace_context
284 }
285
286 pub fn set_guard(&mut self, config: WorkflowGuardConfig, state: SharedGuardState) {
303 self.guard_config = Some(config);
304 self.guard_state = Some(state);
305 }
306
307 pub fn guard_config(&self) -> Option<&WorkflowGuardConfig> {
309 self.guard_config.as_ref()
310 }
311
312 pub fn set_event_bus(&mut self, bus: WorkflowEventBus) {
317 self.event_bus = Some(bus);
318 }
319
320 pub fn set_decision_provider(&mut self, provider: Arc<dyn DecisionProvider>) {
325 self.decision_provider = Some(provider);
326 }
327
328 fn artifact_sink(&self) -> Result<&Arc<dyn ArtifactSink>, EngineError> {
330 self.artifact_sink.as_ref().ok_or_else(|| {
331 EngineError::ArtifactsUnavailable(
332 "no artifact storage is attached to this run".to_string(),
333 )
334 })
335 }
336
337 pub async fn put_artifact(
367 &self,
368 step_id: Uuid,
369 name: &str,
370 content_type: Option<&str>,
371 content: Vec<u8>,
372 ) -> Result<Artifact, EngineError> {
373 let sink = self.artifact_sink()?;
374 sink.put(
375 ArtifactUpload {
376 run_id: self.run_id,
377 step_id,
378 name: name.to_string(),
379 content_type: content_type
380 .map(str::to_string)
381 .unwrap_or_else(|| guess_content_type(name)),
382 },
383 stream_from_bytes(content),
384 )
385 .await
386 }
387
388 pub async fn get_artifact(&self, step: &str, name: &str) -> Result<Vec<u8>, EngineError> {
413 let sink = self.artifact_sink()?;
414
415 let artifact = self
416 .store
417 .find_artifact_for_input(ArtifactLookup {
418 run_id: self.run_id,
419 attempt: self.attempt,
420 before_position: self.position,
421 step_name: step.to_string(),
422 name: name.to_string(),
423 })
424 .await?
425 .ok_or_else(|| EngineError::ArtifactNotFound {
426 step: step.to_string(),
427 name: name.to_string(),
428 })?;
429
430 let mut content = sink.get(&artifact).await?;
431 let mut buffer = Vec::with_capacity(artifact.size_bytes as usize);
432 while let Some(chunk) = content.next().await {
433 let chunk = chunk?;
434 buffer.extend_from_slice(chunk.as_ref());
435 }
436
437 Ok(buffer)
438 }
439
440 async fn prepare_step_inputs(
445 &self,
446 config: &StepConfig,
447 position: u32,
448 ) -> Result<(), EngineError> {
449 let StepConfig::Shell(shell) = config else {
450 return Ok(());
451 };
452 if shell.inputs.is_empty() {
453 return Ok(());
454 }
455
456 materialize_inputs(
457 self.artifact_sink()?,
458 &self.store,
459 shell,
460 StepLocation {
461 run_id: self.run_id,
462 attempt: self.attempt,
463 position,
464 },
465 )
466 .await
467 }
468
469 async fn store_step_outputs(
474 &self,
475 config: &StepConfig,
476 step_id: Uuid,
477 step_name: &str,
478 step_succeeded: bool,
479 ) -> Result<(), EngineError> {
480 let StepConfig::Shell(shell) = config else {
481 return Ok(());
482 };
483 if shell.outputs.is_empty() {
484 return Ok(());
485 }
486
487 let sink = match self.artifact_sink() {
488 Ok(sink) => sink,
489 Err(err) if step_succeeded => return Err(err),
490 Err(err) => {
491 warn!(
492 run_id = %self.run_id,
493 step = %step_name,
494 error = %err,
495 "cannot collect outputs of a failed step"
496 );
497 return Ok(());
498 }
499 };
500
501 let collected =
502 collect_outputs(sink, shell, self.run_id, step_id, step_name, step_succeeded).await;
503
504 match collected {
505 Ok(()) => Ok(()),
506 Err(err) if step_succeeded => Err(err),
507 Err(err) => {
508 warn!(
509 run_id = %self.run_id,
510 step = %step_name,
511 error = %err,
512 "failed to collect outputs of a failed step"
513 );
514 Ok(())
515 }
516 }
517 }
518
519 pub(crate) fn carry_over_run_totals(
526 &mut self,
527 attempt: u32,
528 cost_usd: Decimal,
529 duration_ms: u64,
530 ) {
531 self.attempt = attempt;
532 self.total_cost_usd = cost_usd;
533 self.carried_duration_ms = duration_ms;
534 }
535
536 pub(crate) fn carried_duration_ms(&self) -> u64 {
538 self.carried_duration_ms
539 }
540
541 pub fn attempt(&self) -> u32 {
543 self.attempt
544 }
545
546 pub fn set_max_cost_usd(&mut self, cap: Option<Decimal>) {
562 self.max_cost_usd = cap;
563 }
564
565 pub fn max_cost_usd(&self) -> Option<Decimal> {
567 self.max_cost_usd
568 }
569
570 pub fn charged_cost_usd(&self) -> Decimal {
575 self.inherited_cost_usd + self.total_cost_usd
576 }
577
578 fn check_run_budget(&self, step_budget: Decimal) -> Result<(), EngineError> {
589 let Some(limit) = self.max_cost_usd else {
590 return Ok(());
591 };
592
593 let spent = self.charged_cost_usd();
594 if spent + step_budget <= limit {
595 return Ok(());
596 }
597
598 error!(
599 run_id = %self.run_id,
600 limit_usd = %limit,
601 spent_usd = %spent,
602 step_budget_usd = %step_budget,
603 "run cost cap reached, refusing agent step"
604 );
605
606 Err(EngineError::RunBudgetExceeded {
607 run_id: self.run_id,
608 limit_usd: limit,
609 spent_usd: spent,
610 step_budget_usd: step_budget,
611 })
612 }
613
614 pub(crate) async fn load_replay_steps(&mut self) -> Result<(), EngineError> {
626 let steps = self.store.list_steps(self.run_id).await?;
627 for step in steps {
628 let dominated = matches!(
629 step.status.state,
630 StepStatus::Completed | StepStatus::Running | StepStatus::AwaitingApproval
631 );
632 if !dominated {
633 continue;
634 }
635
636 if step.attempt == self.attempt {
637 self.replay_steps.insert(step.position, step);
638 } else if step.kind == StepKind::Approval && step.status.state == StepStatus::Completed
639 {
640 self.granted_approvals.insert(step.position, step.attempt);
641 }
642 }
643 Ok(())
644 }
645
646 pub fn run_id(&self) -> Uuid {
648 self.run_id
649 }
650
651 pub fn workflow_name(&self) -> &str {
653 &self.workflow_name
654 }
655
656 pub fn total_cost_usd(&self) -> Decimal {
658 self.total_cost_usd
659 }
660
661 pub fn has_allowed_failure(&self) -> bool {
663 self.has_allowed_failure
664 }
665
666 pub fn total_duration_ms(&self) -> u64 {
668 self.total_duration_ms
669 }
670
671 pub fn step_results(&self) -> &[StepResult] {
673 &self.step_results
674 }
675
676 #[cfg(feature = "secret-store")]
698 pub fn secrets(&self) -> ironflow_store::workflow_secrets::ScopedSecretStore {
699 let workflow_uuid = Uuid::new_v5(&Uuid::NAMESPACE_OID, self.workflow_name.as_bytes());
700 ironflow_store::workflow_secrets::ScopedSecretStore::for_workflow(
701 workflow_uuid,
702 self.store.clone(),
703 )
704 }
705
706 fn ensure_operation_ctx(&mut self) -> &OperationContext {
707 self.operation_ctx.get_or_insert_with(|| {
708 #[cfg(feature = "secret-store")]
709 let secrets: Arc<dyn SecretResolver> = {
710 let workflow_uuid =
711 Uuid::new_v5(&Uuid::NAMESPACE_OID, self.workflow_name.as_bytes());
712 Arc::new(ScopedSecretStore::for_workflow(
713 workflow_uuid,
714 self.store.clone(),
715 ))
716 };
717 #[cfg(not(feature = "secret-store"))]
718 let secrets: Arc<dyn SecretResolver> = Arc::new(NoopSecretResolver);
719
720 OperationContext::new(secrets)
721 })
722 }
723
724 async fn persist_progress(&self) {
730 if let Err(err) = self
731 .store
732 .update_run(
733 self.run_id,
734 RunUpdate {
735 cost_usd: Some(self.total_cost_usd),
736 duration_ms: Some(self.total_duration_ms),
737 ..RunUpdate::default()
738 },
739 )
740 .await
741 {
742 warn!(
743 run_id = %self.run_id,
744 error = %err,
745 "failed to persist run progress snapshot"
746 );
747 }
748 }
749
750 pub async fn parallel(
787 &mut self,
788 steps: Vec<(&str, StepConfig)>,
789 fail_fast: bool,
790 ) -> Result<Vec<ParallelStepResult>, EngineError> {
791 if steps.is_empty() {
792 return Ok(Vec::new());
793 }
794
795 self.check_guard_timeout()?;
797
798 let wave_budget: Decimal = steps
801 .iter()
802 .filter_map(|(_, config)| match config {
803 StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
804 _ => None,
805 })
806 .map(step_budget_usd)
807 .sum();
808 self.check_run_budget(wave_budget)?;
809
810 let wave_position = self.position;
811 self.position += 1;
812
813 let now = Utc::now();
814 let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
815 Vec::with_capacity(steps.len());
816
817 for (name, config) in &steps {
818 let kind = config.kind();
819 let trace_id = step_trace_id(self.run_id, name, wave_position);
820 let step = self
821 .store
822 .create_step(NewStep {
823 run_id: self.run_id,
824 trace_id,
825 name: name.to_string(),
826 kind,
827 position: wave_position,
828 input: Some(serde_json::to_value(config)?),
829 is_error_handler: false,
830 })
831 .await?;
832
833 self.start_step(step.id, now).await?;
834
835 if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
838 self.fail_step(step.id, &err).await;
839 if !config.allow_failure() {
840 return Err(err);
841 }
842 self.has_allowed_failure = true;
843 info!(
844 run_id = %self.run_id,
845 step = %name,
846 error = %err,
847 "parallel step input preparation failed but allow_failure is set, skipping"
848 );
849 continue;
850 }
851
852 let mut config_with_trace = config.clone();
853 let step_trace = self.trace_context.child();
854 match config_with_trace {
855 StepConfig::Agent(ref mut agent_config) => {
856 agent_config.trace_context = Some(step_trace);
857 }
858 StepConfig::Http(ref mut http_config) => {
859 http_config.trace_context = Some(step_trace);
860 }
861 _ => {}
862 }
863 step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
864 }
865
866 let mut join_set = JoinSet::new();
867 let mut task_index: HashMap<Id, usize> = HashMap::new();
868 let parallel_timeout = self.guard_remaining_timeout();
869 for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
870 let provider = self.provider.clone();
871 let config = config.clone();
872 let step_log_sender = self
873 .log_sender
874 .as_ref()
875 .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
876 let handle = join_set.spawn(async move {
877 let result = match parallel_timeout {
878 Some(dur) => {
879 match tokio::time::timeout(
880 dur,
881 execute_step_config(&config, &provider, step_log_sender),
882 )
883 .await
884 {
885 Ok(r) => r,
886 Err(_elapsed) => {
887 Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
888 elapsed_secs: 0,
889 max: 0,
890 }))
891 }
892 }
893 }
894 None => execute_step_config(&config, &provider, step_log_sender).await,
895 };
896 (idx, result)
897 });
898 task_index.insert(handle.id(), idx);
899 }
900
901 let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
903 vec![None; step_records.len()];
904 let mut first_error: Option<EngineError> = None;
905
906 while let Some(join_result) = join_set.join_next().await {
907 let (idx, step_result) = match join_result {
908 Ok(r) => r,
909 Err(e) => {
910 let error_msg = format!("join error: {e}");
911 if let Some(&idx) = task_index.get(&e.id()) {
912 let (step_id, _, step_name, _) = &step_records[idx];
913 let completed_at = Utc::now();
914 error!(
915 run_id = %self.run_id,
916 step = %step_name,
917 error = %error_msg,
918 "parallel step panicked or was cancelled"
919 );
920 if let Err(store_err) = self
921 .store
922 .update_step(
923 *step_id,
924 StepUpdate {
925 status: Some(StepStatus::Failed),
926 error: Some(error_msg.clone()),
927 completed_at: Some(completed_at),
928 ..StepUpdate::default()
929 },
930 )
931 .await
932 {
933 error!(
934 run_id = %self.run_id,
935 step_id = %step_id,
936 error = %store_err,
937 "failed to persist JoinError for step"
938 );
939 }
940 indexed_results[idx] = Some(Err(error_msg.clone()));
941 }
942 if first_error.is_none() {
943 first_error = Some(EngineError::StepConfig(error_msg));
944 }
945 if fail_fast {
946 join_set.abort_all();
947 }
948 continue;
949 }
950 };
951
952 let (step_id, step_trace, step_name, step_config) = &step_records[idx];
953 let completed_at = Utc::now();
954
955 if let Err(err) = self
956 .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
957 .await
958 {
959 self.fail_step(*step_id, &err).await;
960 indexed_results[idx] = Some(Err(err.to_string()));
961 if first_error.is_none() {
962 first_error = Some(err);
963 }
964 if fail_fast {
965 join_set.abort_all();
966 }
967 continue;
968 }
969
970 match step_result {
971 Ok(output) => {
972 self.total_cost_usd += output.cost_usd;
973 self.total_duration_ms += output.duration_ms;
974
975 if matches!(step_config, StepConfig::Agent(_)) {
977 let tokens = output
978 .input_tokens
979 .unwrap_or(0)
980 .saturating_add(output.output_tokens.unwrap_or(0));
981 if tokens > 0
982 && let Err(guard_err) = self.guard_record_tokens(tokens)
983 {
984 if first_error.is_none() {
985 first_error = Some(guard_err);
986 }
987 if fail_fast {
988 join_set.abort_all();
989 }
990 }
991 }
992
993 let debug_messages_json = output.debug_messages_json();
994
995 self.store
996 .update_step(
997 *step_id,
998 StepUpdate {
999 status: Some(StepStatus::Completed),
1000 output: Some(output.output.clone()),
1001 duration_ms: Some(output.duration_ms),
1002 cost_usd: Some(output.cost_usd),
1003 input_tokens: output.input_tokens,
1004 output_tokens: output.output_tokens,
1005 completed_at: Some(completed_at),
1006 debug_messages: debug_messages_json,
1007 ..StepUpdate::default()
1008 },
1009 )
1010 .await?;
1011
1012 self.step_results.push(StepResult::from_success(
1013 *step_trace,
1014 step_name,
1015 &output,
1016 ));
1017
1018 if let Some(ref bus) = self.event_bus
1019 && matches!(step_config, StepConfig::Agent(_))
1020 {
1021 let tokens = output
1022 .input_tokens
1023 .unwrap_or(0)
1024 .saturating_add(output.output_tokens.unwrap_or(0));
1025 bus.publish(
1026 self.run_id,
1027 WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
1028 step_name: step_name.clone(),
1029 tokens,
1030 cost_usd: output.cost_usd,
1031 }),
1032 );
1033 }
1034
1035 info!(
1036 run_id = %self.run_id,
1037 step = %step_name,
1038 trace_id = %step_trace,
1039 duration_ms = output.duration_ms,
1040 "parallel step completed"
1041 );
1042
1043 indexed_results[idx] = Some(Ok(output));
1044 }
1045 Err(err) => {
1046 let err_msg = err.to_string();
1047 let debug_messages_json = extract_debug_messages_from_error(&err);
1048 let partial = extract_partial_usage_from_error(&err);
1049 let raw_response_output = extract_raw_response_from_error(&err);
1050
1051 if let Some(ref usage) = partial {
1052 if let Some(cost) = usage.cost_usd {
1053 self.total_cost_usd += cost;
1054 }
1055 if let Some(dur) = usage.duration_ms {
1056 self.total_duration_ms += dur;
1057 }
1058 }
1059
1060 if let Err(store_err) = self
1061 .store
1062 .update_step(
1063 *step_id,
1064 StepUpdate {
1065 status: Some(StepStatus::Failed),
1066 error: Some(err_msg.clone()),
1067 output: raw_response_output.clone(),
1068 completed_at: Some(completed_at),
1069 debug_messages: debug_messages_json,
1070 duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
1071 cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
1072 input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
1073 output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
1074 ..StepUpdate::default()
1075 },
1076 )
1077 .await
1078 {
1079 tracing::error!(
1080 step_id = %step_id,
1081 error = %store_err,
1082 "failed to persist parallel step failure"
1083 );
1084 }
1085
1086 let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
1087 let err_cost = partial
1088 .as_ref()
1089 .and_then(|p| p.cost_usd)
1090 .unwrap_or(Decimal::ZERO);
1091 self.step_results.push(StepResult::from_failure(
1092 *step_trace,
1093 step_name,
1094 &err_msg,
1095 err_duration,
1096 err_cost,
1097 ));
1098
1099 if step_config.allow_failure() {
1100 self.has_allowed_failure = true;
1101 info!(
1102 run_id = %self.run_id,
1103 step = %step_name,
1104 error = %err_msg,
1105 "parallel step failed but allow_failure is set, continuing"
1106 );
1107 indexed_results[idx] = Some(Ok(allowed_failure_output(
1108 &err_msg,
1109 raw_response_output,
1110 partial.as_ref(),
1111 )));
1112 } else {
1113 indexed_results[idx] = Some(Err(err_msg.clone()));
1114
1115 if first_error.is_none() {
1116 first_error = Some(err);
1117 }
1118
1119 if fail_fast {
1120 join_set.abort_all();
1121 }
1122 }
1123 }
1124 }
1125 }
1126
1127 if let Some(err) = first_error {
1128 return Err(err);
1129 }
1130
1131 self.persist_progress().await;
1132
1133 self.last_step_ids = step_records.iter().map(|(id, _, _, _)| *id).collect();
1134
1135 let results: Vec<ParallelStepResult> = step_records
1137 .iter()
1138 .enumerate()
1139 .map(|(idx, (step_id, _trace_id, name, _))| {
1140 let output = match indexed_results[idx].take() {
1141 Some(Ok(o)) => o,
1142 _ => unreachable!("all steps succeeded if no error returned"),
1143 };
1144 ParallelStepResult {
1145 name: name.clone(),
1146 output,
1147 step_id: *step_id,
1148 }
1149 })
1150 .collect();
1151
1152 Ok(results)
1153 }
1154
1155 pub async fn shell(
1178 &mut self,
1179 name: &str,
1180 config: ShellConfig,
1181 ) -> Result<StepOutput, EngineError> {
1182 self.execute_step(name, StepKind::Shell, StepConfig::Shell(config))
1183 .await
1184 }
1185
1186 pub async fn http(
1206 &mut self,
1207 name: &str,
1208 config: HttpConfig,
1209 ) -> Result<StepOutput, EngineError> {
1210 self.execute_step(name, StepKind::Http, StepConfig::Http(config))
1211 .await
1212 }
1213
1214 pub async fn agent(
1234 &mut self,
1235 name: &str,
1236 config: impl Into<AgentStepConfig>,
1237 ) -> Result<StepOutput, EngineError> {
1238 self.execute_step(name, StepKind::Agent, StepConfig::Agent(config.into()))
1239 .await
1240 }
1241
1242 pub async fn approval(
1273 &mut self,
1274 name: &str,
1275 config: ApprovalConfig,
1276 ) -> Result<(), EngineError> {
1277 let position = self.position;
1278 self.position += 1;
1279
1280 if let Some(existing) = self.replay_steps.get(&position)
1283 && existing.kind == StepKind::Approval
1284 {
1285 if existing.status.state == StepStatus::AwaitingApproval {
1286 self.store
1287 .update_step(
1288 existing.id,
1289 StepUpdate {
1290 status: Some(StepStatus::Completed),
1291 completed_at: Some(Utc::now()),
1292 ..StepUpdate::default()
1293 },
1294 )
1295 .await?;
1296 }
1297
1298 self.last_step_ids = vec![existing.id];
1299 info!(
1300 run_id = %self.run_id,
1301 step = %name,
1302 position,
1303 "approval step replayed (approved)"
1304 );
1305 return Ok(());
1306 }
1307
1308 if let Some(&granted_in) = self.granted_approvals.get(&position) {
1312 let trace_id = step_trace_id(self.run_id, name, position);
1313 let step = self
1314 .store
1315 .create_step(NewStep {
1316 run_id: self.run_id,
1317 trace_id,
1318 name: name.to_string(),
1319 kind: StepKind::Approval,
1320 position,
1321 input: Some(serde_json::to_value(&config)?),
1322 is_error_handler: false,
1323 })
1324 .await?;
1325
1326 let now = Utc::now();
1327 self.start_step(step.id, now).await?;
1328 self.store
1329 .update_step(
1330 step.id,
1331 StepUpdate {
1332 status: Some(StepStatus::Completed),
1333 output: Some(json!({"approved_in_attempt": granted_in})),
1334 completed_at: Some(now),
1335 ..StepUpdate::default()
1336 },
1337 )
1338 .await?;
1339
1340 self.last_step_ids = vec![step.id];
1341 info!(
1342 run_id = %self.run_id,
1343 step = %name,
1344 position,
1345 granted_in_attempt = granted_in,
1346 attempt = self.attempt,
1347 "approval carried over from a previous attempt"
1348 );
1349 return Ok(());
1350 }
1351
1352 let trace_id = step_trace_id(self.run_id, name, position);
1354 let step = self
1355 .store
1356 .create_step(NewStep {
1357 run_id: self.run_id,
1358 trace_id,
1359 name: name.to_string(),
1360 kind: StepKind::Approval,
1361 position,
1362 input: Some(serde_json::to_value(&config)?),
1363 is_error_handler: false,
1364 })
1365 .await?;
1366
1367 self.start_step(step.id, Utc::now()).await?;
1368
1369 self.store
1372 .update_step(
1373 step.id,
1374 StepUpdate {
1375 status: Some(StepStatus::AwaitingApproval),
1376 ..StepUpdate::default()
1377 },
1378 )
1379 .await?;
1380
1381 self.last_step_ids = vec![step.id];
1382
1383 if let Some(ref bus) = self.event_bus {
1384 bus.publish(
1385 self.run_id,
1386 WorkflowEvent::ApprovalRequired(WorkflowApprovalRequiredEvent {
1387 step_name: name.to_string(),
1388 step_index: position,
1389 approval_id: step.id,
1390 }),
1391 );
1392 }
1393
1394 Err(EngineError::ApprovalRequired {
1395 run_id: self.run_id,
1396 step_id: step.id,
1397 message: config.message().to_string(),
1398 })
1399 }
1400
1401 pub async fn decision(
1414 &mut self,
1415 name: &str,
1416 config: DecisionConfig,
1417 ) -> Result<DecisionOutput, EngineError> {
1418 if let Some(output) = self.decision_replay(name, &config).await? {
1419 return Ok(output);
1420 }
1421 self.decision_execute(name, config).await
1422 }
1423
1424 pub async fn skip(&mut self, name: &str, reason: &str) -> Result<(), EngineError> {
1453 let position = self.position;
1454 self.position += 1;
1455
1456 let trace_id = step_trace_id(self.run_id, name, position);
1457 let step = self
1458 .store
1459 .create_step(NewStep {
1460 run_id: self.run_id,
1461 trace_id,
1462 name: name.to_string(),
1463 kind: StepKind::Custom("skip".to_string()),
1464 position,
1465 input: None,
1466 is_error_handler: false,
1467 })
1468 .await?;
1469
1470 if !self.last_step_ids.is_empty() {
1471 let deps: Vec<NewStepDependency> = self
1472 .last_step_ids
1473 .iter()
1474 .map(|&depends_on| NewStepDependency {
1475 step_id: step.id,
1476 depends_on,
1477 })
1478 .collect();
1479 self.store.create_step_dependencies(deps).await?;
1480 }
1481
1482 let now = Utc::now();
1483 self.store
1484 .update_step(
1485 step.id,
1486 StepUpdate {
1487 status: Some(StepStatus::Skipped),
1488 output: Some(serde_json::json!({"reason": reason})),
1489 completed_at: Some(now),
1490 ..StepUpdate::default()
1491 },
1492 )
1493 .await?;
1494
1495 self.last_step_ids = vec![step.id];
1496
1497 info!(
1498 run_id = %self.run_id,
1499 step = %name,
1500 reason,
1501 "step skipped"
1502 );
1503
1504 Ok(())
1505 }
1506
1507 pub async fn operation(
1546 &mut self,
1547 name: &str,
1548 op: &dyn Operation,
1549 ) -> Result<StepOutput, EngineError> {
1550 let kind = StepKind::Custom(op.kind().to_string());
1551 let position = self.position;
1552 self.position += 1;
1553
1554 let trace_id = step_trace_id(self.run_id, name, position);
1555 let step = self
1556 .store
1557 .create_step(NewStep {
1558 run_id: self.run_id,
1559 trace_id,
1560 name: name.to_string(),
1561 kind,
1562 position,
1563 input: op.input(),
1564 is_error_handler: false,
1565 })
1566 .await?;
1567
1568 self.start_step(step.id, Utc::now()).await?;
1569
1570 let start = Instant::now();
1571
1572 let op_ctx = self.ensure_operation_ctx();
1573
1574 match op.execute(op_ctx).await {
1575 Ok(output_value) => {
1576 let duration_ms = start.elapsed().as_millis() as u64;
1577 self.total_duration_ms += duration_ms;
1578
1579 let completed_at = Utc::now();
1580 self.store
1581 .update_step(
1582 step.id,
1583 StepUpdate {
1584 status: Some(StepStatus::Completed),
1585 output: Some(output_value.clone()),
1586 duration_ms: Some(duration_ms),
1587 cost_usd: Some(Decimal::ZERO),
1588 completed_at: Some(completed_at),
1589 ..StepUpdate::default()
1590 },
1591 )
1592 .await?;
1593
1594 info!(
1595 run_id = %self.run_id,
1596 step = %name,
1597 kind = op.kind(),
1598 duration_ms,
1599 "operation step completed"
1600 );
1601
1602 self.last_step_ids = vec![step.id];
1603
1604 Ok(StepOutput {
1605 output: output_value,
1606 duration_ms,
1607 cost_usd: Decimal::ZERO,
1608 input_tokens: None,
1609 output_tokens: None,
1610 model: None,
1611 debug_messages: None,
1612 })
1613 }
1614 Err(err) => {
1615 let completed_at = Utc::now();
1616 let engine_err = EngineError::Operation(err);
1617 if let Err(store_err) = self
1618 .store
1619 .update_step(
1620 step.id,
1621 StepUpdate {
1622 status: Some(StepStatus::Failed),
1623 error: Some(engine_err.to_string()),
1624 completed_at: Some(completed_at),
1625 ..StepUpdate::default()
1626 },
1627 )
1628 .await
1629 {
1630 error!(step_id = %step.id, error = %store_err, "failed to persist step failure");
1631 }
1632
1633 Err(engine_err)
1634 }
1635 }
1636 }
1637
1638 pub async fn workflow(
1665 &mut self,
1666 handler: &dyn WorkflowHandler,
1667 payload: Value,
1668 ) -> Result<StepOutput, EngineError> {
1669 if let (Some(guard_config), Some(guard_state)) = (&self.guard_config, &self.guard_state) {
1671 let state = guard_state
1672 .lock()
1673 .map_err(|_| WorkflowRejection::GuardUnavailable)?;
1674 state.check(guard_config, handler.name())?;
1675 }
1676
1677 let config = WorkflowStepConfig::new(handler.name(), payload);
1678 let position = self.position;
1679 self.position += 1;
1680
1681 let trace_id = step_trace_id(self.run_id, &config.workflow_name, position);
1682 let step = self
1683 .store
1684 .create_step(NewStep {
1685 run_id: self.run_id,
1686 trace_id,
1687 name: config.workflow_name.clone(),
1688 kind: StepKind::Workflow,
1689 position,
1690 input: Some(serde_json::to_value(&config)?),
1691 is_error_handler: false,
1692 })
1693 .await?;
1694
1695 self.start_step(step.id, Utc::now()).await?;
1696
1697 if let Some(guard_state) = &self.guard_state {
1699 let mut state = guard_state
1700 .lock()
1701 .map_err(|_| WorkflowRejection::GuardUnavailable)?;
1702 state.record_invocation(handler.name());
1703 }
1704
1705 match self.execute_child_workflow(&config).await {
1706 Ok((output, child_had_allowed_failure)) => {
1707 self.total_cost_usd += output.cost_usd;
1708 self.total_duration_ms += output.duration_ms;
1709 if child_had_allowed_failure {
1710 self.has_allowed_failure = true;
1711 }
1712
1713 let completed_at = Utc::now();
1714 self.store
1715 .update_step(
1716 step.id,
1717 StepUpdate {
1718 status: Some(StepStatus::Completed),
1719 output: Some(output.output.clone()),
1720 duration_ms: Some(output.duration_ms),
1721 cost_usd: Some(output.cost_usd),
1722 completed_at: Some(completed_at),
1723 ..StepUpdate::default()
1724 },
1725 )
1726 .await?;
1727
1728 info!(
1729 run_id = %self.run_id,
1730 child_workflow = %config.workflow_name,
1731 duration_ms = output.duration_ms,
1732 "workflow step completed"
1733 );
1734
1735 self.last_step_ids = vec![step.id];
1736
1737 self.guard_record_return();
1738 Ok(output)
1739 }
1740 Err(err) => {
1741 let completed_at = Utc::now();
1742 if let Err(store_err) = self
1743 .store
1744 .update_step(
1745 step.id,
1746 StepUpdate {
1747 status: Some(StepStatus::Failed),
1748 error: Some(err.to_string()),
1749 completed_at: Some(completed_at),
1750 ..StepUpdate::default()
1751 },
1752 )
1753 .await
1754 {
1755 error!(step_id = %step.id, error = %store_err, "failed to persist step failure");
1756 }
1757
1758 self.guard_record_return();
1759 Err(err)
1760 }
1761 }
1762 }
1763
1764 fn guard_record_return(&self) {
1769 if let Some(guard_state) = &self.guard_state {
1770 match guard_state.lock() {
1771 Ok(mut state) => state.record_return(),
1772 Err(_) => {
1773 error!(
1774 run_id = %self.run_id,
1775 "guard state mutex poisoned in record_return"
1776 );
1777 }
1778 }
1779 }
1780 }
1781
1782 async fn execute_with_guard_timeout(
1786 &self,
1787 config: &StepConfig,
1788 step_log_sender: Option<StepLogSender>,
1789 ) -> Result<StepOutput, EngineError> {
1790 let remaining = self.guard_remaining_timeout();
1791 match remaining {
1792 Some(dur) => {
1793 use tokio::time::timeout;
1794 match timeout(
1795 dur,
1796 execute_step_config(config, &self.provider, step_log_sender),
1797 )
1798 .await
1799 {
1800 Ok(result) => result,
1801 Err(_elapsed) => {
1802 let config_secs = self
1803 .guard_config
1804 .as_ref()
1805 .map_or(0, |c| c.workflow_timeout_secs);
1806 Err(WorkflowRejection::WorkflowTimeout {
1807 elapsed_secs: config_secs,
1808 max: config_secs,
1809 }
1810 .into())
1811 }
1812 }
1813 }
1814 None => execute_step_config(config, &self.provider, step_log_sender).await,
1815 }
1816 }
1817
1818 fn guard_remaining_timeout(&self) -> Option<std::time::Duration> {
1820 let config = self.guard_config.as_ref()?;
1821 let guard_state = self.guard_state.as_ref()?;
1822 let state = guard_state.lock().ok()?;
1823 let elapsed = state.elapsed_secs();
1824 let max = config.workflow_timeout_secs;
1825 if elapsed >= max {
1826 Some(std::time::Duration::ZERO)
1827 } else {
1828 Some(std::time::Duration::from_secs(max - elapsed))
1829 }
1830 }
1831
1832 fn check_guard_timeout(&self) -> Result<(), EngineError> {
1834 if let (Some(config), Some(guard_state)) = (&self.guard_config, &self.guard_state) {
1835 let state = guard_state
1836 .lock()
1837 .map_err(|_| WorkflowRejection::GuardUnavailable)?;
1838 let elapsed = state.elapsed_secs();
1839 if elapsed >= config.workflow_timeout_secs {
1840 return Err(WorkflowRejection::WorkflowTimeout {
1841 elapsed_secs: elapsed,
1842 max: config.workflow_timeout_secs,
1843 }
1844 .into());
1845 }
1846 }
1847 Ok(())
1848 }
1849
1850 fn guard_record_tokens(&self, tokens: u64) -> Result<(), EngineError> {
1852 if let (Some(config), Some(guard_state)) = (&self.guard_config, &self.guard_state) {
1853 let mut state = guard_state
1854 .lock()
1855 .map_err(|_| WorkflowRejection::GuardUnavailable)?;
1856 state.record_tokens(config, tokens)?;
1857 }
1858 Ok(())
1859 }
1860
1861 async fn execute_child_workflow(
1864 &self,
1865 config: &WorkflowStepConfig,
1866 ) -> Result<(StepOutput, bool), EngineError> {
1867 let resolver = self.handler_resolver.as_ref().ok_or_else(|| {
1868 EngineError::InvalidWorkflow(
1869 "sub-workflow requires a handler resolver (use Engine to execute)".to_string(),
1870 )
1871 })?;
1872
1873 let handler = resolver(&config.workflow_name).ok_or_else(|| {
1874 EngineError::InvalidWorkflow(format!("no handler registered: {}", config.workflow_name))
1875 })?;
1876
1877 let parent = self.store.get_run(self.run_id).await?;
1880 let (parent_labels, parent_author) =
1881 parent.map(|r| (r.labels, r.created_by)).unwrap_or_default();
1882
1883 let child_run = self
1884 .store
1885 .create_run(NewRun {
1886 workflow_name: config.workflow_name.clone(),
1887 trigger: TriggerKind::Workflow,
1888 payload: config.payload.clone(),
1889 max_retries: 0,
1890 handler_version: None,
1891 labels: parent_labels,
1892 scheduled_at: None,
1893 created_by: parent_author,
1894 idempotency_key: None,
1895 max_cost_usd: self.max_cost_usd,
1897 })
1898 .await?
1899 .into_run();
1900
1901 let child_run_id = child_run.id;
1902 info!(
1903 parent_run_id = %self.run_id,
1904 child_run_id = %child_run_id,
1905 workflow = %config.workflow_name,
1906 "child run created"
1907 );
1908
1909 self.store
1910 .update_run_status(child_run_id, RunStatus::Running)
1911 .await?;
1912
1913 let run_start = Instant::now();
1914 let mut child_ctx = WorkflowContext {
1915 run_id: child_run_id,
1916 workflow_name: config.workflow_name.clone(),
1917 store: self.store.clone(),
1918 provider: self.provider.clone(),
1919 decision_provider: self.decision_provider.clone(),
1920 handler_resolver: self.handler_resolver.clone(),
1921 position: 0,
1922 last_step_ids: Vec::new(),
1923 total_cost_usd: Decimal::ZERO,
1924 total_duration_ms: 0,
1925 max_cost_usd: self.max_cost_usd,
1926 inherited_cost_usd: self.charged_cost_usd(),
1929 replay_steps: HashMap::new(),
1930 granted_approvals: HashMap::new(),
1931 attempt: 1,
1933 carried_duration_ms: 0,
1934 log_sender: self.log_sender.clone(),
1935 artifact_sink: self.artifact_sink.clone(),
1938 has_allowed_failure: false,
1939 error_handlers: Vec::new(),
1940 guard_state: self.guard_state.clone(),
1941 guard_config: self.guard_config.clone(),
1942 step_results: Vec::new(),
1943 event_bus: self.event_bus.clone(),
1944 trace_context: self.trace_context.child(),
1945 operation_ctx: None,
1946 };
1947
1948 let result = handler.execute(&mut child_ctx).await;
1949 let total_duration = run_start.elapsed().as_millis() as u64;
1950 let completed_at = Utc::now();
1951
1952 match result {
1953 Ok(()) => {
1954 let child_status = if child_ctx.has_allowed_failure {
1955 RunStatus::Warning
1956 } else {
1957 RunStatus::Completed
1958 };
1959 self.store
1960 .update_run(
1961 child_run_id,
1962 RunUpdate {
1963 status: Some(child_status),
1964 cost_usd: Some(child_ctx.total_cost_usd),
1965 duration_ms: Some(total_duration),
1966 completed_at: Some(completed_at),
1967 ..RunUpdate::default()
1968 },
1969 )
1970 .await?;
1971
1972 let child_had_allowed_failure = child_ctx.has_allowed_failure;
1973 Ok((
1974 StepOutput {
1975 output: serde_json::json!({
1976 "run_id": child_run_id,
1977 "workflow_name": config.workflow_name,
1978 "status": child_status,
1979 "cost_usd": child_ctx.total_cost_usd,
1980 "duration_ms": total_duration,
1981 }),
1982 duration_ms: total_duration,
1983 cost_usd: child_ctx.total_cost_usd,
1984 input_tokens: None,
1985 output_tokens: None,
1986 model: None,
1987 debug_messages: None,
1988 },
1989 child_had_allowed_failure,
1990 ))
1991 }
1992 Err(err) => {
1993 if let Err(store_err) = self
1994 .store
1995 .update_run(
1996 child_run_id,
1997 RunUpdate {
1998 status: Some(RunStatus::Failed),
1999 error: Some(err.to_string()),
2000 cost_usd: Some(child_ctx.total_cost_usd),
2001 duration_ms: Some(total_duration),
2002 completed_at: Some(completed_at),
2003 ..RunUpdate::default()
2004 },
2005 )
2006 .await
2007 {
2008 error!(
2009 child_run_id = %child_run_id,
2010 store_error = %store_err,
2011 "failed to persist child run failure"
2012 );
2013 }
2014
2015 Err(err)
2016 }
2017 }
2018 }
2019
2020 fn try_replay_step(&mut self, position: u32) -> Option<StepOutput> {
2025 let step = self.replay_steps.get(&position)?;
2026 if step.status.state != StepStatus::Completed {
2027 return None;
2028 }
2029 let output = StepOutput {
2030 output: step.output.clone().unwrap_or(Value::Null),
2031 duration_ms: step.duration_ms,
2032 cost_usd: step.cost_usd,
2033 input_tokens: step.input_tokens,
2034 output_tokens: step.output_tokens,
2035 model: None,
2036 debug_messages: None,
2037 };
2038 self.total_cost_usd += output.cost_usd;
2039 self.total_duration_ms += output.duration_ms;
2040 self.last_step_ids = vec![step.id];
2041 info!(
2042 run_id = %self.run_id,
2043 step = %step.name,
2044 position,
2045 "step replayed from previous execution"
2046 );
2047 Some(output)
2048 }
2049
2050 #[tracing::instrument(
2052 name = "context.execute_step",
2053 skip_all,
2054 fields(
2055 run_id = %self.run_id,
2056 step.name = %name,
2057 step.kind,
2058 step.position = self.position,
2059 step.trace_id,
2060 )
2061 )]
2062 pub(crate) async fn execute_step(
2063 &mut self,
2064 name: &str,
2065 kind: StepKind,
2066 config: StepConfig,
2067 ) -> Result<StepOutput, EngineError> {
2068 let kind_str: &'static str = match kind {
2069 StepKind::Shell => "shell",
2070 StepKind::Http => "http",
2071 StepKind::Agent => "agent",
2072 StepKind::Workflow => "workflow",
2073 StepKind::Approval => "approval",
2074 StepKind::Decision => "decision",
2075 StepKind::Custom(_) => "custom",
2076 };
2077 Span::current().record("step.kind", kind_str);
2078
2079 self.check_guard_timeout()?;
2081
2082 let position = self.position;
2083 self.position += 1;
2084
2085 if let Some(output) = self.try_replay_step(position) {
2087 return Ok(output);
2088 }
2089
2090 if let StepConfig::Agent(ref agent_config) = config {
2093 self.check_run_budget(step_budget_usd(agent_config.max_budget_usd))?;
2094 }
2095
2096 let trace_id = step_trace_id(self.run_id, name, position);
2098 Span::current().record("step.trace_id", trace_id.to_string().as_str());
2099 let step = self
2100 .store
2101 .create_step(NewStep {
2102 run_id: self.run_id,
2103 trace_id,
2104 name: name.to_string(),
2105 kind,
2106 position,
2107 input: Some(serde_json::to_value(&config)?),
2108 is_error_handler: false,
2109 })
2110 .await?;
2111
2112 self.start_step(step.id, Utc::now()).await?;
2113
2114 if let Some(ref bus) = self.event_bus {
2115 bus.publish(
2116 self.run_id,
2117 WorkflowEvent::StepStarted(WorkflowStepStartedEvent {
2118 step_name: name.to_string(),
2119 step_index: position,
2120 timestamp: Utc::now(),
2121 }),
2122 );
2123 }
2124
2125 if let Err(err) = self.prepare_step_inputs(&config, position).await {
2128 self.fail_step(step.id, &err).await;
2129 if config.allow_failure() {
2130 self.has_allowed_failure = true;
2131 self.last_step_ids = vec![step.id];
2132 info!(
2133 run_id = %self.run_id,
2134 step = %name,
2135 error = %err,
2136 "step input preparation failed but allow_failure is set, continuing"
2137 );
2138 return Ok(StepOutput {
2139 output: json!({"error": err.to_string()}),
2140 duration_ms: 0,
2141 cost_usd: Decimal::ZERO,
2142 input_tokens: None,
2143 output_tokens: None,
2144 model: None,
2145 debug_messages: None,
2146 });
2147 }
2148 return Err(err);
2149 }
2150
2151 let mut config = config;
2152 let step_trace = self.trace_context.child();
2153 match config {
2154 StepConfig::Agent(ref mut agent_config) => {
2155 agent_config.trace_context = Some(step_trace);
2156 }
2157 StepConfig::Http(ref mut http_config) => {
2158 http_config.trace_context = Some(step_trace);
2159 }
2160 _ => {}
2161 }
2162
2163 let step_log_sender = self
2164 .log_sender
2165 .as_ref()
2166 .map(|s| StepLogSender::new(s.clone(), self.run_id, step.id, name.to_string()));
2167
2168 let execution = self
2169 .execute_with_guard_timeout(&config, step_log_sender)
2170 .await;
2171
2172 let execution = self
2173 .retry_step_if_configured(name, kind_str, &config, step.id, execution)
2174 .await;
2175
2176 if let Err(err) = self
2177 .store_step_outputs(&config, step.id, name, execution.is_ok())
2178 .await
2179 {
2180 self.fail_step(step.id, &err).await;
2181 return Err(err);
2182 }
2183
2184 match execution {
2185 Ok(output) => {
2186 self.total_cost_usd += output.cost_usd;
2187 self.total_duration_ms += output.duration_ms;
2188
2189 if matches!(config, StepConfig::Agent(_)) {
2191 let tokens = output
2192 .input_tokens
2193 .unwrap_or(0)
2194 .saturating_add(output.output_tokens.unwrap_or(0));
2195 if tokens > 0 {
2196 self.guard_record_tokens(tokens)?;
2197 }
2198 }
2199
2200 let debug_messages_json = output.debug_messages_json();
2201
2202 let completed_at = Utc::now();
2203 self.store
2204 .update_step(
2205 step.id,
2206 StepUpdate {
2207 status: Some(StepStatus::Completed),
2208 output: Some(output.output.clone()),
2209 duration_ms: Some(output.duration_ms),
2210 cost_usd: Some(output.cost_usd),
2211 input_tokens: output.input_tokens,
2212 output_tokens: output.output_tokens,
2213 completed_at: Some(completed_at),
2214 debug_messages: debug_messages_json,
2215 ..StepUpdate::default()
2216 },
2217 )
2218 .await?;
2219
2220 self.step_results
2221 .push(StepResult::from_success(trace_id, name, &output));
2222 self.persist_progress().await;
2223
2224 info!(
2225 run_id = %self.run_id,
2226 step = %name,
2227 trace_id = %trace_id,
2228 duration_ms = output.duration_ms,
2229 "step completed"
2230 );
2231
2232 if let Some(ref bus) = self.event_bus {
2233 bus.publish(
2234 self.run_id,
2235 WorkflowEvent::StepCompleted(WorkflowStepCompletedEvent {
2236 step_name: name.to_string(),
2237 step_index: position,
2238 duration_ms: output.duration_ms,
2239 output_summary: None,
2240 }),
2241 );
2242
2243 if matches!(config, StepConfig::Agent(_)) {
2244 let tokens = output
2245 .input_tokens
2246 .unwrap_or(0)
2247 .saturating_add(output.output_tokens.unwrap_or(0));
2248 bus.publish(
2249 self.run_id,
2250 WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
2251 step_name: name.to_string(),
2252 tokens,
2253 cost_usd: output.cost_usd,
2254 }),
2255 );
2256 }
2257 }
2258
2259 self.last_step_ids = vec![step.id];
2260
2261 Ok(output)
2262 }
2263 Err(err) => {
2264 let completed_at = Utc::now();
2265 let debug_messages_json = extract_debug_messages_from_error(&err);
2266 let partial = extract_partial_usage_from_error(&err);
2267 let raw_response_output = extract_raw_response_from_error(&err);
2268
2269 if let Some(ref usage) = partial {
2270 if let Some(cost) = usage.cost_usd {
2271 self.total_cost_usd += cost;
2272 }
2273 if let Some(dur) = usage.duration_ms {
2274 self.total_duration_ms += dur;
2275 }
2276 }
2277
2278 if let Err(store_err) = self
2279 .store
2280 .update_step(
2281 step.id,
2282 StepUpdate {
2283 status: Some(StepStatus::Failed),
2284 error: Some(err.to_string()),
2285 output: raw_response_output.clone(),
2286 completed_at: Some(completed_at),
2287 debug_messages: debug_messages_json,
2288 duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
2289 cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
2290 input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
2291 output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
2292 ..StepUpdate::default()
2293 },
2294 )
2295 .await
2296 {
2297 tracing::error!(step_id = %step.id, error = %store_err, "failed to persist step failure");
2298 }
2299
2300 let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
2301 let err_cost = partial
2302 .as_ref()
2303 .and_then(|p| p.cost_usd)
2304 .unwrap_or(Decimal::ZERO);
2305 self.step_results.push(StepResult::from_failure(
2306 trace_id,
2307 name,
2308 &err.to_string(),
2309 err_duration,
2310 err_cost,
2311 ));
2312 self.persist_progress().await;
2313
2314 if let Some(ref bus) = self.event_bus {
2315 bus.publish(
2316 self.run_id,
2317 WorkflowEvent::StepFailed(WorkflowStepFailedEvent {
2318 step_name: name.to_string(),
2319 step_index: position,
2320 error: err.to_string(),
2321 duration_ms: err_duration,
2322 }),
2323 );
2324 }
2325
2326 self.fire_error_handlers(name, &err.to_string(), err_duration)
2327 .await;
2328
2329 if config.allow_failure() {
2330 self.has_allowed_failure = true;
2331 self.last_step_ids = vec![step.id];
2332 info!(
2333 run_id = %self.run_id,
2334 step = %name,
2335 error = %err,
2336 "step failed but allow_failure is set, continuing"
2337 );
2338 Ok(allowed_failure_output(
2339 &err.to_string(),
2340 raw_response_output,
2341 partial.as_ref(),
2342 ))
2343 } else {
2344 Err(err)
2345 }
2346 }
2347 }
2348 }
2349
2350 #[cfg_attr(not(feature = "prometheus"), allow(unused_variables))]
2356 async fn retry_step_if_configured(
2357 &self,
2358 name: &str,
2359 kind_str: &str,
2360 config: &StepConfig,
2361 step_id: Uuid,
2362 first_result: Result<StepOutput, EngineError>,
2363 ) -> Result<StepOutput, EngineError> {
2364 let policy = match config.retry() {
2365 Some(p) => p,
2366 None => return first_result,
2367 };
2368
2369 let mut last_result = match first_result {
2370 Ok(output) => return Ok(output),
2371 Err(err) if !is_step_retryable(&err) => return Err(err),
2372 Err(err) => Err(err),
2373 };
2374
2375 let step_log_sender = self
2376 .log_sender
2377 .as_ref()
2378 .map(|s| StepLogSender::new(s.clone(), self.run_id, step_id, name.to_string()));
2379
2380 for attempt in 0..policy.max_retries() {
2381 if let StepConfig::Agent(agent_config) = config {
2382 self.check_run_budget(step_budget_usd(agent_config.max_budget_usd))?;
2383 }
2384
2385 let delay = policy.delay_for_attempt(attempt);
2386 info!(
2387 run_id = %self.run_id,
2388 step = %name,
2389 attempt = attempt + 1,
2390 max_retries = policy.max_retries(),
2391 delay_ms = delay.as_millis() as u64,
2392 "retrying step after transient failure"
2393 );
2394 tokio::time::sleep(delay).await;
2395
2396 record_retry_metric(kind_str, "retry");
2397
2398 match execute_step_config(config, &self.provider, step_log_sender.clone()).await {
2399 Ok(output) => return Ok(output),
2400 Err(err) if !is_step_retryable(&err) => return Err(err),
2401 err => last_result = err,
2402 }
2403 }
2404
2405 record_retry_metric(kind_str, "exhausted");
2406
2407 info!(
2408 run_id = %self.run_id,
2409 step = %name,
2410 max_retries = policy.max_retries(),
2411 "step retries exhausted"
2412 );
2413
2414 last_result
2415 }
2416
2417 pub(crate) async fn start_step(
2422 &self,
2423 step_id: Uuid,
2424 now: DateTime<Utc>,
2425 ) -> Result<(), EngineError> {
2426 if !self.last_step_ids.is_empty() {
2427 let deps: Vec<NewStepDependency> = self
2428 .last_step_ids
2429 .iter()
2430 .map(|&depends_on| NewStepDependency {
2431 step_id,
2432 depends_on,
2433 })
2434 .collect();
2435 self.store.create_step_dependencies(deps).await?;
2436 }
2437
2438 self.store
2439 .update_step(
2440 step_id,
2441 StepUpdate {
2442 status: Some(StepStatus::Running),
2443 started_at: Some(now),
2444 ..StepUpdate::default()
2445 },
2446 )
2447 .await?;
2448
2449 Ok(())
2450 }
2451
2452 async fn fail_step(&self, step_id: Uuid, err: &EngineError) {
2459 if let Err(store_err) = self
2460 .store
2461 .update_step(
2462 step_id,
2463 StepUpdate {
2464 status: Some(StepStatus::Failed),
2465 error: Some(err.to_string()),
2466 completed_at: Some(Utc::now()),
2467 ..StepUpdate::default()
2468 },
2469 )
2470 .await
2471 {
2472 error!(
2473 step_id = %step_id,
2474 error = %store_err,
2475 "failed to persist step failure"
2476 );
2477 }
2478 }
2479
2480 pub fn store(&self) -> &Arc<dyn Store> {
2482 &self.store
2483 }
2484
2485 pub(crate) fn next_position(&mut self) -> u32 {
2487 let pos = self.position;
2488 self.position += 1;
2489 pos
2490 }
2491
2492 pub(crate) fn replay_steps(&self) -> &HashMap<u32, Step> {
2494 &self.replay_steps
2495 }
2496
2497 pub(crate) fn set_last_step_ids(&mut self, ids: Vec<Uuid>) {
2499 self.last_step_ids = ids;
2500 }
2501
2502 pub async fn payload(&self) -> Result<Value, EngineError> {
2510 let run = self
2511 .store
2512 .get_run(self.run_id)
2513 .await?
2514 .ok_or(EngineError::Store(
2515 ironflow_store::error::StoreError::RunNotFound(self.run_id),
2516 ))?;
2517 Ok(run.payload)
2518 }
2519
2520 pub async fn input<T: serde::de::DeserializeOwned>(&self) -> Result<T, EngineError> {
2548 let payload = self.payload().await?;
2549 serde_json::from_value(payload).map_err(EngineError::Serialization)
2550 }
2551
2552 pub fn on_error(&mut self, name: &str, config: impl Into<StepConfig>) {
2575 self.error_handlers.push(OnErrorHandler {
2576 name: name.to_string(),
2577 config: config.into(),
2578 });
2579 }
2580
2581 pub fn clear_error_handlers(&mut self) {
2600 self.error_handlers.clear();
2601 }
2602
2603 async fn fire_error_handlers(
2609 &mut self,
2610 failed_step_name: &str,
2611 error_msg: &str,
2612 duration_ms: u64,
2613 ) {
2614 let handlers = std::mem::take(&mut self.error_handlers);
2615 if handlers.is_empty() {
2616 return;
2617 }
2618
2619 let error_context = json!({
2620 "failed_step": failed_step_name,
2621 "error": error_msg,
2622 "duration_ms": duration_ms,
2623 });
2624
2625 for handler in handlers {
2626 let mut config = handler.config.clone();
2627 inject_error_context(&mut config, failed_step_name, error_msg, duration_ms);
2628
2629 let position = self.position;
2630 self.position += 1;
2631
2632 let trace_id = step_trace_id(self.run_id, &handler.name, position);
2633 let step = match self
2634 .store
2635 .create_step(NewStep {
2636 run_id: self.run_id,
2637 trace_id,
2638 name: handler.name.clone(),
2639 kind: config.kind(),
2640 position,
2641 input: Some(error_context.clone()),
2642 is_error_handler: true,
2643 })
2644 .await
2645 {
2646 Ok(step) => step,
2647 Err(err) => {
2648 warn!(
2649 run_id = %self.run_id,
2650 handler = %handler.name,
2651 error = %err,
2652 "failed to create error handler step"
2653 );
2654 continue;
2655 }
2656 };
2657
2658 if let Err(err) = self.start_step(step.id, Utc::now()).await {
2659 warn!(
2660 run_id = %self.run_id,
2661 handler = %handler.name,
2662 error = %err,
2663 "failed to start error handler step"
2664 );
2665 continue;
2666 }
2667
2668 let step_log_sender = self
2669 .log_sender
2670 .as_ref()
2671 .map(|s| StepLogSender::new(s.clone(), self.run_id, step.id, handler.name.clone()));
2672
2673 let start = Instant::now();
2674 let result = execute_step_config(&config, &self.provider, step_log_sender).await;
2675 let handler_duration = start.elapsed().as_millis() as u64;
2676 let completed_at = Utc::now();
2677
2678 match result {
2679 Ok(output) => {
2680 if let Err(store_err) = self
2681 .store
2682 .update_step(
2683 step.id,
2684 StepUpdate {
2685 status: Some(StepStatus::Completed),
2686 output: Some(output.output),
2687 duration_ms: Some(handler_duration),
2688 cost_usd: Some(output.cost_usd),
2689 completed_at: Some(completed_at),
2690 ..StepUpdate::default()
2691 },
2692 )
2693 .await
2694 {
2695 warn!(
2696 run_id = %self.run_id,
2697 handler = %handler.name,
2698 error = %store_err,
2699 "failed to persist error handler completion"
2700 );
2701 }
2702
2703 info!(
2704 run_id = %self.run_id,
2705 handler = %handler.name,
2706 duration_ms = handler_duration,
2707 "error handler completed"
2708 );
2709 }
2710 Err(err) => {
2711 if let Err(store_err) = self
2712 .store
2713 .update_step(
2714 step.id,
2715 StepUpdate {
2716 status: Some(StepStatus::Failed),
2717 error: Some(err.to_string()),
2718 duration_ms: Some(handler_duration),
2719 completed_at: Some(completed_at),
2720 ..StepUpdate::default()
2721 },
2722 )
2723 .await
2724 {
2725 warn!(
2726 run_id = %self.run_id,
2727 handler = %handler.name,
2728 error = %store_err,
2729 "failed to persist error handler failure"
2730 );
2731 }
2732
2733 warn!(
2734 run_id = %self.run_id,
2735 handler = %handler.name,
2736 error = %err,
2737 "error handler failed (original error preserved)"
2738 );
2739 }
2740 }
2741 }
2742 }
2743}
2744
2745fn inject_error_context(
2747 config: &mut StepConfig,
2748 failed_step: &str,
2749 error_msg: &str,
2750 duration_ms: u64,
2751) {
2752 match config {
2753 StepConfig::Shell(shell) => {
2754 shell
2755 .env
2756 .push(("IRONFLOW_ERROR_STEP".to_string(), failed_step.to_string()));
2757 shell
2758 .env
2759 .push(("IRONFLOW_ERROR_MESSAGE".to_string(), error_msg.to_string()));
2760 shell.env.push((
2761 "IRONFLOW_ERROR_DURATION_MS".to_string(),
2762 duration_ms.to_string(),
2763 ));
2764 }
2765 StepConfig::Agent(agent) => {
2766 agent.prompt = format!(
2767 "[Error Context]\nStep \"{}\" failed after {}ms:\n{}\n\n{}",
2768 failed_step, duration_ms, error_msg, agent.prompt
2769 );
2770 }
2771 StepConfig::Http(http) => {
2772 http.headers
2773 .push(("X-Ironflow-Error-Step".to_string(), failed_step.to_string()));
2774 http.headers.push((
2775 "X-Ironflow-Error-Message".to_string(),
2776 error_msg.to_string(),
2777 ));
2778 }
2779 StepConfig::Workflow(_)
2780 | StepConfig::Approval(_)
2781 | StepConfig::Decision(_)
2782 | StepConfig::Delay(_) => {}
2783 }
2784}
2785
2786#[cfg(feature = "prometheus")]
2787fn record_retry_metric(kind: &str, outcome: &str) {
2788 use ironflow_core::metric_names::STEP_RETRIES_TOTAL;
2789 use metrics::counter;
2790 counter!(STEP_RETRIES_TOTAL, "kind" => kind.to_string(), "outcome" => outcome.to_string())
2791 .increment(1);
2792}
2793
2794#[cfg(not(feature = "prometheus"))]
2795fn record_retry_metric(_kind: &str, _outcome: &str) {}
2796
2797fn is_step_retryable(err: &EngineError) -> bool {
2801 use ironflow_core::error::{AgentError, OperationError};
2802
2803 match err {
2804 EngineError::Operation(op) => match op {
2805 OperationError::Agent(AgentError::PromptTooLarge { .. }) => false,
2806 OperationError::Agent(AgentError::BudgetExceeded { .. }) => false,
2807 OperationError::Deserialize { .. } => false,
2808 OperationError::Http {
2809 status: Some(code), ..
2810 } if (400..500).contains(code) && *code != 429 => false,
2811 _ => true,
2812 },
2813 _ => false,
2814 }
2815}
2816
2817fn allowed_failure_output(
2818 error_msg: &str,
2819 raw_response: Option<Value>,
2820 partial: Option<&StepPartialUsage>,
2821) -> StepOutput {
2822 StepOutput {
2823 output: raw_response.unwrap_or_else(|| json!({"error": error_msg})),
2824 duration_ms: partial.and_then(|p| p.duration_ms).unwrap_or(0),
2825 cost_usd: partial.and_then(|p| p.cost_usd).unwrap_or(Decimal::ZERO),
2826 input_tokens: partial.and_then(|p| p.input_tokens),
2827 output_tokens: partial.and_then(|p| p.output_tokens),
2828 model: None,
2829 debug_messages: None,
2830 }
2831}
2832
2833impl fmt::Debug for WorkflowContext {
2834 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
2835 f.debug_struct("WorkflowContext")
2836 .field("run_id", &self.run_id)
2837 .field("position", &self.position)
2838 .field("total_cost_usd", &self.total_cost_usd)
2839 .field("inherited_cost_usd", &self.inherited_cost_usd)
2840 .field("max_cost_usd", &self.max_cost_usd)
2841 .finish_non_exhaustive()
2842 }
2843}
2844
2845fn extract_debug_messages_from_error(err: &EngineError) -> Option<Value> {
2848 if let EngineError::Operation(OperationError::Agent(AgentError::SchemaValidation {
2849 debug_messages,
2850 ..
2851 })) = err
2852 && !debug_messages.is_empty()
2853 {
2854 return serde_json::to_value(debug_messages).ok();
2855 }
2856 None
2857}
2858
2859struct StepPartialUsage {
2865 cost_usd: Option<Decimal>,
2866 duration_ms: Option<u64>,
2867 input_tokens: Option<u64>,
2868 output_tokens: Option<u64>,
2869}
2870
2871fn extract_raw_response_from_error(err: &EngineError) -> Option<Value> {
2877 if let EngineError::Operation(OperationError::Agent(AgentError::SchemaValidation {
2878 raw_response: Some(text),
2879 ..
2880 })) = err
2881 {
2882 return Some(Value::String(text.clone()));
2883 }
2884 None
2885}
2886
2887fn extract_partial_usage_from_error(err: &EngineError) -> Option<StepPartialUsage> {
2888 if let EngineError::Operation(OperationError::Agent(AgentError::SchemaValidation {
2889 partial_usage,
2890 ..
2891 })) = err
2892 && (partial_usage.cost_usd.is_some() || partial_usage.duration_ms.is_some())
2893 {
2894 return Some(StepPartialUsage {
2895 cost_usd: partial_usage
2896 .cost_usd
2897 .and_then(|c| Decimal::try_from(c).ok()),
2898 duration_ms: partial_usage.duration_ms,
2899 input_tokens: partial_usage.input_tokens,
2900 output_tokens: partial_usage.output_tokens,
2901 });
2902 }
2903 None
2904}
2905
2906#[cfg(test)]
2907mod tests {
2908 use super::*;
2909 use ironflow_core::providers::claude::ClaudeCodeProvider;
2910 use ironflow_core::providers::record_replay::RecordReplayProvider;
2911 use ironflow_store::memory::InMemoryStore;
2912 use ironflow_store::models::{Run, RunActor, RunFilter};
2913 use ironflow_store::store::RunStore;
2914 use serde_json::json;
2915 use std::sync::Arc;
2916 use std::sync::atomic::{AtomicBool, Ordering};
2917 use uuid::Uuid;
2918
2919 fn create_test_provider() -> Arc<dyn ironflow_core::provider::AgentProvider> {
2921 let inner = ClaudeCodeProvider::new();
2922 Arc::new(RecordReplayProvider::replay(
2923 inner,
2924 "/tmp/ironflow-fixtures",
2925 ))
2926 }
2927
2928 fn create_test_context() -> WorkflowContext {
2930 let store = Arc::new(InMemoryStore::new());
2931 let provider = create_test_provider();
2932 let run_id = Uuid::now_v7();
2933 WorkflowContext::new(run_id, "test".to_string(), store, provider)
2934 }
2935
2936 #[test]
2937 fn context_new_initializes_correctly() {
2938 let ctx = create_test_context();
2939 assert_eq!(ctx.position, 0);
2940 assert_eq!(ctx.total_cost_usd, Decimal::ZERO);
2941 assert_eq!(ctx.total_duration_ms, 0);
2942 assert!(ctx.last_step_ids.is_empty());
2943 assert!(ctx.replay_steps.is_empty());
2944 assert!(ctx.log_sender.is_none());
2945 }
2946
2947 #[test]
2948 fn context_run_id_returns_correct_id() {
2949 let run_id = Uuid::now_v7();
2950 let store = Arc::new(InMemoryStore::new());
2951 let provider = create_test_provider();
2952 let ctx = WorkflowContext::new(run_id, "test".to_string(), store, provider);
2953 assert_eq!(ctx.run_id(), run_id);
2954 }
2955
2956 #[test]
2957 fn context_total_cost_usd_initially_zero() {
2958 let ctx = create_test_context();
2959 assert_eq!(ctx.total_cost_usd(), Decimal::ZERO);
2960 }
2961
2962 #[test]
2963 fn context_total_duration_ms_initially_zero() {
2964 let ctx = create_test_context();
2965 assert_eq!(ctx.total_duration_ms(), 0);
2966 }
2967
2968 #[test]
2969 fn context_with_handler_resolver_creates_context_with_resolver() {
2970 let store = Arc::new(InMemoryStore::new());
2971 let provider = create_test_provider();
2972 let run_id = Uuid::now_v7();
2973
2974 let called = Arc::new(AtomicBool::new(false));
2975 let called_clone = called.clone();
2976
2977 let resolver: HandlerResolver = Arc::new(move |_name: &str| {
2978 called_clone.store(true, Ordering::SeqCst);
2979 None
2980 });
2981
2982 let ctx = WorkflowContext::with_handler_resolver(
2983 run_id,
2984 "test".to_string(),
2985 store,
2986 provider,
2987 resolver,
2988 );
2989
2990 assert_eq!(ctx.run_id(), run_id);
2991 assert!(ctx.handler_resolver.is_some());
2992 }
2993
2994 #[tokio::test]
2995 async fn context_set_log_sender_attaches_sender() {
2996 let mut ctx = create_test_context();
2997 let (sender, _receiver) = crate::log_sender::channel();
2998 ctx.set_log_sender(sender);
2999 assert!(ctx.log_sender.is_some());
3000 }
3001
3002 #[tokio::test]
3003 async fn context_skip_creates_skipped_step() {
3004 let store = Arc::new(InMemoryStore::new());
3005 let provider = create_test_provider();
3006
3007 store
3009 .create_run(NewRun {
3010 created_by: None,
3011 workflow_name: "test".to_string(),
3012 trigger: TriggerKind::Manual,
3013 payload: json!({}),
3014 max_retries: 0,
3015 handler_version: None,
3016 labels: Default::default(),
3017 scheduled_at: None,
3018 idempotency_key: None,
3019 max_cost_usd: None,
3020 })
3021 .await
3022 .expect("failed to create run")
3023 .into_run();
3024
3025 let runs = store
3027 .list_runs(RunFilter::default(), 1, 10)
3028 .await
3029 .expect("failed to list runs");
3030 let created_run_id = runs.items[0].id;
3031
3032 let mut ctx =
3033 WorkflowContext::new(created_run_id, "test".to_string(), store.clone(), provider);
3034 let initial_position = ctx.position;
3035
3036 ctx.skip("skip-step", "condition not met")
3037 .await
3038 .expect("skip failed");
3039
3040 assert_eq!(ctx.position, initial_position + 1);
3041 assert!(!ctx.last_step_ids.is_empty());
3042
3043 let steps = store
3045 .list_steps(created_run_id)
3046 .await
3047 .expect("failed to list steps");
3048 assert_eq!(steps.len(), 1);
3049 assert_eq!(steps[0].status.state, StepStatus::Skipped);
3050 }
3051
3052 struct NoopSubWorkflow;
3055
3056 impl WorkflowHandler for NoopSubWorkflow {
3057 fn name(&self) -> &str {
3058 "noop-sub"
3059 }
3060
3061 fn execute<'a>(
3062 &'a self,
3063 _ctx: &'a mut WorkflowContext,
3064 ) -> crate::handler::HandlerFuture<'a> {
3065 Box::pin(async move { Ok(()) })
3066 }
3067 }
3068
3069 async fn child_run_of_parent_authored_by(created_by: Option<RunActor>) -> Run {
3072 let store = Arc::new(InMemoryStore::new());
3073 let provider = create_test_provider();
3074
3075 let parent = store
3076 .create_run(NewRun {
3077 workflow_name: "parent".to_string(),
3078 trigger: TriggerKind::Api,
3079 payload: json!({}),
3080 max_retries: 0,
3081 handler_version: None,
3082 labels: Default::default(),
3083 scheduled_at: None,
3084 created_by,
3085 idempotency_key: None,
3086 max_cost_usd: None,
3087 })
3088 .await
3089 .expect("failed to create parent run")
3090 .into_run();
3091
3092 let resolver: HandlerResolver = Arc::new(|name: &str| match name {
3093 "noop-sub" => Some(Arc::new(NoopSubWorkflow) as Arc<dyn WorkflowHandler>),
3094 _ => None,
3095 });
3096
3097 let mut ctx = WorkflowContext::with_handler_resolver(
3098 parent.id,
3099 "parent".to_string(),
3100 store.clone(),
3101 provider,
3102 resolver,
3103 );
3104 ctx.workflow(&NoopSubWorkflow, json!({}))
3105 .await
3106 .expect("sub-workflow failed");
3107
3108 let runs = store
3109 .list_runs(RunFilter::default(), 1, 10)
3110 .await
3111 .expect("failed to list runs");
3112 runs.items
3113 .into_iter()
3114 .find(|r| r.workflow_name == "noop-sub")
3115 .expect("child run was created")
3116 }
3117
3118 #[tokio::test]
3119 async fn child_run_inherits_the_parent_author() {
3120 let user_id = Uuid::now_v7();
3121 let child = child_run_of_parent_authored_by(Some(RunActor::User { user_id })).await;
3122
3123 assert_eq!(child.created_by, Some(RunActor::User { user_id }));
3124 }
3125
3126 #[tokio::test]
3127 async fn child_run_of_an_unattributed_parent_has_no_author() {
3128 let child = child_run_of_parent_authored_by(None).await;
3129
3130 assert!(child.created_by.is_none());
3131 }
3132
3133 #[tokio::test]
3134 async fn context_parallel_empty_steps_returns_empty_vec() {
3135 let mut ctx = create_test_context();
3136 let results = ctx
3137 .parallel(vec![], true)
3138 .await
3139 .expect("parallel should not fail on empty input");
3140 assert!(results.is_empty());
3141 }
3142
3143 #[tokio::test]
3144 async fn context_approval_first_execution_returns_error() {
3145 let store = Arc::new(InMemoryStore::new());
3146 let provider = create_test_provider();
3147
3148 store
3150 .create_run(NewRun {
3151 created_by: None,
3152 workflow_name: "test".to_string(),
3153 trigger: TriggerKind::Manual,
3154 payload: json!({}),
3155 max_retries: 0,
3156 handler_version: None,
3157 labels: Default::default(),
3158 scheduled_at: None,
3159 idempotency_key: None,
3160 max_cost_usd: None,
3161 })
3162 .await
3163 .expect("failed to create run")
3164 .into_run();
3165
3166 let runs = store
3168 .list_runs(RunFilter::default(), 1, 10)
3169 .await
3170 .expect("failed to list runs");
3171 let created_run_id = runs.items[0].id;
3172
3173 let mut ctx =
3174 WorkflowContext::new(created_run_id, "test".to_string(), store.clone(), provider);
3175
3176 let result = ctx
3177 .approval(
3178 "approve-step",
3179 crate::config::ApprovalConfig::new("Continue?"),
3180 )
3181 .await;
3182
3183 assert!(matches!(result, Err(EngineError::ApprovalRequired { .. })));
3185
3186 assert_eq!(ctx.position, 1);
3188
3189 let steps = store
3191 .list_steps(created_run_id)
3192 .await
3193 .expect("failed to list steps");
3194 assert_eq!(steps.len(), 1);
3195 assert_eq!(steps[0].status.state, StepStatus::AwaitingApproval);
3196 }
3197
3198 #[tokio::test]
3199 async fn context_approval_replay_returns_ok() {
3200 let store = Arc::new(InMemoryStore::new());
3201 let provider = create_test_provider();
3202
3203 store
3205 .create_run(NewRun {
3206 created_by: None,
3207 workflow_name: "test".to_string(),
3208 trigger: TriggerKind::Manual,
3209 payload: json!({}),
3210 max_retries: 0,
3211 handler_version: None,
3212 labels: Default::default(),
3213 scheduled_at: None,
3214 idempotency_key: None,
3215 max_cost_usd: None,
3216 })
3217 .await
3218 .expect("failed to create run")
3219 .into_run();
3220
3221 let runs = store
3223 .list_runs(RunFilter::default(), 1, 10)
3224 .await
3225 .expect("failed to list runs");
3226 let created_run_id = runs.items[0].id;
3227
3228 let step = store
3230 .create_step(NewStep {
3231 run_id: created_run_id,
3232 trace_id: step_trace_id(created_run_id, "approval", 0),
3233 name: "approval".to_string(),
3234 kind: StepKind::Approval,
3235 position: 0,
3236 input: None,
3237 is_error_handler: false,
3238 })
3239 .await
3240 .expect("failed to create step");
3241
3242 store
3244 .update_step(
3245 step.id,
3246 StepUpdate {
3247 status: Some(StepStatus::Running),
3248 started_at: Some(Utc::now()),
3249 ..StepUpdate::default()
3250 },
3251 )
3252 .await
3253 .expect("failed to update step to Running");
3254
3255 store
3256 .update_step(
3257 step.id,
3258 StepUpdate {
3259 status: Some(StepStatus::AwaitingApproval),
3260 ..StepUpdate::default()
3261 },
3262 )
3263 .await
3264 .expect("failed to update step to AwaitingApproval");
3265
3266 let mut ctx =
3268 WorkflowContext::new(created_run_id, "test".to_string(), store.clone(), provider);
3269 ctx.load_replay_steps()
3270 .await
3271 .expect("failed to load replay steps");
3272
3273 let result = ctx
3275 .approval("approval", crate::config::ApprovalConfig::new("Continue?"))
3276 .await;
3277
3278 assert!(result.is_ok());
3279
3280 let steps = store
3282 .list_steps(created_run_id)
3283 .await
3284 .expect("failed to list steps");
3285 assert_eq!(steps.len(), 1);
3286 assert_eq!(steps[0].status.state, StepStatus::Completed);
3287 }
3288
3289 #[tokio::test]
3290 async fn context_load_replay_steps_loads_completed_steps() {
3291 let store = Arc::new(InMemoryStore::new());
3292 let provider = create_test_provider();
3293
3294 store
3296 .create_run(NewRun {
3297 created_by: None,
3298 workflow_name: "test".to_string(),
3299 trigger: TriggerKind::Manual,
3300 payload: json!({}),
3301 max_retries: 0,
3302 handler_version: None,
3303 labels: Default::default(),
3304 scheduled_at: None,
3305 idempotency_key: None,
3306 max_cost_usd: None,
3307 })
3308 .await
3309 .expect("failed to create run")
3310 .into_run();
3311
3312 let runs = store
3314 .list_runs(RunFilter::default(), 1, 10)
3315 .await
3316 .expect("failed to list runs");
3317 let created_run_id = runs.items[0].id;
3318
3319 let completed_step = store
3321 .create_step(NewStep {
3322 run_id: created_run_id,
3323 trace_id: step_trace_id(created_run_id, "completed", 0),
3324 name: "completed".to_string(),
3325 kind: StepKind::Shell,
3326 position: 0,
3327 input: None,
3328 is_error_handler: false,
3329 })
3330 .await
3331 .expect("failed to create step");
3332
3333 store
3335 .update_step(
3336 completed_step.id,
3337 StepUpdate {
3338 status: Some(StepStatus::Running),
3339 started_at: Some(Utc::now()),
3340 ..StepUpdate::default()
3341 },
3342 )
3343 .await
3344 .expect("failed to update step to Running");
3345
3346 store
3347 .update_step(
3348 completed_step.id,
3349 StepUpdate {
3350 status: Some(StepStatus::Completed),
3351 completed_at: Some(Utc::now()),
3352 ..StepUpdate::default()
3353 },
3354 )
3355 .await
3356 .expect("failed to update step to Completed");
3357
3358 let _pending_step = store
3359 .create_step(NewStep {
3360 run_id: created_run_id,
3361 trace_id: step_trace_id(created_run_id, "pending", 1),
3362 name: "pending".to_string(),
3363 kind: StepKind::Shell,
3364 position: 1,
3365 input: None,
3366 is_error_handler: false,
3367 })
3368 .await
3369 .expect("failed to create step");
3370
3371 let mut ctx = WorkflowContext::new(created_run_id, "test".to_string(), store, provider);
3373 ctx.load_replay_steps()
3374 .await
3375 .expect("failed to load replay steps");
3376
3377 assert_eq!(ctx.replay_steps.len(), 1);
3379 assert!(ctx.replay_steps.contains_key(&0));
3380 assert!(!ctx.replay_steps.contains_key(&1));
3381 }
3382
3383 #[tokio::test]
3384 async fn context_payload_returns_run_payload() {
3385 let store = Arc::new(InMemoryStore::new());
3386 let provider = create_test_provider();
3387 let test_payload = json!({"key": "value", "number": 42});
3388
3389 store
3391 .create_run(NewRun {
3392 created_by: None,
3393 workflow_name: "test".to_string(),
3394 trigger: TriggerKind::Manual,
3395 payload: test_payload.clone(),
3396 max_retries: 0,
3397 handler_version: None,
3398 labels: Default::default(),
3399 scheduled_at: None,
3400 idempotency_key: None,
3401 max_cost_usd: None,
3402 })
3403 .await
3404 .expect("failed to create run")
3405 .into_run();
3406
3407 let runs = store
3409 .list_runs(RunFilter::default(), 1, 10)
3410 .await
3411 .expect("failed to list runs");
3412 let created_run_id = runs.items[0].id;
3413
3414 let ctx = WorkflowContext::new(created_run_id, "test".to_string(), store, provider);
3415 let payload = ctx.payload().await.expect("failed to get payload");
3416
3417 assert_eq!(payload, test_payload);
3418 }
3419
3420 #[tokio::test]
3421 async fn context_payload_returns_error_for_nonexistent_run() {
3422 let store = Arc::new(InMemoryStore::new());
3423 let provider = create_test_provider();
3424 let run_id = Uuid::now_v7();
3425
3426 let ctx = WorkflowContext::new(run_id, "test".to_string(), store, provider);
3427 let result = ctx.payload().await;
3428
3429 assert!(result.is_err());
3430 }
3431
3432 #[tokio::test]
3433 async fn context_store_returns_reference() {
3434 let ctx = create_test_context();
3435 let _store = ctx.store();
3436 }
3438
3439 #[test]
3440 fn context_debug_formatting() {
3441 let ctx = create_test_context();
3442 let debug_str = format!("{:?}", ctx);
3443 assert!(debug_str.contains("WorkflowContext"));
3444 assert!(debug_str.contains("run_id"));
3445 }
3446
3447 #[tokio::test]
3448 async fn context_last_step_ids_tracks_executed_steps() {
3449 let store = Arc::new(InMemoryStore::new());
3450 let provider = create_test_provider();
3451
3452 store
3454 .create_run(NewRun {
3455 created_by: None,
3456 workflow_name: "test".to_string(),
3457 trigger: TriggerKind::Manual,
3458 payload: json!({}),
3459 max_retries: 0,
3460 handler_version: None,
3461 labels: Default::default(),
3462 scheduled_at: None,
3463 idempotency_key: None,
3464 max_cost_usd: None,
3465 })
3466 .await
3467 .expect("failed to create run")
3468 .into_run();
3469
3470 let runs = store
3472 .list_runs(RunFilter::default(), 1, 10)
3473 .await
3474 .expect("failed to list runs");
3475 let created_run_id = runs.items[0].id;
3476
3477 let mut ctx = WorkflowContext::new(created_run_id, "test".to_string(), store, provider);
3478 assert!(ctx.last_step_ids.is_empty());
3479
3480 ctx.skip("step1", "reason").await.expect("skip failed");
3481
3482 assert_eq!(ctx.last_step_ids.len(), 1);
3483
3484 ctx.skip("step2", "reason").await.expect("skip failed");
3485
3486 assert_eq!(ctx.last_step_ids.len(), 1);
3488 }
3489}