Skip to main content

agent_base/engine/
repeat_tool_limit.rs

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
9/// Default nudge threshold: identical call count at which the model gets a
10/// one-time "stop polling" follow-up.
11pub const DEFAULT_NUDGE_AFTER: usize = 5;
12/// Default block threshold: identical call count at which pending calls are
13/// discarded outright (same hard block as `TurnToolLimitMiddleware`).
14pub const DEFAULT_BLOCK_AFTER: usize = 10;
15
16/// Configuration for [`RepeatToolLimitMiddleware`].
17#[derive(Clone, Debug)]
18pub struct RepeatToolLimitConfig {
19    /// Identical-call count that triggers the one-time nudge. `0` disables
20    /// nudging (blocking still applies at `block_after`).
21    pub nudge_after: usize,
22    /// Identical-call count that hard-blocks the pending calls.
23    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
46/// Break poll loops — repeated identical tool calls — hard.
47///
48/// Session 20260914_50cf809d: a root spawned two sub-agents, then polled
49/// `list_agents` 124 times over 15 minutes (57% of wall clock, interleaved
50/// with real work) instead of ending its turn to receive the pushed reports.
51/// The system prompt forbade polling; the model did it anyway, and no
52/// mechanism stopped it. This middleware is that mechanism.
53///
54/// **Counting is cumulative per run, not consecutive**: the incident's polls
55/// were interleaved with `read_file` etc., so a consecutive-counter would
56/// reset forever and never fire. The fingerprint is tool name + canonical
57/// (key-sorted) arguments — a model legitimately reading 51 *different*
58/// files never accumulates a fingerprint.
59///
60/// Two-stage response, mirroring existing middleware precedents:
61/// - at `nudge_after` (once per fingerprint): `follow_up_message` nudge
62///   (soft, like `MaxTurnsNudgeMiddleware`);
63/// - at `block_after` (and every attempt beyond): pending `tool_calls` are
64///   discarded and a follow-up forces a summary (hard, like
65///   `TurnToolLimitMiddleware`).
66///
67/// State is per-session and reset by `on_user_message` — a run's counts must
68/// not leak into the next run (the middleware instance is process-long in
69/// phimint). Clearing all pending calls on block (not just the offending
70/// one) follows the `TurnToolLimitMiddleware` precedent: the breaker only
71/// trips after `block_after` identical calls, so collateral is negligible.
72pub 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    /// Fingerprint of one tool call: name + canonical (key-sorted) args.
86    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    /// Bump the run counter for a fingerprint; returns the new count.
94    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        // New run: last run's call counts are stale by definition.
107        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            // Fire exactly once per fingerprint: the count equals the
140            // threshold only on the crossing call.
141            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        // The incident shape: polls interleaved with other work. A
192        // consecutive counter would reset on every read_file; cumulative
193        // counting must still reach the nudge at 5 identical calls.
194        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        // 5th identical list_agents → nudge.
204        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        // A model legitimately reading many different files must never trip.
249        // Each probe uses a unique path so the probe itself does not
250        // accumulate to the same fingerprint (a repeated probe would count
251        // as an identical call on its own).
252        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        // Same call, keys in the other order → still the 5th identical call.
279        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        // Fresh run: 4 more calls are below the threshold again.
297        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}