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::{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 output_tokens = result.output_tokens();
78
79        info!(
80            step_kind = "agent",
81            model = %self.config.model,
82            cost_usd = %cost,
83            input_tokens = ?input_tokens,
84            output_tokens = ?output_tokens,
85            duration_ms,
86            "agent step completed"
87        );
88
89        let pricing = StaticPricing::new();
90        let breakdown = CostBreakdown::compute(
91            &pricing,
92            &self.config.model,
93            input_tokens.unwrap_or(0),
94            output_tokens.unwrap_or(0),
95        );
96        spawn_log("agent", &self.config.model, breakdown);
97
98        #[cfg(feature = "prometheus")]
99        {
100            use ironflow_core::metric_names::{
101                AGENT_COST_USD_TOTAL, AGENT_DURATION_SECONDS, AGENT_TOKENS_INPUT_TOTAL,
102                AGENT_TOKENS_OUTPUT_TOTAL, AGENT_TOTAL, STATUS_SUCCESS,
103            };
104            use metrics::{counter, gauge, histogram};
105            let model_label = self.config.model.clone();
106            counter!(AGENT_TOTAL, "model" => model_label.clone(), "status" => STATUS_SUCCESS)
107                .increment(1);
108            histogram!(AGENT_DURATION_SECONDS, "model" => model_label.clone())
109                .record(duration_ms as f64 / 1000.0);
110            gauge!(AGENT_COST_USD_TOTAL, "model" => model_label.clone())
111                .increment(cost.to_string().parse::<f64>().unwrap_or(0.0));
112            if let Some(inp) = input_tokens {
113                counter!(AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(inp);
114            }
115            if let Some(out) = output_tokens {
116                counter!(AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label).increment(out);
117            }
118        }
119
120        if let Some(ref sender) = self.log_sender {
121            sender.emit(
122                LogStream::System,
123                &format!(
124                    "agent step completed (cost=${cost}, tokens_in={}, tokens_out={})",
125                    input_tokens.unwrap_or(0),
126                    output_tokens.unwrap_or(0),
127                ),
128            );
129        }
130
131        let debug_messages = result.debug_messages().map(|msgs| msgs.to_vec());
132
133        Ok(StepOutput {
134            output: result.value().clone(),
135            duration_ms,
136            cost_usd: cost,
137            input_tokens,
138            output_tokens,
139            model: result.model().map(String::from),
140            debug_messages,
141        })
142    }
143}
144
145#[cfg(test)]
146mod tests {
147    use ironflow_core::operations::agent::PermissionMode;
148
149    #[test]
150    fn parse_permission_mode_via_serde() {
151        let json = r#""auto""#;
152        let mode: PermissionMode = serde_json::from_str(json).unwrap();
153        assert!(matches!(mode, PermissionMode::Auto));
154    }
155
156    #[test]
157    fn parse_permission_mode_dont_ask() {
158        let json = r#""dont_ask""#;
159        let mode: PermissionMode = serde_json::from_str(json).unwrap();
160        assert!(matches!(mode, PermissionMode::DontAsk));
161    }
162
163    #[test]
164    fn parse_permission_mode_bypass() {
165        let json = r#""bypass""#;
166        let mode: PermissionMode = serde_json::from_str(json).unwrap();
167        assert!(matches!(mode, PermissionMode::BypassPermissions));
168    }
169
170    #[test]
171    fn parse_permission_mode_case_insensitive() {
172        let json = r#""AUTO""#;
173        let mode: PermissionMode = serde_json::from_str(json).unwrap();
174        assert!(matches!(mode, PermissionMode::Auto));
175
176        let json = r#""DONT_ASK""#;
177        let mode: PermissionMode = serde_json::from_str(json).unwrap();
178        assert!(matches!(mode, PermissionMode::DontAsk));
179    }
180
181    #[test]
182    fn parse_permission_mode_unknown_defaults() {
183        let json = r#""unknown""#;
184        let mode: PermissionMode = serde_json::from_str(json).unwrap();
185        assert!(matches!(mode, PermissionMode::Default));
186    }
187}