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