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}