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::operations::agent::Agent;
11use ironflow_core::pricing::{CostBreakdown, StaticPricing, spawn_log};
12use ironflow_core::provider::{AgentConfig, AgentProvider, LogSink};
13use ironflow_store::entities::StepKind;
14
15use crate::error::EngineError;
16use crate::log_sender::StepLogSender;
17use crate::notify::LogStream;
18
19use super::{StepArtifacts, StepExecutor, StepOutput};
20
21/// Executor for agent (AI) steps.
22///
23/// Runs an AI agent with the given prompt and configuration, capturing
24/// the response value, cost, and token counts. When a [`StepLogSender`]
25/// is attached, emits system log lines for step start/end.
26pub struct AgentExecutor<'a> {
27    config: &'a AgentConfig,
28    log_sender: Option<StepLogSender>,
29}
30
31impl<'a> AgentExecutor<'a> {
32    /// Create a new agent executor from a config reference.
33    pub fn new(config: &'a AgentConfig) -> Self {
34        Self {
35            config,
36            log_sender: None,
37        }
38    }
39
40    /// Attach a log sender for system-level log lines.
41    pub fn with_log_sender(mut self, sender: StepLogSender) -> Self {
42        self.log_sender = Some(sender);
43        self
44    }
45}
46
47impl StepExecutor for AgentExecutor<'_> {
48    fn kind(&self) -> StepKind {
49        StepKind::Agent
50    }
51
52    async fn execute(&self, provider: &Arc<dyn AgentProvider>) -> Result<StepOutput, EngineError> {
53        let start = Instant::now();
54
55        if let Some(ref sender) = self.log_sender {
56            sender.emit(
57                LogStream::System,
58                &format!("agent step started (model={})", self.config.model),
59            );
60        }
61
62        if self.config.json_schema.is_some() && self.config.max_turns == Some(1) {
63            warn!(
64                "structured output (json_schema) requires max_turns >= 2; \
65                 max_turns is set to 1, the agent will likely fail with error_max_turns"
66            );
67        }
68
69        let mut agent = Agent::from_config(self.config.clone());
70        if let Some(ref sender) = self.log_sender {
71            agent = agent.log_sink(Arc::new(sender.clone()) as Arc<dyn LogSink>);
72        }
73        let result = agent.run(provider.as_ref()).await?;
74
75        let duration_ms = start.elapsed().as_millis() as u64;
76        let cost = Decimal::try_from(result.cost_usd().unwrap_or(0.0)).unwrap_or(Decimal::ZERO);
77        let input_tokens = result.input_tokens();
78        let cache_read_tokens = result.cache_read_input_tokens();
79        let cache_creation_tokens = result.cache_creation_input_tokens();
80        let output_tokens = result.output_tokens();
81
82        info!(
83            step_kind = "agent",
84            model = %self.config.model,
85            cost_usd = %cost,
86            input_tokens = ?input_tokens,
87            cache_read_input_tokens = ?cache_read_tokens,
88            cache_creation_input_tokens = ?cache_creation_tokens,
89            output_tokens = ?output_tokens,
90            duration_ms,
91            "agent step completed"
92        );
93
94        let pricing = StaticPricing::new();
95        let breakdown = CostBreakdown::compute_with_cache(
96            &pricing,
97            &self.config.model,
98            input_tokens.unwrap_or(0),
99            cache_read_tokens.unwrap_or(0),
100            cache_creation_tokens.unwrap_or(0),
101            output_tokens.unwrap_or(0),
102        );
103        spawn_log("agent", &self.config.model, breakdown);
104
105        #[cfg(feature = "prometheus")]
106        {
107            use ironflow_core::metric_names::{
108                AGENT_COST_USD_TOTAL, AGENT_DURATION_SECONDS, AGENT_TOKENS_CACHE_READ_TOTAL,
109                AGENT_TOKENS_CACHE_WRITE_TOTAL, AGENT_TOKENS_INPUT_TOTAL,
110                AGENT_TOKENS_OUTPUT_TOTAL, AGENT_TOTAL, STATUS_SUCCESS,
111            };
112            use metrics::{counter, gauge, histogram};
113            let model_label = self.config.model.clone();
114            counter!(AGENT_TOTAL, "model" => model_label.clone(), "status" => STATUS_SUCCESS)
115                .increment(1);
116            histogram!(AGENT_DURATION_SECONDS, "model" => model_label.clone())
117                .record(duration_ms as f64 / 1000.0);
118            gauge!(AGENT_COST_USD_TOTAL, "model" => model_label.clone())
119                .increment(cost.to_string().parse::<f64>().unwrap_or(0.0));
120            if let Some(inp) = input_tokens {
121                counter!(AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(inp);
122            }
123            if let Some(t) = cache_read_tokens {
124                counter!(AGENT_TOKENS_CACHE_READ_TOTAL, "model" => model_label.clone())
125                    .increment(t);
126            }
127            if let Some(t) = cache_creation_tokens {
128                counter!(AGENT_TOKENS_CACHE_WRITE_TOTAL, "model" => model_label.clone())
129                    .increment(t);
130            }
131            if let Some(out) = output_tokens {
132                counter!(AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label).increment(out);
133            }
134        }
135
136        if let Some(ref sender) = self.log_sender {
137            sender.emit(
138                LogStream::System,
139                &format!(
140                    "agent step completed (cost=${cost}, tokens_in={}, cache_read={}, cache_write={}, tokens_out={})",
141                    input_tokens.unwrap_or(0),
142                    cache_read_tokens.unwrap_or(0),
143                    cache_creation_tokens.unwrap_or(0),
144                    output_tokens.unwrap_or(0),
145                ),
146            );
147        }
148
149        let debug_messages = result.debug_messages().map(|msgs| msgs.to_vec());
150        let account_id = match result.account_id() {
151            Some(raw) => match Uuid::parse_str(raw) {
152                Ok(id) => Some(id),
153                Err(e) => {
154                    warn!(account_id = raw, error = %e, "agent output carries an invalid account id");
155                    None
156                }
157            },
158            None => None,
159        };
160
161        Ok(StepOutput {
162            output: result.value().clone(),
163            duration_ms,
164            cost_usd: cost,
165            input_tokens,
166            cache_read_input_tokens: cache_read_tokens,
167            cache_creation_input_tokens: cache_creation_tokens,
168            output_tokens,
169            model: result.model().map(String::from),
170            debug_messages,
171            artifacts: StepArtifacts::default(),
172            account_id,
173            environment_id: result.environment_id().map(String::from),
174        })
175    }
176}
177
178#[cfg(test)]
179mod tests {
180    use std::sync::Arc;
181    use std::time::Duration;
182
183    use ironflow_core::operations::agent::PermissionMode;
184    use ironflow_core::provider::{AgentConfig, AgentOutput, AgentProvider, InvokeFuture};
185    use serde_json::json;
186    use tokio::time::timeout;
187    use uuid::Uuid;
188
189    use super::{AgentExecutor, StepExecutor};
190
191    /// Real provider returning a fixed output, used to exercise usage propagation.
192    struct FixedUsageProvider {
193        output: AgentOutput,
194    }
195
196    impl AgentProvider for FixedUsageProvider {
197        fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
198            Box::pin(async move { Ok(self.output.clone()) })
199        }
200    }
201
202    fn budget_config() -> AgentConfig {
203        let mut config = AgentConfig::new("hi");
204        config.max_budget_usd = Some(0.10);
205        config
206    }
207
208    #[tokio::test]
209    async fn agent_executor_propagates_cache_tokens() {
210        timeout(Duration::from_secs(10), async {
211            let mut output = AgentOutput::new(json!("ok"));
212            output.input_tokens = Some(100);
213            output.cache_read_input_tokens = Some(5000);
214            output.cache_creation_input_tokens = Some(200);
215            output.output_tokens = Some(50);
216            output.cost_usd = Some(0.02);
217            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
218
219            let config = budget_config();
220            let step = AgentExecutor::new(&config)
221                .execute(&provider)
222                .await
223                .expect("agent step succeeds");
224
225            assert_eq!(step.input_tokens, Some(100));
226            assert_eq!(step.cache_read_input_tokens, Some(5000));
227            assert_eq!(step.cache_creation_input_tokens, Some(200));
228            assert_eq!(step.output_tokens, Some(50));
229            assert_eq!(step.total_tokens(), 5350);
230        })
231        .await
232        .expect("test timed out");
233    }
234
235    #[tokio::test]
236    async fn agent_executor_propagates_account_id() {
237        timeout(Duration::from_secs(10), async {
238            let account_id = Uuid::now_v7();
239            let mut output = AgentOutput::new(json!("ok"));
240            output.cost_usd = Some(0.02);
241            output.account_id = Some(account_id.to_string());
242            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
243
244            let step = AgentExecutor::new(&budget_config())
245                .execute(&provider)
246                .await
247                .expect("agent step succeeds");
248            assert_eq!(step.account_id, Some(account_id));
249
250            let mut invalid = AgentOutput::new(json!("ok"));
251            invalid.cost_usd = Some(0.02);
252            invalid.account_id = Some("not-a-uuid".to_string());
253            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output: invalid });
254            let step = AgentExecutor::new(&budget_config())
255                .execute(&provider)
256                .await
257                .expect("agent step succeeds");
258            assert_eq!(step.account_id, None);
259        })
260        .await
261        .expect("test timed out");
262    }
263
264    #[tokio::test]
265    async fn agent_executor_propagates_environment_id() {
266        timeout(Duration::from_secs(10), async {
267            let mut output = AgentOutput::new(json!("ok"));
268            output.cost_usd = Some(0.02);
269            output.environment_id = Some("ironflow-env-0a1b2c".to_string());
270            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
271            let step = AgentExecutor::new(&budget_config())
272                .execute(&provider)
273                .await
274                .expect("agent step succeeds");
275            assert_eq!(step.environment_id.as_deref(), Some("ironflow-env-0a1b2c"));
276
277            let mut output = AgentOutput::new(json!("ok"));
278            output.cost_usd = Some(0.02);
279            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
280            let step = AgentExecutor::new(&budget_config())
281                .execute(&provider)
282                .await
283                .expect("agent step succeeds");
284            assert_eq!(step.environment_id, None);
285        })
286        .await
287        .expect("test timed out");
288    }
289
290    #[tokio::test]
291    async fn agent_executor_without_cache_tokens_yields_none() {
292        timeout(Duration::from_secs(10), async {
293            let mut output = AgentOutput::new(json!("ok"));
294            output.input_tokens = Some(100);
295            output.output_tokens = Some(50);
296            output.cost_usd = Some(0.02);
297            let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
298
299            let config = budget_config();
300            let step = AgentExecutor::new(&config)
301                .execute(&provider)
302                .await
303                .expect("agent step succeeds");
304
305            assert_eq!(step.input_tokens, Some(100));
306            assert!(step.cache_read_input_tokens.is_none());
307            assert!(step.cache_creation_input_tokens.is_none());
308            assert_eq!(step.total_tokens(), 150);
309        })
310        .await
311        .expect("test timed out");
312    }
313
314    #[test]
315    fn parse_permission_mode_via_serde() {
316        let json = r#""auto""#;
317        let mode: PermissionMode = serde_json::from_str(json).unwrap();
318        assert!(matches!(mode, PermissionMode::Auto));
319    }
320
321    #[test]
322    fn parse_permission_mode_dont_ask() {
323        let json = r#""dont_ask""#;
324        let mode: PermissionMode = serde_json::from_str(json).unwrap();
325        assert!(matches!(mode, PermissionMode::DontAsk));
326    }
327
328    #[test]
329    fn parse_permission_mode_bypass() {
330        let json = r#""bypass""#;
331        let mode: PermissionMode = serde_json::from_str(json).unwrap();
332        assert!(matches!(mode, PermissionMode::BypassPermissions));
333    }
334
335    #[test]
336    fn parse_permission_mode_case_insensitive() {
337        let json = r#""AUTO""#;
338        let mode: PermissionMode = serde_json::from_str(json).unwrap();
339        assert!(matches!(mode, PermissionMode::Auto));
340
341        let json = r#""DONT_ASK""#;
342        let mode: PermissionMode = serde_json::from_str(json).unwrap();
343        assert!(matches!(mode, PermissionMode::DontAsk));
344    }
345
346    #[test]
347    fn parse_permission_mode_unknown_defaults() {
348        let json = r#""unknown""#;
349        let mode: PermissionMode = serde_json::from_str(json).unwrap();
350        assert!(matches!(mode, PermissionMode::Default));
351    }
352}