1use std::sync::Arc;
2
3use crate::acp;
4use crate::acp::Error as SdkError;
5use crate::zed::connection::ConnectionHandle;
6use async_trait::async_trait;
7use serde_json::Value;
8use tracing::{error, warn};
9
10use crate::reports::{
11 TOOL_PERMISSION_ALLOW_ALWAYS_OPTION_ID, TOOL_PERMISSION_ALLOW_OPTION_ID, TOOL_PERMISSION_ALLOW_PREFIX,
12 TOOL_PERMISSION_CANCELLED_MESSAGE, TOOL_PERMISSION_DENIED_MESSAGE, TOOL_PERMISSION_DENY_ALWAYS_OPTION_ID,
13 TOOL_PERMISSION_DENY_OPTION_ID, TOOL_PERMISSION_DENY_PREFIX, TOOL_PERMISSION_REQUEST_FAILURE_LOG,
14 TOOL_PERMISSION_REQUEST_FAILURE_MESSAGE, TOOL_PERMISSION_UNKNOWN_OPTION_LOG, ToolExecutionReport,
15};
16
17use super::tooling::{SupportedTool, ToolDescriptor, ToolRegistryProvider};
18
19#[derive(Clone, Copy, Debug)]
20pub struct PermissionToolContext<'a> {
21 name: &'a str,
22 kind: acp::ToolKind,
23 action_label: &'a str,
24}
25
26impl<'a> PermissionToolContext<'a> {
27 #[must_use]
28 pub(crate) fn new(name: &'a str, kind: acp::ToolKind, action_label: &'a str) -> Self {
29 Self { name, kind, action_label }
30 }
31}
32
33#[async_trait]
39pub trait AcpPermissionPrompter: Send + Sync {
40 fn permission_options(&self, tool: SupportedTool, args: Option<&Value>) -> Vec<acp::PermissionOption>;
41
42 async fn request_tool_permission(
43 &self,
44 client: &ConnectionHandle,
45 session_id: &acp::SessionId,
46 call: &acp::ToolCall,
47 tool: SupportedTool,
48 args: &Value,
49 ) -> Result<Option<ToolExecutionReport>, SdkError>;
50
51 async fn request_named_tool_permission(
52 &self,
53 client: &ConnectionHandle,
54 session_id: &acp::SessionId,
55 call: &acp::ToolCall,
56 tool: PermissionToolContext<'_>,
57 args: &Value,
58 ) -> Result<Option<ToolExecutionReport>, SdkError>;
59}
60
61pub struct DefaultPermissionPrompter<P> {
62 registry: P,
63}
64
65impl<P> DefaultPermissionPrompter<P>
66where
67 P: ToolRegistryProvider,
68{
69 pub fn new(registry: P) -> Self {
70 Self { registry }
71 }
72
73 fn render_action_label(&self, tool: SupportedTool, args: Option<&Value>) -> String {
74 if let Some(arguments) = args {
75 self.registry
76 .render_title(ToolDescriptor::Acp(tool), tool.function_name(), arguments)
77 } else {
78 tool.default_title().to_string()
79 }
80 }
81
82 fn permission_options_for_action(&self, action_label: &str) -> Vec<acp::PermissionOption> {
83 let allow_once_option = acp::PermissionOption::new(
84 acp::PermissionOptionId::from(Arc::from(TOOL_PERMISSION_ALLOW_OPTION_ID)),
85 format!("{TOOL_PERMISSION_ALLOW_PREFIX} {action_label} once"),
86 acp::PermissionOptionKind::AllowOnce,
87 );
88
89 let allow_always_option = acp::PermissionOption::new(
90 acp::PermissionOptionId::from(Arc::from(TOOL_PERMISSION_ALLOW_ALWAYS_OPTION_ID)),
91 format!("{TOOL_PERMISSION_ALLOW_PREFIX} {action_label} always"),
92 acp::PermissionOptionKind::AllowAlways,
93 );
94
95 let deny_once_option = acp::PermissionOption::new(
96 acp::PermissionOptionId::from(Arc::from(TOOL_PERMISSION_DENY_OPTION_ID)),
97 format!("{TOOL_PERMISSION_DENY_PREFIX} {action_label} once"),
98 acp::PermissionOptionKind::RejectOnce,
99 );
100
101 let deny_always_option = acp::PermissionOption::new(
102 acp::PermissionOptionId::from(Arc::from(TOOL_PERMISSION_DENY_ALWAYS_OPTION_ID)),
103 format!("{TOOL_PERMISSION_DENY_PREFIX} {action_label} always"),
104 acp::PermissionOptionKind::RejectAlways,
105 );
106
107 vec![
108 allow_once_option,
109 allow_always_option,
110 deny_once_option,
111 deny_always_option,
112 ]
113 }
114}
115
116#[async_trait]
117impl<P> AcpPermissionPrompter for DefaultPermissionPrompter<P>
118where
119 P: ToolRegistryProvider + Send + Sync,
120{
121 fn permission_options(&self, tool: SupportedTool, args: Option<&Value>) -> Vec<acp::PermissionOption> {
122 let action_label = self.render_action_label(tool, args);
123 self.permission_options_for_action(&action_label)
124 }
125
126 async fn request_tool_permission(
127 &self,
128 client: &ConnectionHandle,
129 session_id: &acp::SessionId,
130 call: &acp::ToolCall,
131 tool: SupportedTool,
132 args: &Value,
133 ) -> Result<Option<ToolExecutionReport>, SdkError> {
134 let action_label = self.render_action_label(tool, Some(args));
135 self.request_named_tool_permission(
136 client,
137 session_id,
138 call,
139 PermissionToolContext::new(tool.function_name(), tool.kind(), &action_label),
140 args,
141 )
142 .await
143 }
144
145 async fn request_named_tool_permission(
146 &self,
147 client: &ConnectionHandle,
148 session_id: &acp::SessionId,
149 call: &acp::ToolCall,
150 tool: PermissionToolContext<'_>,
151 args: &Value,
152 ) -> Result<Option<ToolExecutionReport>, SdkError> {
153 let fields = acp::ToolCallUpdateFields::default()
154 .title(call.title.clone())
155 .kind(tool.kind)
156 .status(acp::ToolCallStatus::Pending)
157 .raw_input(args.clone());
158
159 let request = acp::RequestPermissionRequest::new(
160 session_id.clone(),
161 acp::ToolCallUpdate::new(call.tool_call_id.clone(), fields),
162 self.permission_options_for_action(tool.action_label),
163 );
164
165 match client.request_permission(request).await {
166 Ok(response) => match response.outcome {
167 acp::RequestPermissionOutcome::Cancelled => {
168 Ok(Some(ToolExecutionReport::failure(tool.name, TOOL_PERMISSION_CANCELLED_MESSAGE)))
169 }
170 acp::RequestPermissionOutcome::Selected(outcome) => {
171 let option_id_str = outcome.option_id.0.as_ref();
172 if option_id_str == TOOL_PERMISSION_ALLOW_OPTION_ID
173 || option_id_str == TOOL_PERMISSION_ALLOW_ALWAYS_OPTION_ID
174 {
175 Ok(None)
176 } else if option_id_str == TOOL_PERMISSION_DENY_OPTION_ID
177 || option_id_str == TOOL_PERMISSION_DENY_ALWAYS_OPTION_ID
178 {
179 Ok(Some(ToolExecutionReport::failure(tool.name, TOOL_PERMISSION_DENIED_MESSAGE)))
180 } else {
181 warn!("{}", TOOL_PERMISSION_UNKNOWN_OPTION_LOG);
182 Ok(Some(ToolExecutionReport::failure(tool.name, TOOL_PERMISSION_DENIED_MESSAGE)))
183 }
184 }
185 _ => Ok(Some(ToolExecutionReport::failure(tool.name, TOOL_PERMISSION_DENIED_MESSAGE))),
186 },
187 Err(error) => {
188 error!(%error, "{}", TOOL_PERMISSION_REQUEST_FAILURE_LOG);
189 Ok(Some(ToolExecutionReport::failure(tool.name, TOOL_PERMISSION_REQUEST_FAILURE_MESSAGE)))
190 }
191 }
192 }
193}
194
195#[cfg(test)]
196mod tests {
197 use super::*;
198 use crate::reports::{
199 TOOL_PERMISSION_ALLOW_OPTION_ID, TOOL_PERMISSION_CANCELLED_MESSAGE, TOOL_PERMISSION_DENIED_MESSAGE,
200 TOOL_PERMISSION_REQUEST_FAILURE_MESSAGE,
201 };
202 use crate::tooling::{AcpToolRegistry, SupportedTool};
203 use crate::zed::connection::ConnectionHandle;
204 use agent_client_protocol::schema::v1::{
205 RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse, SelectedPermissionOutcome,
206 };
207 use agent_client_protocol::{Agent, Channel, Client, ConnectionTo, on_receive_request};
208 use serde_json::json;
209 use std::path::Path;
210 use tokio::sync::oneshot;
211
212 #[derive(Clone, Copy)]
213 enum ClientDecision {
214 Allow,
215 Deny,
216 Cancel,
217 Unknown,
218 RequestFailure,
219 }
220
221 fn test_prompter() -> DefaultPermissionPrompter<AcpToolRegistry> {
222 DefaultPermissionPrompter::new(AcpToolRegistry::new(Path::new("/tmp"), true, true, Vec::new()))
223 }
224
225 async fn run_permission_flow(decision: ClientDecision) -> Option<ToolExecutionReport> {
226 let (agent_channel, client_channel) = Channel::duplex();
227 let (result_tx, result_rx) = oneshot::channel();
228 let session_id = acp::SessionId::new("permission-test-session");
229 let call = acp::ToolCall::new("permission-test-call", "Read file src/lib.rs");
230 let args = json!({ "path": "src/lib.rs" });
231
232 let agent = Agent
233 .builder()
234 .connect_with(agent_channel, async move |cx: ConnectionTo<Client>| {
235 let handle = ConnectionHandle::new(cx);
236 let result = test_prompter()
237 .request_tool_permission(&handle, &session_id, &call, SupportedTool::ReadFile, &args)
238 .await;
239 drop(result_tx.send(result));
240 Ok(())
241 });
242
243 let client = Client
244 .builder()
245 .on_receive_request(
246 async move |request: RequestPermissionRequest, responder, _connection| {
247 assert_eq!(request.options.len(), 4);
248 let response = match decision {
249 ClientDecision::Allow => RequestPermissionResponse::new(RequestPermissionOutcome::Selected(
250 SelectedPermissionOutcome::new(TOOL_PERMISSION_ALLOW_OPTION_ID),
251 )),
252 ClientDecision::Deny => RequestPermissionResponse::new(RequestPermissionOutcome::Selected(
253 SelectedPermissionOutcome::new(TOOL_PERMISSION_DENY_OPTION_ID),
254 )),
255 ClientDecision::Cancel => RequestPermissionResponse::new(RequestPermissionOutcome::Cancelled),
256 ClientDecision::Unknown => RequestPermissionResponse::new(RequestPermissionOutcome::Selected(
257 SelectedPermissionOutcome::new("unsupported-option"),
258 )),
259 ClientDecision::RequestFailure => {
260 return responder.respond_with_internal_error("simulated permission request failure");
261 }
262 };
263 responder.respond(response)
264 },
265 on_receive_request!(),
266 )
267 .connect_to(client_channel);
268
269 let (agent_result, client_result) = tokio::join!(agent, client);
270 agent_result.expect("agent duplex connection should complete");
271 client_result.expect("client duplex connection should complete");
272 result_rx
273 .await
274 .expect("agent should report the permission result")
275 .expect("prompter should return a result")
276 }
277
278 #[tokio::test]
279 async fn permission_allow_flow_returns_no_failure() {
280 assert!(run_permission_flow(ClientDecision::Allow).await.is_none());
281 }
282
283 #[tokio::test]
284 async fn permission_deny_flow_returns_denied_report() {
285 let report = run_permission_flow(ClientDecision::Deny)
286 .await
287 .expect("deny should produce a report");
288 assert!(report.llm_response.contains(TOOL_PERMISSION_DENIED_MESSAGE));
289 }
290
291 #[tokio::test]
292 async fn permission_cancel_flow_returns_cancelled_report() {
293 let report = run_permission_flow(ClientDecision::Cancel)
294 .await
295 .expect("cancel should produce a report");
296 assert!(report.llm_response.contains(TOOL_PERMISSION_CANCELLED_MESSAGE));
297 }
298
299 #[tokio::test]
300 async fn permission_unknown_option_fails_closed() {
301 let report = run_permission_flow(ClientDecision::Unknown)
302 .await
303 .expect("unknown option should be denied");
304 assert!(report.llm_response.contains(TOOL_PERMISSION_DENIED_MESSAGE));
305 }
306
307 #[tokio::test]
308 async fn permission_request_failure_returns_failure_report() {
309 let report = run_permission_flow(ClientDecision::RequestFailure)
310 .await
311 .expect("request failure should produce a report");
312 assert!(report.llm_response.contains(TOOL_PERMISSION_REQUEST_FAILURE_MESSAGE));
313 }
314}