Skip to main content

roder_core/runtime/
tool_approvals.rs

1use super::*;
2
3impl Runtime {
4    pub(crate) async fn auto_resolve_pending_tool_approvals(&self) {
5        let pending: Vec<_> = self
6            .pending_tool_approvals
7            .lock()
8            .await
9            .iter()
10            .map(|(id, approval)| (id.clone(), approval.call.clone(), approval.context.clone()))
11            .collect();
12        let gate = DefaultPolicyGate::new();
13        for (id, call, context) in pending {
14            // A runtime mode change must not widen a child's independently
15            // restricted permissions or skip an extension's explicit denial.
16            let mode = self.effective_policy_mode_for_thread(&call.thread_id).await;
17            let config = self.status().await;
18            let mut ctx = context.unwrap_or_else(|| {
19                self.tool_execution_context(
20                    call.thread_id.clone(),
21                    call.turn_id.clone(),
22                    mode,
23                    config.workspace.as_deref(),
24                    Some(&config.command_shell),
25                )
26            });
27            ctx.effective_mode = mode;
28            if mode != PolicyMode::Bypass
29                && !matches!(
30                    gate.decide(&call, mode, &ctx),
31                    PolicyDecision::AutoApproved { .. }
32                )
33            {
34                continue;
35            }
36            let Ok(decision) = gate
37                .decide_with_contributors(&call, mode, &ctx, &self.registry.policy_contributors)
38                .await
39            else {
40                continue;
41            };
42            let approved = matches!(
43                decision,
44                PolicyDecision::AutoApproved { .. } | PolicyDecision::Allowed
45            );
46            if !(approved || matches!(decision, PolicyDecision::Denied { .. }))
47                || self.effective_policy_mode_for_thread(&call.thread_id).await != mode
48            {
49                continue;
50            }
51            let Some(approval) = self.pending_tool_approvals.lock().await.remove(&id) else {
52                continue;
53            };
54            self.emit(RoderEvent::PolicyDecisionRecorded(PolicyDecisionRecorded {
55                thread_id: approval.thread_id.clone(),
56                turn_id: approval.turn_id.clone(),
57                tool_id: approval.tool_id.clone(),
58                tool_name: approval.tool_name.clone(),
59                mode,
60                decision,
61                timestamp: OffsetDateTime::now_utc(),
62            }))
63            .await;
64            if approved && mode == PolicyMode::Bypass {
65                self.emit(RoderEvent::PolicyBypassActive(PolicyBypassActive {
66                    thread_id: approval.thread_id.clone(),
67                    turn_id: approval.turn_id.clone(),
68                    tool_id: approval.tool_id.clone(),
69                    tool_name: approval.tool_name.clone(),
70                    timestamp: OffsetDateTime::now_utc(),
71                }))
72                .await;
73            }
74            self.emit(RoderEvent::ApprovalResolved(ApprovalResolved {
75                thread_id: approval.thread_id,
76                turn_id: approval.turn_id,
77                approval_id: id,
78                tool_id: approval.tool_id,
79                tool_name: approval.tool_name,
80                approved,
81                timestamp: OffsetDateTime::now_utc(),
82            }))
83            .await;
84            let _ = approval.tx.send(approved);
85        }
86    }
87
88    pub(crate) async fn request_tool_approval(
89        &self,
90        thread_id: &ThreadId,
91        turn_id: &TurnId,
92        call: &ToolCall,
93        reason: Option<String>,
94        context: &ToolExecutionContext,
95    ) -> anyhow::Result<bool> {
96        let approval_id = call.id.clone();
97        let (tx, rx) = tokio::sync::oneshot::channel();
98        self.pending_tool_approvals.lock().await.insert(
99            approval_id.clone(),
100            crate::runtime::PendingToolApproval {
101                thread_id: thread_id.clone(),
102                turn_id: turn_id.clone(),
103                tool_id: call.id.clone(),
104                tool_name: call.name.clone(),
105                call: call.clone(),
106                context: Some(context.clone()),
107                tx,
108            },
109        );
110        // Close the race where permissions changed after policy review but
111        // before the approval was registered. Full Access must not strand it.
112        self.auto_resolve_pending_tool_approvals().await;
113        if !self
114            .pending_tool_approvals
115            .lock()
116            .await
117            .contains_key(&approval_id)
118        {
119            return Ok(rx.await.unwrap_or(false));
120        }
121        let runtime_config = self.status().await;
122        crate::hooks::run_lifecycle(
123            self,
124            thread_id,
125            turn_id,
126            runtime_config.workspace.as_deref(),
127            "PermissionRequest",
128            Some(&call.name),
129            serde_json::json!({"toolName": call.name, "toolInput": call.arguments, "reason": reason}),
130        )
131        .await;
132        self.emit(RoderEvent::ApprovalRequested(ApprovalRequested {
133            thread_id: thread_id.clone(),
134            turn_id: turn_id.clone(),
135            approval_id,
136            tool_id: call.id.clone(),
137            tool_name: call.name.clone(),
138            reason,
139            timestamp: OffsetDateTime::now_utc(),
140        }))
141        .await;
142        Ok(rx.await.unwrap_or(false))
143    }
144
145    pub async fn request_app_server_tool_approval(
146        &self,
147        call: ToolCall,
148        reason: Option<String>,
149    ) -> anyhow::Result<bool> {
150        let approval_id = call.id.clone();
151        let (tx, rx) = oneshot::channel();
152        self.pending_tool_approvals.lock().await.insert(
153            approval_id.clone(),
154            PendingToolApproval {
155                thread_id: call.thread_id.clone(),
156                turn_id: call.turn_id.clone(),
157                tool_id: call.id.clone(),
158                tool_name: call.name.clone(),
159                call: call.clone(),
160                context: None,
161                tx,
162            },
163        );
164        self.auto_resolve_pending_tool_approvals().await;
165        if !self
166            .pending_tool_approvals
167            .lock()
168            .await
169            .contains_key(&approval_id)
170        {
171            return Ok(rx.await.unwrap_or(false));
172        }
173        self.emit(RoderEvent::ApprovalRequested(ApprovalRequested {
174            thread_id: call.thread_id.clone(),
175            turn_id: call.turn_id.clone(),
176            approval_id,
177            tool_id: call.id.clone(),
178            tool_name: call.name.clone(),
179            reason,
180            timestamp: OffsetDateTime::now_utc(),
181        }))
182        .await;
183        Ok(rx.await.unwrap_or(false))
184    }
185}
186
187#[cfg(test)]
188mod tests {
189    use super::*;
190    use crate::teams::{TeamMemberStartRequest, TeamStartRequest};
191
192    fn call(thread: &str) -> ToolCall {
193        ToolCall {
194            id: "pending-permission".into(),
195            name: "shell".into(),
196            arguments: serde_json::json!({"command":"cargo test"}),
197            raw_arguments: String::new(),
198            thread_id: thread.into(),
199            turn_id: "turn".into(),
200        }
201    }
202
203    async fn pending_request(
204        runtime: &Arc<Runtime>,
205        thread: &str,
206    ) -> tokio::task::JoinHandle<anyhow::Result<bool>> {
207        let mut events = runtime.subscribe_events();
208        let request = call(thread);
209        let runtime = runtime.clone();
210        let pending = tokio::spawn(async move {
211            runtime
212                .request_app_server_tool_approval(request, None)
213                .await
214        });
215        tokio::time::timeout(std::time::Duration::from_secs(5), async {
216            while !matches!(
217                events.recv().await.unwrap().event,
218                RoderEvent::ApprovalRequested(_)
219            ) {}
220        })
221        .await
222        .unwrap();
223        pending
224    }
225
226    #[tokio::test]
227    async fn full_access_switch_resolves_pending_internal_tool() {
228        let runtime = Arc::new(Runtime::fake().unwrap());
229        let mut events = runtime.subscribe_events();
230        let mut request = call("lead");
231        request.name = "spawn_agent".into();
232        let pending = {
233            let runtime = runtime.clone();
234            tokio::spawn(async move {
235                runtime
236                    .request_app_server_tool_approval(request, None)
237                    .await
238            })
239        };
240        tokio::time::timeout(std::time::Duration::from_secs(5), async {
241            while !matches!(
242                events.recv().await.unwrap().event,
243                RoderEvent::ApprovalRequested(_)
244            ) {}
245        })
246        .await
247        .unwrap();
248        runtime
249            .set_policy_mode(PolicyMode::Bypass, None)
250            .await
251            .unwrap();
252        assert!(pending.await.unwrap().unwrap());
253    }
254
255    #[tokio::test]
256    async fn full_access_auto_resolution_does_not_elevate_restricted_child() {
257        let data = tempfile::tempdir().unwrap();
258        let mut builder = roder_api::extension::ExtensionRegistryBuilder::new();
259        builder.inference_engine(Arc::new(crate::fake_provider::FakeInferenceEngine));
260        let runtime = Arc::new(
261            Runtime::new(
262                builder.build().unwrap(),
263                RuntimeConfig {
264                    team_data_dir: Some(data.path().into()),
265                    ..Default::default()
266                },
267            )
268            .unwrap(),
269        );
270        let team = runtime
271            .start_team(TeamStartRequest {
272                lead_thread_id: None,
273                display_mode: roder_api::teams::AgentTeamDisplayMode::InProcess,
274                members: vec![TeamMemberStartRequest {
275                    name: "Restricted".into(),
276                    model_provider: None,
277                    model: None,
278                }],
279            })
280            .await
281            .unwrap();
282        let child = &team.members[1].thread_id;
283        let pending = pending_request(&runtime, child).await;
284        runtime
285            .set_policy_mode(PolicyMode::Bypass, None)
286            .await
287            .unwrap();
288        assert!(
289            runtime
290                .pending_tool_approvals
291                .lock()
292                .await
293                .contains_key("pending-permission")
294        );
295        assert_eq!(
296            runtime.effective_policy_mode_for_thread(child).await,
297            PolicyMode::Default
298        );
299        runtime
300            .resolve_tool_approval("pending-permission", false)
301            .await
302            .unwrap();
303        assert!(!pending.await.unwrap().unwrap());
304    }
305
306    struct Deny;
307    #[async_trait::async_trait]
308    impl roder_api::context::PolicyContributor for Deny {
309        fn id(&self) -> String {
310            "mandatory-deny".into()
311        }
312        async fn review_tool(
313            &self,
314            _: roder_api::context::PolicyReview,
315        ) -> anyhow::Result<roder_api::context::PolicyContribution> {
316            Ok(roder_api::context::PolicyContribution::Deny {
317                reason: "host restriction".into(),
318            })
319        }
320    }
321
322    #[tokio::test]
323    async fn full_access_switch_rejects_pending_call_when_extension_denies() {
324        let mut builder = roder_api::extension::ExtensionRegistryBuilder::new();
325        builder.inference_engine(Arc::new(crate::fake_provider::FakeInferenceEngine));
326        builder.policy_contributor(Arc::new(Deny));
327        let runtime = Arc::new(Runtime::new(builder.build().unwrap(), Default::default()).unwrap());
328        let pending = pending_request(&runtime, "lead").await;
329        runtime
330            .set_policy_mode(PolicyMode::Bypass, None)
331            .await
332            .unwrap();
333        assert!(!pending.await.unwrap().unwrap());
334    }
335}