1use std::collections::HashMap;
6use std::collections::hash_map::Entry;
7use std::panic::AssertUnwindSafe;
8use std::sync::{Arc, OnceLock, Weak};
9
10use async_trait::async_trait;
11use futures_util::FutureExt;
12use parking_lot::{Mutex, RwLock};
13use tokio_util::sync::CancellationToken;
14use tracing::warn;
15
16pub use crate::rpc::{
17 InstallationConfirmationRequest, InstallationConfirmationResponse, InstallationDecision,
18 InstallationReview, McpInstallationReview,
19};
20use crate::{
21 Client, ClientInner, JsonRpcError, JsonRpcRequest, JsonRpcResponse, Result, error_codes,
22};
23
24pub(crate) const CONFIRM_METHOD: &str = "installations.confirm";
25const REQUEST_CANCELLED: i32 = -32800;
26
27#[derive(Clone)]
33pub struct InstallationConfirmationContext {
34 cancellation: CancellationToken,
35}
36
37impl InstallationConfirmationContext {
38 pub fn cancellation(&self) -> CancellationToken {
42 self.cancellation.child_token()
43 }
44}
45
46#[async_trait]
60pub trait InstallationConfirmationHandler: Send + Sync + 'static {
61 async fn confirm(
63 &self,
64 request: InstallationConfirmationRequest,
65 context: InstallationConfirmationContext,
66 ) -> Result<InstallationDecision>;
67}
68
69#[derive(Default)]
72pub(crate) struct ConfirmationRequests {
73 pending: Mutex<HashMap<u64, Arc<CancellationToken>>>,
74}
75
76impl ConfirmationRequests {
77 pub(crate) fn register(&self, id: u64) -> bool {
78 let mut pending = self.pending.lock();
79 match pending.entry(id) {
80 Entry::Vacant(entry) => {
81 entry.insert(Arc::new(CancellationToken::new()));
82 true
83 }
84 Entry::Occupied(_) => false,
85 }
86 }
87
88 pub(crate) fn cancel(&self, id: u64) {
89 if let Some(token) = self.pending.lock().get(&id) {
90 token.cancel();
91 }
92 }
93
94 pub(crate) fn clear(&self) {
95 self.pending.lock().clear();
96 }
97
98 fn claim(self: &Arc<Self>, id: u64) -> Option<PendingConfirmation> {
99 let cancellation = self.pending.lock().get(&id)?.clone();
100 Some(PendingConfirmation {
101 requests: self.clone(),
102 id,
103 cancellation,
104 })
105 }
106}
107
108struct PendingConfirmation {
109 requests: Arc<ConfirmationRequests>,
110 id: u64,
111 cancellation: Arc<CancellationToken>,
112}
113
114impl Drop for PendingConfirmation {
115 fn drop(&mut self) {
116 let mut pending = self.requests.pending.lock();
117 if pending
118 .get(&self.id)
119 .is_some_and(|token| Arc::ptr_eq(token, &self.cancellation))
120 {
121 pending.remove(&self.id);
122 }
123 }
124}
125
126pub(crate) struct InstallationConfirmationDispatcher {
127 handler: RwLock<Option<Arc<dyn InstallationConfirmationHandler>>>,
128 client: OnceLock<Weak<ClientInner>>,
129}
130
131impl InstallationConfirmationDispatcher {
132 pub(crate) fn new() -> Self {
133 Self {
134 handler: RwLock::new(None),
135 client: OnceLock::new(),
136 }
137 }
138
139 pub(crate) fn set_client(&self, client: Weak<ClientInner>) {
140 let _ = self.client.set(client);
141 }
142
143 pub(crate) fn set_handler(&self, handler: Option<Arc<dyn InstallationConfirmationHandler>>) {
144 *self.handler.write() = handler;
145 }
146
147 pub(crate) fn clear(&self) {
148 self.handler.write().take();
149 }
150
151 pub(crate) fn dispatch(self: &Arc<Self>, request: JsonRpcRequest) {
152 let Some(client) = self.client.get().and_then(Weak::upgrade) else {
153 return;
154 };
155 let Some(pending) = client.rpc.confirmation_requests.claim(request.id) else {
156 warn!("confirmation request retired before dispatch");
157 return;
158 };
159 let request_cancelled = pending.cancellation.as_ref().clone();
160 let connection_closed = client.rpc.connection_closed_token();
161 let context = InstallationConfirmationContext {
162 cancellation: connection_closed.child_token(),
163 };
164 let handler = self.handler.read().clone();
165 let dispatcher = self.clone();
166 tokio::spawn(async move {
167 let outcome = tokio::select! {
168 biased;
169 _ = connection_closed.cancelled() => return,
170 _ = request_cancelled.cancelled() => {
171 context.cancellation.cancel();
172 Err((REQUEST_CANCELLED, "Installation confirmation request cancelled"))
173 }
174 outcome = Self::handle(handler, request.params, context.clone()) => outcome,
175 };
176 if connection_closed.is_cancelled() {
177 return;
178 }
179 let outcome = if request_cancelled.is_cancelled() {
180 context.cancellation.cancel();
181 Err((
182 REQUEST_CANCELLED,
183 "Installation confirmation request cancelled",
184 ))
185 } else {
186 outcome
187 };
188 dispatcher.respond(request.id, outcome).await;
189 drop(pending);
190 });
191 }
192
193 async fn handle(
194 handler: Option<Arc<dyn InstallationConfirmationHandler>>,
195 params: Option<serde_json::Value>,
196 context: InstallationConfirmationContext,
197 ) -> std::result::Result<InstallationConfirmationResponse, (i32, &'static str)> {
198 let Some(handler) = handler else {
199 return Err((
200 error_codes::METHOD_NOT_FOUND,
201 "No installations client-global handler registered",
202 ));
203 };
204 let request: InstallationConfirmationRequest =
205 serde_json::from_value(params.unwrap_or(serde_json::Value::Null)).map_err(|_| {
206 (
207 error_codes::INVALID_PARAMS,
208 "Invalid installation confirmation review",
209 )
210 })?;
211 let confirmation_id = request.confirmation_id.clone();
212 let review_fingerprint = request.review_fingerprint.clone();
213 let outcome = AssertUnwindSafe(handler.confirm(request, context))
214 .catch_unwind()
215 .await;
216 let decision = match outcome {
217 Ok(Ok(
218 decision @ (InstallationDecision::Confirm
219 | InstallationDecision::Decline
220 | InstallationDecision::Cancel),
221 )) => decision,
222 Ok(Ok(InstallationDecision::Unknown)) => {
223 return Err((
224 error_codes::INTERNAL_ERROR,
225 "Invalid installation confirmation decision",
226 ));
227 }
228 Ok(Err(_)) => {
229 return Err((
230 error_codes::INTERNAL_ERROR,
231 "Installation confirmation handler failed",
232 ));
233 }
234 Err(_) => {
235 return Err((
236 error_codes::INTERNAL_ERROR,
237 "Installation confirmation handler panicked",
238 ));
239 }
240 };
241 Ok(InstallationConfirmationResponse {
242 confirmation_id,
243 review_fingerprint,
244 decision,
245 })
246 }
247
248 async fn respond(
249 &self,
250 id: u64,
251 outcome: std::result::Result<InstallationConfirmationResponse, (i32, &'static str)>,
252 ) {
253 let Some(client) = self.client.get().and_then(Weak::upgrade) else {
254 return;
255 };
256 let (result, error) = match outcome {
257 Ok(response) => match serde_json::to_value(response) {
258 Ok(value) => (Some(value), None),
259 Err(_) => {
260 warn!("failed to serialise installation confirmation response");
261 (
262 None,
263 Some(JsonRpcError {
264 code: error_codes::INTERNAL_ERROR,
265 message: "Installation confirmation serialisation failed".to_string(),
266 data: None,
267 }),
268 )
269 }
270 },
271 Err((code, message)) => (
272 None,
273 Some(JsonRpcError {
274 code,
275 message: message.to_string(),
276 data: None,
277 }),
278 ),
279 };
280 if Client::from_inner(client)
281 .send_response(&JsonRpcResponse {
282 jsonrpc: "2.0".to_string(),
283 id,
284 result,
285 error,
286 })
287 .await
288 .is_err()
289 {
290 warn!("failed to send installation confirmation response");
291 }
292 }
293}