Skip to main content

ironflow_engine/executor/
agent.rs

1//! Agent step executor.
2
3use std::sync::Arc;
4use std::time::Instant;
5
6use rust_decimal::Decimal;
7use tracing::{info, warn};
8use uuid::Uuid;
9
10use ironflow_core::error::OperationError;
11use ironflow_core::operations::agent::{Agent, AgentResult};
12use ironflow_core::pricing::{CostBreakdown, StaticPricing, spawn_log};
13use ironflow_core::provider::{AgentConfig, AgentProvider, LogSink};
14use ironflow_core::providers::claude::is_session_not_found;
15use ironflow_store::entities::StepKind;
16
17use crate::error::EngineError;
18use crate::log_sender::StepLogSender;
19use crate::notify::LogStream;
20
21use super::{StepArtifacts, StepExecutor, StepOutput};
22
23/// Executor for agent (AI) steps.
24///
25/// Runs an AI agent with the given prompt and configuration, capturing
26/// the response value, cost, and token counts. When a [`StepLogSender`]
27/// is attached, emits system log lines for step start/end.
28pub struct AgentExecutor<'a> {
29    config: &'a AgentConfig,
30    log_sender: Option<StepLogSender>,
31}
32
33impl<'a> AgentExecutor<'a> {
34    /// Create a new agent executor from a config reference.
35    pub fn new(config: &'a AgentConfig) -> Self {
36        Self {
37            config,
38            log_sender: None,
39        }
40    }
41
42    /// Attach a log sender for system-level log lines.
43    pub fn with_log_sender(mut self, sender: StepLogSender) -> Self {
44        self.log_sender = Some(sender);
45        self
46    }
47
48    fn emit_system(&self, line: &str) {
49        if let Some(ref sender) = self.log_sender {
50            sender.emit(LogStream::System, line);
51        }
52    }
53
54    async fn run_agent(
55        &self,
56        config: AgentConfig,
57        provider: &Arc<dyn AgentProvider>,
58    ) -> Result<AgentResult, OperationError> {
59        let mut agent = Agent::from_config(config);
60        if let Some(ref sender) = self.log_sender {
61            agent = agent.log_sink(Arc::new(sender.clone()) as Arc<dyn LogSink>);
62        }
63        agent.run(provider.as_ref()).await
64    }
65
66    /// Resume `session_id` with `resume_prompt`. When the session does not
67    /// exist anymore (ephemeral HOME, other machine, other `cwd`), run the
68    /// original prompt from scratch in a session of the same id: the step
69    /// never fails because its session is gone.
70    async fn resume(
71        &self,
72        provider: &Arc<dyn AgentProvider>,
73        session_id: &str,
74        resume_prompt: &str,
75    ) -> Result<AgentResult, OperationError> {
76        info!(session_id, "agent step resumed from session");
77        self.emit_system(&format!("agent step resumed from session {session_id}"));
78
79        let mut resumed = self.config.clone();
80        resumed.prompt = resume_prompt.to_string();
81        match self.run_agent(resumed, provider).await {
82            Err(OperationError::Agent(ref err)) if is_session_not_found(err) => {
83                warn!(
84                    session_id,
85                    error = %err,
86                    "session not found, restarting the agent from scratch"
87                );
88                self.emit_system(&format!(
89                    "session {session_id} not found, restarting the agent from scratch"
90                ));
91                let mut fresh = self.config.clone();
92                fresh.resume_session_id = None;
93                fresh.session_id = Some(session_id.to_string());
94                self.run_agent(fresh, provider).await
95            }
96            other => other,
97        }
98    }
99}
100
101impl StepExecutor for AgentExecutor<'_> {
102    fn kind(&self) -> StepKind {
103        StepKind::Agent
104    }
105
106    async fn execute(&self, provider: &Arc<dyn AgentProvider>) -> Result<StepOutput, EngineError> {
107        let start = Instant::now();
108
109        if let Some(ref sender) = self.log_sender {
110            sender.emit(
111                LogStream::System,
112                &format!("agent step started (model={})", self.config.model),
113            );
114        }
115
116        if self.config.json_schema.is_some() && self.config.max_turns == Some(1) {
117            warn!(
118                "structured output (json_schema) requires max_turns >= 2; \
119                 max_turns is set to 1, the agent will likely fail with error_max_turns"
120            );
121        }
122
123        let result = match (&self.config.resume_session_id, &self.config.resume_prompt) {
124            (Some(session_id), Some(resume_prompt)) => {
125                self.resume(provider, session_id, resume_prompt).await?
126            }
127            _ => self.run_agent(self.config.clone(), provider).await?,
128        };
129
130        let duration_ms = start.elapsed().as_millis() as u64;
131        let cost = Decimal::try_from(result.cost_usd().unwrap_or(0.0)).unwrap_or(Decimal::ZERO);
132        let input_tokens = result.input_tokens();
133        let cache_read_tokens = result.cache_read_input_tokens();
134        let cache_creation_tokens = result.cache_creation_input_tokens();
135        let output_tokens = result.output_tokens();
136
137        info!(
138            step_kind = "agent",
139            model = %self.config.model,
140            cost_usd = %cost,
141            input_tokens = ?input_tokens,
142            cache_read_input_tokens = ?cache_read_tokens,
143            cache_creation_input_tokens = ?cache_creation_tokens,
144            output_tokens = ?output_tokens,
145            duration_ms,
146            "agent step completed"
147        );
148
149        let pricing = StaticPricing::new();
150        let breakdown = CostBreakdown::compute_with_cache(
151            &pricing,
152            &self.config.model,
153            input_tokens.unwrap_or(0),
154            cache_read_tokens.unwrap_or(0),
155            cache_creation_tokens.unwrap_or(0),
156            output_tokens.unwrap_or(0),
157        );
158        spawn_log("agent", &self.config.model, breakdown);
159
160        #[cfg(feature = "prometheus")]
161        {
162            use ironflow_core::metric_names::{
163                AGENT_COST_USD_TOTAL, AGENT_DURATION_SECONDS, AGENT_TOKENS_CACHE_READ_TOTAL,
164                AGENT_TOKENS_CACHE_WRITE_TOTAL, AGENT_TOKENS_INPUT_TOTAL,
165                AGENT_TOKENS_OUTPUT_TOTAL, AGENT_TOTAL, STATUS_SUCCESS,
166            };
167            use metrics::{counter, gauge, histogram};
168            let model_label = self.config.model.clone();
169            counter!(AGENT_TOTAL, "model" => model_label.clone(), "status" => STATUS_SUCCESS)
170                .increment(1);
171            histogram!(AGENT_DURATION_SECONDS, "model" => model_label.clone())
172                .record(duration_ms as f64 / 1000.0);
173            gauge!(AGENT_COST_USD_TOTAL, "model" => model_label.clone())
174                .increment(cost.to_string().parse::<f64>().unwrap_or(0.0));
175            if let Some(inp) = input_tokens {
176                counter!(AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(inp);
177            }
178            if let Some(t) = cache_read_tokens {
179                counter!(AGENT_TOKENS_CACHE_READ_TOTAL, "model" => model_label.clone())
180                    .increment(t);
181            }
182            if let Some(t) = cache_creation_tokens {
183                counter!(AGENT_TOKENS_CACHE_WRITE_TOTAL, "model" => model_label.clone())
184                    .increment(t);
185            }
186            if let Some(out) = output_tokens {
187                counter!(AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label).increment(out);
188            }
189        }
190
191        if let Some(ref sender) = self.log_sender {
192            sender.emit(
193                LogStream::System,
194                &format!(
195                    "agent step completed (cost=${cost}, tokens_in={}, cache_read={}, cache_write={}, tokens_out={})",
196                    input_tokens.unwrap_or(0),
197                    cache_read_tokens.unwrap_or(0),
198                    cache_creation_tokens.unwrap_or(0),
199                    output_tokens.unwrap_or(0),
200                ),
201            );
202        }
203
204        let debug_messages = result.debug_messages().map(|msgs| msgs.to_vec());
205        let account_id = match result.account_id() {
206            Some(raw) => match Uuid::parse_str(raw) {
207                Ok(id) => Some(id),
208                Err(e) => {
209                    warn!(account_id = raw, error = %e, "agent output carries an invalid account id");
210                    None
211                }
212            },
213            None => None,
214        };
215
216        Ok(StepOutput {
217            output: result.value().clone(),
218            duration_ms,
219            cost_usd: cost,
220            input_tokens,
221            cache_read_input_tokens: cache_read_tokens,
222            cache_creation_input_tokens: cache_creation_tokens,
223            output_tokens,
224            model: result.model().map(String::from),
225            debug_messages,
226            artifacts: StepArtifacts::default(),
227            account_id,
228            environment_id: result.environment_id().map(String::from),
229        })
230    }
231}
232
233#[cfg(test)]
234mod tests {
235    use std::sync::{Arc, Mutex};
236    use std::time::Duration;
237
238    use ironflow_core::error::AgentError;
239    use ironflow_core::operations::agent::PermissionMode;
240    use ironflow_core::provider::{AgentConfig, AgentOutput, AgentProvider, InvokeFuture};
241    use serde_json::json;
242    use tokio::time::timeout;
243    use uuid::Uuid;
244
245    use super::{AgentExecutor, StepExecutor};
246
247    /// Real provider returning a fixed output, used to exercise usage propagation.
248    struct FixedUsageProvider {
249        output: AgentOutput,
250    }
251
252    impl AgentProvider for FixedUsageProvider {
253        fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
254            Box::pin(async move { Ok(self.output.clone()) })
255        }
256    }
257
258    /// Provider recording every config it receives, failing a resume with
259    /// `resume_error` when set.
260    struct SessionProvider {
261        seen: Mutex<Vec<AgentConfig>>,
262        resume_error: Option<String>,
263    }
264
265    impl SessionProvider {
266        fn new(resume_error: Option<&str>) -> Arc<Self> {
267            Arc::new(Self {
268                seen: Mutex::new(Vec::new()),
269                resume_error: resume_error.map(String::from),
270            })
271        }
272
273        fn seen(&self) -> Vec<AgentConfig> {
274            self.seen.lock().unwrap().clone()
275        }
276    }
277
278    impl AgentProvider for SessionProvider {
279        fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
280            Box::pin(async move {
281                self.seen.lock().unwrap().push(config.clone());
282                match (&config.resume_session_id, &self.resume_error) {
283                    (Some(_), Some(stderr)) => Err(AgentError::ProcessFailed {
284                        exit_code: 1,
285                        stderr: stderr.clone(),
286                    }),
287                    _ => {
288                        let mut output = AgentOutput::new(json!("ok"));
289                        output.cost_usd = Some(0.02);
290                        Ok(output)
291                    }
292                }
293            })
294        }
295    }
296
297    const SID: &str = "0192f0c1-7d2e-7a4b-9c3d-1e2f3a4b5c6d";
298
299    #[tokio::test]
300    async fn agent_resume_executor_sends_resume_prompt() {
301        timeout(Duration::from_secs(10), async {
302            let recorder = SessionProvider::new(None);
303            let provider: Arc<dyn AgentProvider> = recorder.clone();
304            let config = budget_config().resume(SID).resume_prompt("go on");
305
306            AgentExecutor::new(&config)
307                .execute(&provider)
308                .await
309                .expect("resumed step succeeds");
310
311            let seen = recorder.seen();
312            assert_eq!(seen.len(), 1);
313            assert_eq!(seen[0].prompt, "go on");
314            assert_eq!(seen[0].resume_session_id.as_deref(), Some(SID));
315        })
316        .await
317        .expect("test timed out");
318    }
319
320    #[tokio::test]
321    async fn agent_resume_executor_falls_back_when_session_is_missing() {
322        timeout(Duration::from_secs(10), async {
323            let recorder =
324                SessionProvider::new(Some("No conversation found with session ID: 0192f0c1"));
325            let provider: Arc<dyn AgentProvider> = recorder.clone();
326            let config = budget_config().resume(SID).resume_prompt("go on");
327
328            AgentExecutor::new(&config)
329                .execute(&provider)
330                .await
331                .expect("a missing session never fails the step");
332
333            let seen = recorder.seen();
334            assert_eq!(seen.len(), 2);
335            assert_eq!(seen[1].prompt, "hi");
336            assert_eq!(seen[1].resume_session_id, None);
337            assert_eq!(seen[1].session_id.as_deref(), Some(SID));
338        })
339        .await
340        .expect("test timed out");
341    }
342
343    #[tokio::test]
344    async fn agent_resume_executor_propagates_other_errors() {
345        timeout(Duration::from_secs(10), async {
346            let recorder = SessionProvider::new(Some("permission denied"));
347            let provider: Arc<dyn AgentProvider> = recorder.clone();
348            let config = budget_config().resume(SID).resume_prompt("go on");
349
350            let err = AgentExecutor::new(&config)
351                .execute(&provider)
352                .await
353                .expect_err("other errors are not masked");
354
355            assert!(err.to_string().contains("permission denied"), "{err}");
356            assert_eq!(recorder.seen().len(), 1);
357        })
358        .await
359        .expect("test timed out");
360    }
361
362    #[tokio::test]
363    async fn agent_resume_executor_leaves_user_resume_alone() {
364        timeout(Duration::from_secs(10), async {
365            let recorder = SessionProvider::new(None);
366            let provider: Arc<dyn AgentProvider> = recorder.clone();
367            // A user-set resume without resume prompt keeps the original prompt.
368            let config = budget_config().resume(SID);
369
370            AgentExecutor::new(&config)
371                .execute(&provider)
372                .await
373                .expect("step succeeds");
374
375            let seen = recorder.seen();
376            assert_eq!(seen.len(), 1);
377            assert_eq!(seen[0].prompt, "hi");
378            assert_eq!(seen[0].resume_session_id.as_deref(), Some(SID));
379        })
380        .await
381        .expect("test timed out");
382    }
383
384    fn budget_config() -> AgentConfig {
385        let mut config = AgentConfig::new("hi");
386        config.max_budget_usd = Some(0.10);
387        config
388    }
389
390    #[tokio::test]
391    async fn agent_executor_propagates_cache_tokens() {
392        timeout(Duration::from_secs(10), async {
393            let mut output = AgentOutput::new(json!("ok"));
394            output.input_tokens = Some(100);
395            output.cache_read_input_tokens = Some(5000);
396            output.cache_creation_input_tokens = Some(200);
397            output.output_tokens = Some(50);
398            output.cost_usd = Some(0.02);
399            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
400
401            let config = budget_config();
402            let step = AgentExecutor::new(&config)
403                .execute(&provider)
404                .await
405                .expect("agent step succeeds");
406
407            assert_eq!(step.input_tokens, Some(100));
408            assert_eq!(step.cache_read_input_tokens, Some(5000));
409            assert_eq!(step.cache_creation_input_tokens, Some(200));
410            assert_eq!(step.output_tokens, Some(50));
411            assert_eq!(step.total_tokens(), 5350);
412        })
413        .await
414        .expect("test timed out");
415    }
416
417    #[tokio::test]
418    async fn agent_executor_propagates_account_id() {
419        timeout(Duration::from_secs(10), async {
420            let account_id = Uuid::now_v7();
421            let mut output = AgentOutput::new(json!("ok"));
422            output.cost_usd = Some(0.02);
423            output.account_id = Some(account_id.to_string());
424            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
425
426            let step = AgentExecutor::new(&budget_config())
427                .execute(&provider)
428                .await
429                .expect("agent step succeeds");
430            assert_eq!(step.account_id, Some(account_id));
431
432            let mut invalid = AgentOutput::new(json!("ok"));
433            invalid.cost_usd = Some(0.02);
434            invalid.account_id = Some("not-a-uuid".to_string());
435            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output: invalid });
436            let step = AgentExecutor::new(&budget_config())
437                .execute(&provider)
438                .await
439                .expect("agent step succeeds");
440            assert_eq!(step.account_id, None);
441        })
442        .await
443        .expect("test timed out");
444    }
445
446    #[tokio::test]
447    async fn agent_executor_propagates_environment_id() {
448        timeout(Duration::from_secs(10), async {
449            let mut output = AgentOutput::new(json!("ok"));
450            output.cost_usd = Some(0.02);
451            output.environment_id = Some("ironflow-env-0a1b2c".to_string());
452            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
453            let step = AgentExecutor::new(&budget_config())
454                .execute(&provider)
455                .await
456                .expect("agent step succeeds");
457            assert_eq!(step.environment_id.as_deref(), Some("ironflow-env-0a1b2c"));
458
459            let mut output = AgentOutput::new(json!("ok"));
460            output.cost_usd = Some(0.02);
461            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
462            let step = AgentExecutor::new(&budget_config())
463                .execute(&provider)
464                .await
465                .expect("agent step succeeds");
466            assert_eq!(step.environment_id, None);
467        })
468        .await
469        .expect("test timed out");
470    }
471
472    #[tokio::test]
473    async fn agent_executor_without_cache_tokens_yields_none() {
474        timeout(Duration::from_secs(10), async {
475            let mut output = AgentOutput::new(json!("ok"));
476            output.input_tokens = Some(100);
477            output.output_tokens = Some(50);
478            output.cost_usd = Some(0.02);
479            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
480
481            let config = budget_config();
482            let step = AgentExecutor::new(&config)
483                .execute(&provider)
484                .await
485                .expect("agent step succeeds");
486
487            assert_eq!(step.input_tokens, Some(100));
488            assert!(step.cache_read_input_tokens.is_none());
489            assert!(step.cache_creation_input_tokens.is_none());
490            assert_eq!(step.total_tokens(), 150);
491        })
492        .await
493        .expect("test timed out");
494    }
495
496    #[test]
497    fn parse_permission_mode_via_serde() {
498        let json = r#""auto""#;
499        let mode: PermissionMode = serde_json::from_str(json).unwrap();
500        assert!(matches!(mode, PermissionMode::Auto));
501    }
502
503    #[test]
504    fn parse_permission_mode_dont_ask() {
505        let json = r#""dont_ask""#;
506        let mode: PermissionMode = serde_json::from_str(json).unwrap();
507        assert!(matches!(mode, PermissionMode::DontAsk));
508    }
509
510    #[test]
511    fn parse_permission_mode_bypass() {
512        let json = r#""bypass""#;
513        let mode: PermissionMode = serde_json::from_str(json).unwrap();
514        assert!(matches!(mode, PermissionMode::BypassPermissions));
515    }
516
517    #[test]
518    fn parse_permission_mode_case_insensitive() {
519        let json = r#""AUTO""#;
520        let mode: PermissionMode = serde_json::from_str(json).unwrap();
521        assert!(matches!(mode, PermissionMode::Auto));
522
523        let json = r#""DONT_ASK""#;
524        let mode: PermissionMode = serde_json::from_str(json).unwrap();
525        assert!(matches!(mode, PermissionMode::DontAsk));
526    }
527
528    #[test]
529    fn parse_permission_mode_unknown_defaults() {
530        let json = r#""unknown""#;
531        let mode: PermissionMode = serde_json::from_str(json).unwrap();
532        assert!(matches!(mode, PermissionMode::Default));
533    }
534}