codei_tools/
approval_gate.rs1use std::sync::Arc;
2
3use async_trait::async_trait;
4use tokio::sync::{oneshot, Mutex};
5
6use crate::{ApprovalHandler, ApprovalRequest, ApprovalResponse};
7
8pub struct SharedApprovalGate {
10 inner: Arc<Mutex<GateState>>,
11}
12
13#[derive(Default)]
14struct GateState {
15 pending: Option<PendingApproval>,
16 always_approve: bool,
17}
18
19struct PendingApproval {
20 request: ApprovalRequest,
21 respond: oneshot::Sender<bool>,
22}
23
24impl Default for SharedApprovalGate {
25 fn default() -> Self {
26 Self::new()
27 }
28}
29
30impl SharedApprovalGate {
31 pub fn new() -> Self {
32 Self {
33 inner: Arc::new(Mutex::new(GateState::default())),
34 }
35 }
36
37 pub fn handler(self: &Arc<Self>) -> GateApprovalHandler {
38 GateApprovalHandler {
39 gate: Arc::clone(self),
40 }
41 }
42
43 pub async fn take_pending(&self) -> Option<ApprovalRequest> {
44 let guard = self.inner.lock().await;
45 guard.pending.as_ref().map(|p| p.request.clone())
46 }
47
48 pub async fn respond(&self, approved: bool) -> bool {
49 let respond = {
50 let mut guard = self.inner.lock().await;
51 guard.pending.take().map(|p| p.respond)
52 };
53 if let Some(tx) = respond {
54 tx.send(approved).is_ok()
55 } else {
56 false
57 }
58 }
59
60 pub async fn approve_always(&self) -> bool {
62 let respond = {
63 let mut guard = self.inner.lock().await;
64 guard.always_approve = true;
65 guard.pending.take().map(|p| p.respond)
66 };
67 if let Some(tx) = respond {
68 tx.send(true).is_ok()
69 } else {
70 false
71 }
72 }
73}
74
75pub struct GateApprovalHandler {
76 gate: Arc<SharedApprovalGate>,
77}
78
79#[async_trait]
80impl ApprovalHandler for GateApprovalHandler {
81 async fn approve(&self, request: ApprovalRequest) -> ApprovalResponse {
82 match request.tool_name.as_str() {
83 "write" | "edit" | "shell" => {}
84 _ => return ApprovalResponse { approved: true },
85 }
86
87 {
88 let guard = self.gate.inner.lock().await;
89 if guard.always_approve {
90 return ApprovalResponse { approved: true };
91 }
92 }
93
94 let (tx, rx) = oneshot::channel();
95 {
96 let mut guard = self.gate.inner.lock().await;
97 guard.pending = Some(PendingApproval {
98 request: request.clone(),
99 respond: tx,
100 });
101 }
102
103 let approved = rx.await.unwrap_or(false);
104 ApprovalResponse { approved }
105 }
106}
107
108#[cfg(test)]
109mod tests {
110 use serde_json::json;
111
112 use super::*;
113
114 #[tokio::test]
115 async fn approve_always_skips_future_prompts() {
116 let gate = Arc::new(SharedApprovalGate::new());
117 let handler = gate.handler();
118
119 let first = tokio::spawn({
120 let handler = gate.handler();
121 async move {
122 handler
123 .approve(ApprovalRequest {
124 tool_name: "shell".into(),
125 arguments: json!({"command": "ls"}),
126 })
127 .await
128 }
129 });
130
131 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
132 assert!(gate.take_pending().await.is_some());
133 assert!(gate.approve_always().await);
134
135 let first = first.await.unwrap();
136 assert!(first.approved);
137
138 let second = handler
139 .approve(ApprovalRequest {
140 tool_name: "write".into(),
141 arguments: json!({"path": "a.txt"}),
142 })
143 .await;
144 assert!(second.approved);
145 assert!(gate.take_pending().await.is_none());
146 }
147}