Skip to main content

codei_tools/
approval_gate.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use tokio::sync::{oneshot, Mutex};
5
6use crate::{ApprovalHandler, ApprovalRequest, ApprovalResponse};
7
8/// Pending approval surfaced to the UI layer.
9pub 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    /// Approve the current request and auto-approve future destructive tool calls.
61    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}