1use std::collections::HashMap;
2use std::sync::Mutex;
3
4use async_trait::async_trait;
5
6use crate::engine::middleware::{Middleware, PostLlmCtx, UserMessageCtx};
7use crate::types::{AgentResult, SessionId};
8
9pub const DEFAULT_NUDGE_AFTER: usize = 5;
12pub const DEFAULT_BLOCK_AFTER: usize = 10;
15
16#[derive(Clone, Debug)]
18pub struct RepeatToolLimitConfig {
19 pub nudge_after: usize,
22 pub block_after: usize,
24 pub nudge_message: String,
25 pub block_message: String,
26}
27
28impl Default for RepeatToolLimitConfig {
29 fn default() -> Self {
30 Self {
31 nudge_after: DEFAULT_NUDGE_AFTER,
32 block_after: DEFAULT_BLOCK_AFTER,
33 nudge_message: "You have repeatedly issued the same tool call with the same arguments \
34 (e.g. polling list_agents for progress). Sub-agent reports are pushed to you \
35 automatically the moment your turn ends — polling cannot speed them up and only \
36 burns tokens. End your turn now with a brief progress note and no further tool calls."
37 .to_string(),
38 block_message: "You have issued this identical tool call far too many times. The \
39 pending calls were discarded. Based on everything you already know, write your \
40 response and end the turn. Do not call this tool again."
41 .to_string(),
42 }
43 }
44}
45
46pub struct RepeatToolLimitMiddleware {
73 config: RepeatToolLimitConfig,
74 counts: Mutex<HashMap<SessionId, HashMap<String, usize>>>,
75}
76
77impl RepeatToolLimitMiddleware {
78 pub fn new(config: RepeatToolLimitConfig) -> Self {
79 Self {
80 config,
81 counts: Mutex::new(HashMap::new()),
82 }
83 }
84
85 fn fingerprint(name: &str, args: &str) -> String {
87 let canonical = serde_json::from_str::<serde_json::Value>(args)
88 .map(|v| v.to_string())
89 .unwrap_or_else(|_| args.to_string());
90 format!("{name}::{canonical}")
91 }
92
93 fn bump(&self, session_id: &SessionId, fingerprint: &str) -> usize {
95 let mut counts = self.counts.lock().unwrap();
96 let session_counts = counts.entry(session_id.clone()).or_default();
97 let entry = session_counts.entry(fingerprint.to_string()).or_insert(0);
98 *entry += 1;
99 *entry
100 }
101}
102
103#[async_trait]
104impl Middleware for RepeatToolLimitMiddleware {
105 async fn on_user_message(&self, ctx: &mut UserMessageCtx) -> AgentResult<()> {
106 self.counts.lock().unwrap().remove(&ctx.session_id);
108 Ok(())
109 }
110
111 async fn on_post_llm(&self, ctx: &mut PostLlmCtx) -> AgentResult<()> {
112 if !ctx.is_tool_call || ctx.tool_calls.is_empty() {
113 return Ok(());
114 }
115
116 let mut max_count = 0usize;
117 let mut offending_tool = String::new();
118 for (_id, name, args) in &ctx.tool_calls {
119 let fingerprint = Self::fingerprint(name, args);
120 let count = self.bump(&ctx.session_id, &fingerprint);
121 if count > max_count {
122 max_count = count;
123 offending_tool = name.clone();
124 }
125 }
126
127 if max_count >= self.config.block_after {
128 tracing::warn!(
129 session_id = ctx.session_id.id,
130 tool = %offending_tool,
131 count = max_count,
132 pending_calls = ctx.tool_calls.len(),
133 "RepeatToolLimit: blocking tool calls — identical-call limit reached"
134 );
135 ctx.tool_calls.clear();
136 ctx.is_tool_call = false;
137 ctx.follow_up_message = Some(self.config.block_message.clone());
138 } else if self.config.nudge_after > 0 && max_count == self.config.nudge_after {
139 tracing::info!(
142 session_id = ctx.session_id.id,
143 tool = %offending_tool,
144 count = max_count,
145 "RepeatToolLimit: nudging — identical-call threshold reached"
146 );
147 ctx.follow_up_message = Some(self.config.nudge_message.clone());
148 }
149
150 Ok(())
151 }
152}
153
154#[cfg(test)]
155mod tests {
156 use super::*;
157 use crate::types::FinishReason;
158
159 fn mw() -> RepeatToolLimitMiddleware {
160 RepeatToolLimitMiddleware::new(RepeatToolLimitConfig::default())
161 }
162
163 fn ctx(session: u64, calls: Vec<(&str, &str)>) -> PostLlmCtx {
164 PostLlmCtx {
165 session_id: SessionId::new(session),
166 full_text: String::new(),
167 is_tool_call: !calls.is_empty(),
168 tool_calls: calls
169 .into_iter()
170 .enumerate()
171 .map(|(i, (name, args))| (format!("call_{i}"), name.to_string(), args.to_string()))
172 .collect(),
173 available_tools: vec![],
174 turn_count: 1,
175 total_tool_calls: 0,
176 nudge_count: 0,
177 turn_tool_calls: 0,
178 skip_push: false,
179 follow_up_message: None,
180 finish_reason: FinishReason::ToolUse,
181 }
182 }
183
184 async fn call(mw: &RepeatToolLimitMiddleware, session: u64, calls: Vec<(&str, &str)>) {
185 let mut c = ctx(session, calls);
186 mw.on_post_llm(&mut c).await.unwrap();
187 }
188
189 #[tokio::test]
190 async fn interleaved_identical_calls_still_count() {
191 let mw = mw();
195 for _ in 0..4 {
196 call(
197 &mw,
198 1,
199 vec![("list_agents", "{}"), ("read_file", r#"{"path":"a.rs"}"#)],
200 )
201 .await;
202 }
203 let mut c = ctx(
205 1,
206 vec![("list_agents", "{}"), ("read_file", r#"{"path":"b.rs"}"#)],
207 );
208 mw.on_post_llm(&mut c).await.unwrap();
209 assert!(c.follow_up_message.is_some(), "5th identical call nudges");
210 assert!(
211 !c.tool_calls.is_empty() && c.is_tool_call,
212 "nudge is soft — calls still execute"
213 );
214 }
215
216 #[tokio::test]
217 async fn nudge_fires_only_once_per_fingerprint() {
218 let mw = mw();
219 let mut fired = 0;
220 for _ in 0..8 {
221 let mut c = ctx(1, vec![("list_agents", "{}")]);
222 mw.on_post_llm(&mut c).await.unwrap();
223 if c.follow_up_message.is_some() {
224 fired += 1;
225 }
226 }
227 assert_eq!(fired, 1, "nudge must fire exactly once, got {fired}");
228 }
229
230 #[tokio::test]
231 async fn blocks_at_threshold_and_discards_pending_calls() {
232 let mw = mw();
233 for _ in 0..9 {
234 call(&mw, 1, vec![("list_agents", "{}")]).await;
235 }
236 let mut c = ctx(
237 1,
238 vec![("list_agents", "{}"), ("read_file", r#"{"path":"x"}"#)],
239 );
240 mw.on_post_llm(&mut c).await.unwrap();
241 assert!(c.tool_calls.is_empty(), "pending calls discarded");
242 assert!(!c.is_tool_call, "response demoted to text-only");
243 assert!(c.follow_up_message.is_some(), "block carries a follow-up");
244 }
245
246 #[tokio::test]
247 async fn different_args_are_different_fingerprints() {
248 let mw = mw();
253 for i in 0..20 {
254 call(
255 &mw,
256 1,
257 vec![("read_file", &format!(r#"{{"path":"f{i}.rs"}}"#))],
258 )
259 .await;
260 let mut probe = ctx(
261 1,
262 vec![("read_file", &format!(r#"{{"path":"probe{i}.rs"}}"#))],
263 );
264 mw.on_post_llm(&mut probe).await.unwrap();
265 assert!(
266 probe.follow_up_message.is_none() && probe.is_tool_call,
267 "legitimate distinct calls must never be limited"
268 );
269 }
270 }
271
272 #[tokio::test]
273 async fn same_args_reordered_keys_share_fingerprint() {
274 let mw = mw();
275 for _ in 0..4 {
276 call(&mw, 1, vec![("wait_agent", r#"{"name":"a","timeout":5}"#)]).await;
277 }
278 let mut c = ctx(1, vec![("wait_agent", r#"{"timeout":5,"name":"a"}"#)]);
280 mw.on_post_llm(&mut c).await.unwrap();
281 assert!(c.follow_up_message.is_some(), "canonical args must unify");
282 }
283
284 #[tokio::test]
285 async fn on_user_message_resets_counts() {
286 let mw = mw();
287 for _ in 0..4 {
288 call(&mw, 1, vec![("list_agents", "{}")]).await;
289 }
290 mw.on_user_message(&mut UserMessageCtx {
291 session_id: SessionId::new(1),
292 user_input: "next run".to_string(),
293 })
294 .await
295 .unwrap();
296 for _ in 0..4 {
298 let mut c = ctx(1, vec![("list_agents", "{}")]);
299 mw.on_post_llm(&mut c).await.unwrap();
300 assert!(c.follow_up_message.is_none(), "counts must reset per run");
301 }
302 }
303
304 #[tokio::test]
305 async fn sessions_are_isolated() {
306 let mw = mw();
307 for _ in 0..4 {
308 call(&mw, 1, vec![("list_agents", "{}")]).await;
309 }
310 let mut other = ctx(2, vec![("list_agents", "{}")]);
311 mw.on_post_llm(&mut other).await.unwrap();
312 assert!(
313 other.follow_up_message.is_none(),
314 "session 2's first call must not see session 1's count"
315 );
316 }
317}