Skip to main content

vv_agent/
approval.rs

1use std::collections::{HashMap, HashSet};
2use std::future::Future;
3use std::pin::Pin;
4use std::sync::{Arc, Condvar, Mutex};
5use std::time::{Duration, Instant};
6
7use serde_json::Value;
8use uuid::Uuid;
9
10use crate::tools::ApprovalDecision;
11use crate::types::{Metadata, ToolCall};
12
13pub type ApprovalFuture<T> = Pin<Box<dyn Future<Output = Result<T, ApprovalError>> + Send>>;
14
15pub trait ApprovalProvider: Send + Sync {
16    fn should_request(&self, request: &ApprovalRequest) -> bool;
17    fn decide(&self, request: &ApprovalRequest) -> ApprovalFuture<Option<ApprovalDecision>>;
18}
19
20#[derive(Debug, Clone, PartialEq)]
21pub struct ApprovalRequest {
22    pub request_id: String,
23    pub run_id: String,
24    pub trace_id: String,
25    pub agent_name: String,
26    pub cycle_index: u32,
27    pub tool_call_id: String,
28    pub tool_name: String,
29    pub arguments: Value,
30    pub preview: String,
31    pub metadata: Metadata,
32}
33
34impl ApprovalRequest {
35    pub fn for_tool_call(
36        run_id: impl Into<String>,
37        trace_id: impl Into<String>,
38        agent_name: impl Into<String>,
39        cycle_index: u32,
40        call: &ToolCall,
41    ) -> Self {
42        let run_id = run_id.into();
43        let trace_id = trace_id.into();
44        let agent_name = agent_name.into();
45        let arguments = Value::Object(call.arguments.clone().into_iter().collect());
46        Self {
47            request_id: new_approval_request_id(),
48            run_id,
49            trace_id,
50            agent_name,
51            cycle_index,
52            tool_call_id: call.id.clone(),
53            tool_name: call.name.clone(),
54            preview: format!("{} {}", call.name, arguments),
55            arguments,
56            metadata: Metadata::new(),
57        }
58    }
59}
60
61pub(crate) fn new_approval_request_id() -> String {
62    format!("approval_{}", Uuid::new_v4().simple())
63}
64
65#[derive(Debug, Clone, PartialEq, Eq)]
66pub struct ApprovalError {
67    message: String,
68}
69
70impl ApprovalError {
71    pub fn new(message: impl Into<String>) -> Self {
72        Self {
73            message: message.into(),
74        }
75    }
76
77    pub fn message(&self) -> &str {
78        &self.message
79    }
80}
81
82impl std::fmt::Display for ApprovalError {
83    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
84        formatter.write_str(&self.message)
85    }
86}
87
88impl std::error::Error for ApprovalError {}
89
90#[derive(Clone, Default)]
91pub struct ApprovalBroker {
92    inner: Arc<ApprovalBrokerInner>,
93}
94
95#[derive(Default)]
96struct ApprovalBrokerInner {
97    state: Mutex<ApprovalBrokerState>,
98    changed: Condvar,
99}
100
101#[derive(Default)]
102struct ApprovalBrokerState {
103    pending: HashMap<String, PendingApproval>,
104    session_allowed_tools: HashSet<String>,
105    cancel_decision: Option<ApprovalDecision>,
106}
107
108struct PendingApproval {
109    request: ApprovalRequest,
110    decision: Option<ApprovalDecision>,
111}
112
113impl ApprovalBroker {
114    pub fn register(&self, request: ApprovalRequest) -> Result<(), ApprovalError> {
115        let mut state = self
116            .inner
117            .state
118            .lock()
119            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
120        let decision = state.cancel_decision.clone().or_else(|| {
121            state
122                .session_allowed_tools
123                .contains(&request.tool_name)
124                .then_some(ApprovalDecision::ApprovedForSession)
125        });
126        state.pending.insert(
127            request.request_id.clone(),
128            PendingApproval { request, decision },
129        );
130        self.inner.changed.notify_all();
131        Ok(())
132    }
133
134    pub fn resolve(
135        &self,
136        request_id: impl AsRef<str>,
137        decision: ApprovalDecision,
138    ) -> Result<(), ApprovalError> {
139        let mut state = self
140            .inner
141            .state
142            .lock()
143            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
144        let request_id = request_id.as_ref();
145        let Some(tool_name) = state
146            .pending
147            .get(request_id)
148            .filter(|entry| entry.decision.is_none())
149            .map(|entry| entry.request.tool_name.clone())
150        else {
151            return Err(ApprovalError::new(format!(
152                "unknown approval request: {request_id}"
153            )));
154        };
155        let decision = state.cancel_decision.clone().unwrap_or(decision);
156        if decision.action() == "allow_session" {
157            state.session_allowed_tools.insert(tool_name.clone());
158        }
159        if let Some(entry) = state.pending.get_mut(request_id) {
160            entry.decision = Some(decision);
161        }
162        self.inner.changed.notify_all();
163        Ok(())
164    }
165
166    pub(crate) fn allows_tool_for_session(&self, tool_name: &str) -> Result<bool, ApprovalError> {
167        let state = self
168            .inner
169            .state
170            .lock()
171            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
172        Ok(state.cancel_decision.is_none() && state.session_allowed_tools.contains(tool_name))
173    }
174
175    #[cfg(test)]
176    pub(crate) fn allow_tool_for_session(&self, tool_name: &str) -> Result<(), ApprovalError> {
177        let mut state = self
178            .inner
179            .state
180            .lock()
181            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
182        if state.cancel_decision.is_some() {
183            return Ok(());
184        }
185        state.session_allowed_tools.insert(tool_name.to_string());
186        for entry in state
187            .pending
188            .values_mut()
189            .filter(|entry| entry.request.tool_name == tool_name)
190        {
191            entry.decision = Some(ApprovalDecision::ApprovedForSession);
192        }
193        self.inner.changed.notify_all();
194        Ok(())
195    }
196
197    pub fn pending_request(&self, request_id: impl AsRef<str>) -> Option<ApprovalRequest> {
198        self.inner.state.lock().ok().and_then(|state| {
199            state
200                .pending
201                .get(request_id.as_ref())
202                .filter(|entry| entry.decision.is_none())
203                .map(|entry| entry.request.clone())
204        })
205    }
206
207    pub(crate) fn discard(&self, request_id: &str) -> Result<bool, ApprovalError> {
208        let mut state = self
209            .inner
210            .state
211            .lock()
212            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
213        let removed = state.pending.remove(request_id).is_some();
214        if removed {
215            self.inner.changed.notify_all();
216        }
217        Ok(removed)
218    }
219
220    pub fn cancel_pending(&self, reason: impl Into<String>) -> Result<usize, ApprovalError> {
221        let decision = ApprovalDecision::deny(reason.into());
222        let mut state = self
223            .inner
224            .state
225            .lock()
226            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
227        state.cancel_decision = Some(decision.clone());
228        let pending_count = state
229            .pending
230            .values()
231            .filter(|entry| entry.decision.is_none())
232            .count();
233        for entry in state
234            .pending
235            .values_mut()
236            .filter(|entry| entry.decision.is_none())
237        {
238            entry.decision = Some(decision.clone());
239        }
240        self.inner.changed.notify_all();
241        Ok(pending_count)
242    }
243
244    pub(crate) fn reset_cancelled(&self) -> Result<(), ApprovalError> {
245        let mut state = self
246            .inner
247            .state
248            .lock()
249            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
250        state.cancel_decision = None;
251        Ok(())
252    }
253
254    pub(crate) fn wait_blocking(
255        &self,
256        request_id: &str,
257        timeout: Option<Duration>,
258    ) -> Result<ApprovalDecision, ApprovalError> {
259        let started = Instant::now();
260        let mut state = self
261            .inner
262            .state
263            .lock()
264            .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
265        loop {
266            if let Some(decision) = state
267                .pending
268                .get(request_id)
269                .and_then(|entry| entry.decision.clone())
270            {
271                state.pending.remove(request_id);
272                return Ok(decision);
273            }
274
275            if let Some(timeout) = timeout {
276                let elapsed = started.elapsed();
277                if elapsed >= timeout {
278                    state.pending.remove(request_id);
279                    return Ok(ApprovalDecision::timeout("Approval request timed out."));
280                }
281                let remaining = timeout.saturating_sub(elapsed);
282                let (next_state, wait_result) =
283                    self.inner
284                        .changed
285                        .wait_timeout(state, remaining)
286                        .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
287                state = next_state;
288                if wait_result.timed_out()
289                    && state
290                        .pending
291                        .get(request_id)
292                        .is_none_or(|entry| entry.decision.is_none())
293                {
294                    state.pending.remove(request_id);
295                    return Ok(ApprovalDecision::timeout("Approval request timed out."));
296                }
297            } else {
298                state = self
299                    .inner
300                    .changed
301                    .wait(state)
302                    .map_err(|_| ApprovalError::new("approval broker lock poisoned"))?;
303            }
304        }
305    }
306}
307
308pub(crate) fn block_on_approval_future<T: Send + 'static>(
309    future: ApprovalFuture<T>,
310) -> Result<T, ApprovalError> {
311    if let Ok(handle) = tokio::runtime::Handle::try_current() {
312        if handle.runtime_flavor() == tokio::runtime::RuntimeFlavor::MultiThread {
313            tokio::task::block_in_place(|| handle.block_on(future))
314        } else {
315            std::thread::spawn(move || {
316                tokio::runtime::Builder::new_current_thread()
317                    .enable_all()
318                    .build()
319                    .map_err(|error| ApprovalError::new(error.to_string()))?
320                    .block_on(future)
321            })
322            .join()
323            .map_err(|_| ApprovalError::new("approval future thread panicked"))?
324        }
325    } else {
326        tokio::runtime::Builder::new_current_thread()
327            .enable_all()
328            .build()
329            .map_err(|error| ApprovalError::new(error.to_string()))?
330            .block_on(future)
331    }
332}
333
334#[cfg(test)]
335mod tests {
336    use std::collections::BTreeMap;
337    use std::time::Duration;
338
339    use serde_json::json;
340
341    use super::{ApprovalBroker, ApprovalRequest};
342    use crate::tools::ApprovalDecision;
343    use crate::types::ToolCall;
344
345    fn request(id: &str, tool_name: &str) -> ApprovalRequest {
346        ApprovalRequest::for_tool_call(
347            "run",
348            "trace",
349            "agent",
350            0,
351            &ToolCall::new(
352                id,
353                tool_name,
354                BTreeMap::from([("path".to_string(), json!("file.txt"))]),
355            ),
356        )
357    }
358
359    #[test]
360    fn cancel_pending_wakes_waiters_and_applies_to_future_registrations() {
361        let broker = ApprovalBroker::default();
362        let first = request("first", "dangerous_tool");
363        let first_id = first.request_id.clone();
364        broker.register(first).expect("register first request");
365
366        let waiter = broker.clone();
367        let join = std::thread::spawn(move || waiter.wait_blocking(&first_id, None));
368        assert_eq!(
369            broker
370                .cancel_pending("Run was cancelled.")
371                .expect("cancel pending"),
372            1
373        );
374        assert!(matches!(
375            join.join().expect("join waiter").expect("decision"),
376            ApprovalDecision::Denied(reason) if reason == "Run was cancelled."
377        ));
378
379        let second = request("second", "dangerous_tool");
380        let second_id = second.request_id.clone();
381        broker.register(second).expect("register second request");
382        assert!(matches!(
383            broker
384                .wait_blocking(&second_id, Some(Duration::from_millis(10)))
385                .expect("future cancellation decision"),
386            ApprovalDecision::Denied(reason) if reason == "Run was cancelled."
387        ));
388
389        let late = request("late", "dangerous_tool");
390        let late_id = late.request_id.clone();
391        broker.register(late).expect("register late request");
392        assert!(broker
393            .resolve(&late_id, ApprovalDecision::allow_session())
394            .is_err());
395        assert!(matches!(
396            broker
397                .wait_blocking(&late_id, Some(Duration::from_millis(10)))
398                .expect("late cancellation decision"),
399            ApprovalDecision::Denied(reason) if reason == "Run was cancelled."
400        ));
401        assert!(!broker
402            .allows_tool_for_session("dangerous_tool")
403            .expect("cancelled session grant"));
404    }
405
406    #[test]
407    fn allow_session_grants_only_the_same_tool_for_the_broker_lifetime() {
408        let broker = ApprovalBroker::default();
409        let first = request("first", "dangerous_tool");
410        let first_id = first.request_id.clone();
411        broker.register(first).expect("register first request");
412        broker
413            .resolve(&first_id, ApprovalDecision::allow_session())
414            .expect("allow tool for session");
415        assert_eq!(
416            broker
417                .wait_blocking(&first_id, Some(Duration::from_millis(10)))
418                .expect("session decision"),
419            ApprovalDecision::ApprovedForSession
420        );
421
422        assert!(broker
423            .allows_tool_for_session("dangerous_tool")
424            .expect("session grant"));
425        assert!(!broker
426            .allows_tool_for_session("other_tool")
427            .expect("other tool grant"));
428
429        let repeated = request("repeated", "dangerous_tool");
430        let repeated_id = repeated.request_id.clone();
431        broker
432            .register(repeated)
433            .expect("register repeated request");
434        assert_eq!(
435            broker
436                .wait_blocking(&repeated_id, Some(Duration::from_millis(10)))
437                .expect("repeated decision"),
438            ApprovalDecision::ApprovedForSession
439        );
440    }
441
442    #[test]
443    fn allow_deny_and_timeout_do_not_grant_session_access() {
444        let broker = ApprovalBroker::default();
445        let decisions = [
446            ApprovalDecision::allow(),
447            ApprovalDecision::deny("not allowed"),
448            ApprovalDecision::timeout("too late"),
449        ];
450
451        for (index, decision) in decisions.into_iter().enumerate() {
452            let request = request(&format!("call_{index}"), "dangerous_tool");
453            let request_id = request.request_id.clone();
454            broker.register(request).expect("register request");
455            broker
456                .resolve(&request_id, decision.clone())
457                .expect("resolve request");
458            assert_eq!(
459                broker
460                    .wait_blocking(&request_id, Some(Duration::from_millis(10)))
461                    .expect("decision"),
462                decision
463            );
464            assert!(!broker
465                .allows_tool_for_session("dangerous_tool")
466                .expect("session grant"));
467        }
468    }
469
470    #[test]
471    fn first_resolution_wins_until_the_waiter_consumes_it() {
472        let broker = ApprovalBroker::default();
473        let request = request("first-wins", "dangerous_tool");
474        let request_id = request.request_id.clone();
475        broker.register(request).expect("register request");
476        broker
477            .resolve(&request_id, ApprovalDecision::allow())
478            .expect("resolve request");
479
480        assert!(broker
481            .resolve(&request_id, ApprovalDecision::deny("too late"))
482            .is_err());
483        assert!(broker.pending_request(&request_id).is_none());
484        assert_eq!(
485            broker
486                .wait_blocking(&request_id, Some(Duration::from_millis(10)))
487                .expect("first decision"),
488            ApprovalDecision::Approved
489        );
490    }
491
492    #[test]
493    fn allow_session_does_not_resolve_an_already_pending_same_tool_request() {
494        let broker = ApprovalBroker::default();
495        let first = request("session-first", "dangerous_tool");
496        let first_id = first.request_id.clone();
497        let second = request("session-second", "dangerous_tool");
498        let second_id = second.request_id.clone();
499        broker.register(first).expect("register first request");
500        broker.register(second).expect("register second request");
501
502        broker
503            .resolve(&first_id, ApprovalDecision::allow_session())
504            .expect("resolve first request");
505        assert_eq!(
506            broker
507                .wait_blocking(&first_id, Some(Duration::from_millis(10)))
508                .expect("first decision"),
509            ApprovalDecision::ApprovedForSession
510        );
511        assert!(broker.pending_request(&second_id).is_some());
512        broker
513            .resolve(&second_id, ApprovalDecision::allow())
514            .expect("resolve second request");
515        assert_eq!(
516            broker
517                .wait_blocking(&second_id, Some(Duration::from_millis(10)))
518                .expect("second decision"),
519            ApprovalDecision::Approved
520        );
521    }
522
523    #[test]
524    fn cancellation_preserves_an_existing_resolution_and_closes_future_requests() {
525        let broker = ApprovalBroker::default();
526        let resolved = request("resolved", "dangerous_tool");
527        let resolved_id = resolved.request_id.clone();
528        broker
529            .register(resolved)
530            .expect("register resolved request");
531        broker
532            .resolve(&resolved_id, ApprovalDecision::allow())
533            .expect("resolve request");
534
535        assert_eq!(
536            broker.cancel_pending("cancelled").expect("cancel broker"),
537            0
538        );
539        assert_eq!(
540            broker
541                .wait_blocking(&resolved_id, Some(Duration::from_millis(10)))
542                .expect("existing decision"),
543            ApprovalDecision::Approved
544        );
545
546        let future = request("future", "dangerous_tool");
547        let future_id = future.request_id.clone();
548        broker.register(future).expect("register future request");
549        assert!(broker.pending_request(&future_id).is_none());
550        assert!(broker
551            .resolve(&future_id, ApprovalDecision::allow())
552            .is_err());
553        assert_eq!(
554            broker
555                .wait_blocking(&future_id, Some(Duration::from_millis(10)))
556                .expect("cancellation decision"),
557            ApprovalDecision::Denied("cancelled".to_string())
558        );
559    }
560}