1use std::sync::Arc;
17use std::time::Duration;
18
19use async_trait::async_trait;
20use dashmap::DashMap;
21
22use crate::error::ToolError;
23
24use super::ToolFunction;
25
26#[derive(Debug, Clone, Default)]
28pub struct ToolPolicy {
29 pub timeout: Option<Duration>,
31 pub cache: bool,
33 pub confirm: bool,
35 pub confirm_message: Option<String>,
37}
38
39impl ToolPolicy {
40 pub fn new() -> Self {
42 Self::default()
43 }
44
45 pub fn is_noop(&self) -> bool {
49 self.timeout.is_none() && !self.cache && !self.confirm
50 }
51
52 pub fn with_timeout(mut self, d: Duration) -> Self {
54 self.timeout = Some(d);
55 self
56 }
57
58 pub fn with_cache(mut self) -> Self {
60 self.cache = true;
61 self
62 }
63
64 pub fn with_confirm(mut self, message: Option<String>) -> Self {
66 self.confirm = true;
67 self.confirm_message = message;
68 self
69 }
70
71 pub fn merge(mut self, other: &ToolPolicy) -> Self {
73 if other.timeout.is_some() {
74 self.timeout = other.timeout;
75 }
76 self.cache |= other.cache;
77 if other.confirm {
78 self.confirm = true;
79 if other.confirm_message.is_some() {
80 self.confirm_message = other.confirm_message.clone();
81 }
82 }
83 self
84 }
85}
86
87pub struct PolicyTool {
89 inner: Arc<dyn ToolFunction>,
90 policy: ToolPolicy,
91 cache: Arc<DashMap<String, serde_json::Value>>,
92}
93
94impl PolicyTool {
95 pub fn new(inner: Arc<dyn ToolFunction>, policy: ToolPolicy) -> Self {
97 Self {
98 inner,
99 policy,
100 cache: Arc::new(DashMap::new()),
101 }
102 }
103
104 pub fn wrap(inner: Arc<dyn ToolFunction>, policy: ToolPolicy) -> Arc<dyn ToolFunction> {
106 if policy.is_noop() {
107 inner
108 } else {
109 Arc::new(Self::new(inner, policy))
110 }
111 }
112
113 pub fn requires_confirmation(&self) -> bool {
115 self.policy.confirm
116 }
117
118 pub fn policy(&self) -> &ToolPolicy {
120 &self.policy
121 }
122
123 fn cache_key(&self, args: &serde_json::Value) -> String {
125 format!("{}\u{1}{}", self.inner.name(), canonical_json(args))
126 }
127}
128
129fn canonical_json(value: &serde_json::Value) -> String {
131 match value {
132 serde_json::Value::Object(map) => {
133 let mut keys: Vec<&String> = map.keys().collect();
134 keys.sort();
135 let mut out = String::from("{");
136 for (i, k) in keys.iter().enumerate() {
137 if i > 0 {
138 out.push(',');
139 }
140 out.push_str(&serde_json::to_string(k).unwrap_or_default());
141 out.push(':');
142 out.push_str(&canonical_json(&map[*k]));
143 }
144 out.push('}');
145 out
146 }
147 serde_json::Value::Array(items) => {
148 let mut out = String::from("[");
149 for (i, item) in items.iter().enumerate() {
150 if i > 0 {
151 out.push(',');
152 }
153 out.push_str(&canonical_json(item));
154 }
155 out.push(']');
156 out
157 }
158 other => serde_json::to_string(other).unwrap_or_default(),
159 }
160}
161
162#[async_trait]
163impl ToolFunction for PolicyTool {
164 fn name(&self) -> &str {
165 self.inner.name()
166 }
167
168 fn description(&self) -> &str {
169 self.inner.description()
170 }
171
172 fn parameters(&self) -> Option<serde_json::Value> {
173 self.inner.parameters()
174 }
175
176 fn requires_confirmation(&self) -> bool {
177 self.policy.confirm || self.inner.requires_confirmation()
181 }
182
183 fn confirmation_message(&self) -> Option<&str> {
184 self.policy
185 .confirm_message
186 .as_deref()
187 .or_else(|| self.inner.confirmation_message())
188 }
189
190 async fn call(&self, args: serde_json::Value) -> Result<serde_json::Value, ToolError> {
191 self.call_with_context(args, super::ToolContext::detached())
192 .await
193 }
194
195 async fn call_with_context(
196 &self,
197 args: serde_json::Value,
198 ctx: super::ToolContext,
199 ) -> Result<serde_json::Value, ToolError> {
200 let key = if self.policy.cache {
202 let key = self.cache_key(&args);
203 if let Some(hit) = self.cache.get(&key) {
204 return Ok(hit.clone());
205 }
206 Some(key)
207 } else {
208 None
209 };
210
211 let result = if let Some(timeout) = self.policy.timeout {
213 match tokio::time::timeout(timeout, self.inner.call_with_context(args, ctx)).await {
214 Ok(r) => r,
215 Err(_elapsed) => Err(ToolError::Timeout(timeout)),
216 }
217 } else {
218 self.inner.call_with_context(args, ctx).await
219 };
220
221 if let (Some(key), Ok(value)) = (key, &result) {
223 self.cache.insert(key, value.clone());
224 }
225
226 result
227 }
228}
229
230#[cfg(test)]
231mod tests {
232 use super::*;
233 use crate::tool::SimpleTool;
234 use serde_json::json;
235 use std::sync::atomic::{AtomicU32, Ordering};
236
237 #[tokio::test]
238 async fn timeout_policy_returns_timeout_error() {
239 let slow: Arc<dyn ToolFunction> = Arc::new(SimpleTool::new(
240 "slow",
241 "sleeps too long",
242 None,
243 |_| async move {
244 tokio::time::sleep(Duration::from_secs(3600)).await;
245 Ok(json!({"ok": true}))
246 },
247 ));
248 let tool = PolicyTool::new(
249 slow,
250 ToolPolicy::new().with_timeout(Duration::from_millis(50)),
251 );
252
253 match tool.call(json!({})).await {
254 Err(ToolError::Timeout(d)) => assert_eq!(d, Duration::from_millis(50)),
255 other => panic!("expected Timeout, got {other:?}"),
256 }
257 }
258
259 #[tokio::test]
260 async fn under_timeout_succeeds() {
261 let fast: Arc<dyn ToolFunction> = Arc::new(SimpleTool::new(
262 "fast",
263 "returns quickly",
264 None,
265 |_| async move { Ok(json!({"ok": true})) },
266 ));
267 let tool = PolicyTool::new(fast, ToolPolicy::new().with_timeout(Duration::from_secs(5)));
268 let out = tool.call(json!({})).await.unwrap();
269 assert_eq!(out["ok"], true);
270 }
271
272 #[tokio::test]
273 async fn cache_returns_same_value_and_runs_once() {
274 let counter = Arc::new(AtomicU32::new(0));
275 let c = counter.clone();
276 let counting: Arc<dyn ToolFunction> = Arc::new(SimpleTool::new(
277 "count",
278 "increments a counter",
279 None,
280 move |_| {
281 let c = c.clone();
282 async move {
283 let n = c.fetch_add(1, Ordering::SeqCst) + 1;
284 Ok(json!({"n": n}))
285 }
286 },
287 ));
288 let tool = PolicyTool::new(counting, ToolPolicy::new().with_cache());
289
290 let first = tool.call(json!({"x": 1})).await.unwrap();
291 let second = tool.call(json!({"x": 1})).await.unwrap();
292 assert_eq!(first, second);
293 assert_eq!(first["n"], 1);
294 assert_eq!(counter.load(Ordering::SeqCst), 1);
295
296 let third = tool.call(json!({"x": 2})).await.unwrap();
298 assert_eq!(third["n"], 2);
299 assert_eq!(counter.load(Ordering::SeqCst), 2);
300 }
301
302 #[tokio::test]
303 async fn cache_key_is_order_independent() {
304 let counter = Arc::new(AtomicU32::new(0));
305 let c = counter.clone();
306 let counting: Arc<dyn ToolFunction> = Arc::new(SimpleTool::new(
307 "count2",
308 "increments a counter",
309 None,
310 move |_| {
311 let c = c.clone();
312 async move {
313 c.fetch_add(1, Ordering::SeqCst);
314 Ok(json!({"ok": true}))
315 }
316 },
317 ));
318 let tool = PolicyTool::new(counting, ToolPolicy::new().with_cache());
319
320 tool.call(json!({"a": 1, "b": 2})).await.unwrap();
321 tool.call(json!({"b": 2, "a": 1})).await.unwrap();
323 assert_eq!(counter.load(Ordering::SeqCst), 1);
324 }
325
326 #[tokio::test]
327 async fn errors_are_not_cached() {
328 let counter = Arc::new(AtomicU32::new(0));
329 let c = counter.clone();
330 let failing: Arc<dyn ToolFunction> =
331 Arc::new(SimpleTool::new("fail", "always fails", None, move |_| {
332 let c = c.clone();
333 async move {
334 c.fetch_add(1, Ordering::SeqCst);
335 Err(ToolError::ExecutionFailed("boom".into()))
336 }
337 }));
338 let tool = PolicyTool::new(failing, ToolPolicy::new().with_cache());
339
340 assert!(tool.call(json!({})).await.is_err());
341 assert!(tool.call(json!({})).await.is_err());
342 assert_eq!(counter.load(Ordering::SeqCst), 2);
343 }
344
345 #[tokio::test]
346 async fn wrap_skips_noop_policy() {
347 let inner: Arc<dyn ToolFunction> = Arc::new(SimpleTool::new(
348 "plain",
349 "plain tool",
350 None,
351 |_| async move { Ok(json!({})) },
352 ));
353 let wrapped = PolicyTool::wrap(inner.clone(), ToolPolicy::new());
354 assert_eq!(wrapped.name(), "plain");
355 let confirmed = PolicyTool::wrap(inner, ToolPolicy::new().with_confirm(None));
357 assert_eq!(confirmed.name(), "plain");
358 }
359
360 #[test]
361 fn confirm_flag_is_recorded() {
362 let inner: Arc<dyn ToolFunction> = Arc::new(SimpleTool::new(
363 "danger",
364 "dangerous",
365 None,
366 |_| async move { Ok(json!({})) },
367 ));
368 let tool = PolicyTool::new(
369 inner,
370 ToolPolicy::new().with_confirm(Some("are you sure?".into())),
371 );
372 assert!(tool.requires_confirmation());
373 assert_eq!(
374 tool.policy().confirm_message.as_deref(),
375 Some("are you sure?")
376 );
377 }
378}