ironflow_engine/executor/
agent.rs1use 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
20pub struct AgentExecutor<'a> {
26 config: &'a AgentConfig,
27 log_sender: Option<StepLogSender>,
28}
29
30impl<'a> AgentExecutor<'a> {
31 pub fn new(config: &'a AgentConfig) -> Self {
33 Self {
34 config,
35 log_sender: None,
36 }
37 }
38
39 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 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}