1use crate::events::TraceContext;
2use mcp_utils::client::{
3 McpClient, McpConnectAttempt, McpConnectionAttemptManager, McpError, McpManager, McpServer, McpServerStatusEntry,
4};
5use mcp_utils::display_meta::ToolResultMeta;
6
7use futures::future::Either;
8use futures::stream::{self, StreamExt};
9use llm::{ToolCallError, ToolCallRequest, ToolCallResult};
10use rmcp::RoleClient;
11use rmcp::model::{
12 CallToolRequestParams, CreateElicitationRequestParams, ErrorCode, GetPromptResult, Meta, ProgressNotificationParam,
13 Prompt,
14};
15use rmcp::service::RunningService;
16use std::collections::HashSet;
17use std::sync::Arc;
18use std::time::Duration;
19use tokio::select;
20use tokio::sync::mpsc;
21use tokio::sync::oneshot;
22
23#[derive(Debug)]
25pub enum ToolExecutionEvent {
26 Progress { tool_id: String, progress: ProgressNotificationParam },
27 Complete { tool_id: String, result: Result<ToolCallResult, ToolCallError>, result_meta: Option<ToolResultMeta> },
28}
29
30const MCP_AUTH_TIMEOUT: Duration = Duration::from_mins(3);
31
32#[derive(Debug)]
34pub enum McpCommand {
35 ExecuteTool {
36 request: ToolCallRequest,
37 trace_context: Option<TraceContext>,
38 timeout: Duration,
39 tx: mpsc::Sender<ToolExecutionEvent>,
40 },
41 ListPrompts {
42 tx: oneshot::Sender<Result<Vec<Prompt>, String>>,
43 },
44 GetPrompt {
45 name: String,
46 arguments: Option<serde_json::Map<String, serde_json::Value>>,
47 tx: oneshot::Sender<Result<GetPromptResult, String>>,
48 },
49 GetServerStatuses {
50 tx: oneshot::Sender<Vec<McpServerStatusEntry>>,
51 },
52 AuthenticateServer {
53 name: String,
54 },
55}
56
57pub async fn run_mcp_task(
58 mut mcp: McpManager,
59 mut command_rx: mpsc::Receiver<McpCommand>,
60 pending_servers: Vec<McpServer>,
61) {
62 let mut mcp_connection_attempts = McpConnectionAttemptManager::default();
63 let mut pending_connections: HashSet<String> = pending_servers.iter().map(|server| server.name.clone()).collect();
64 for server in pending_servers {
65 let name = server.name.clone();
66 let task = mcp.connect_pending_task(server);
67 mcp_connection_attempts.spawn(name, task);
68 }
69 if pending_connections.is_empty() {
70 mcp.emit_connection_ready().await;
71 }
72
73 loop {
74 select! {
75 command = command_rx.recv() => {
76 let Some(command) = command else { break; };
77 on_command(command, &mut mcp, &mut mcp_connection_attempts).await;
78 }
79
80 Some(joined) = mcp_connection_attempts.join_next(), if !mcp_connection_attempts.is_empty() => {
81 match joined {
82 Ok(attempt) => {
83 let was_bootstrap = pending_connections.remove(&attempt.name);
84 mcp.apply_connection_attempt(attempt).await;
85 if was_bootstrap && pending_connections.is_empty() {
86 mcp.emit_connection_ready().await;
87 }
88 }
89 Err(e) => tracing::error!("MCP auth task did not complete normally: {e:?}"),
90 }
91 }
92 }
93 }
94
95 mcp_connection_attempts.shutdown().await;
96 mcp.shutdown().await;
97 tracing::debug!("MCP manager task ended");
98}
99
100async fn on_command(command: McpCommand, mcp: &mut McpManager, auth_tasks: &mut McpConnectionAttemptManager) {
101 match command {
102 McpCommand::ExecuteTool { request, trace_context, timeout, tx } => {
103 let tool_id = request.id.clone();
104
105 match mcp.get_client_for_tool(&request.name, &request.arguments) {
106 Ok((client, params)) => {
107 let trace_meta = trace_context.as_ref().map(TraceContext::to_meta);
108 tokio::spawn(async move {
109 let outcome = execute_mcp_call(
110 client,
111 &request,
112 params,
113 trace_meta,
114 timeout,
115 tool_id.clone(),
116 tx.clone(),
117 )
118 .await;
119 let (result, result_meta) = match outcome {
120 Ok((r, m)) => (Ok(r), m),
121 Err(e) => (Err(e), None),
122 };
123 let _ = tx.send(ToolExecutionEvent::Complete { tool_id, result, result_meta }).await;
124 });
125 }
126 Err(e) => {
127 tracing::error!("Failed to get client for tool {}: {e}", request.name);
128 let error = ToolCallError::from_request(&request, format!("Failed to get client: {e}"));
129 let _ =
130 tx.send(ToolExecutionEvent::Complete { tool_id, result: Err(error), result_meta: None }).await;
131 }
132 }
133 }
134
135 McpCommand::ListPrompts { tx } => {
136 let result = mcp.list_prompts().await.map_err(|e| format!("Failed to list prompts: {e}"));
137 let _ = tx.send(result);
138 }
139
140 McpCommand::GetPrompt { name: namespaced_name, arguments, tx } => {
141 let result =
142 mcp.get_prompt(&namespaced_name, arguments).await.map_err(|e| format!("Failed to get prompt: {e}"));
143 let _ = tx.send(result);
144 }
145
146 McpCommand::GetServerStatuses { tx } => {
147 let _ = tx.send(mcp.server_statuses());
148 }
149
150 McpCommand::AuthenticateServer { name } => match mcp.authenticate_server_task(&name).await {
151 Ok(task) => {
152 let server_name = name.clone();
153 auth_tasks.spawn(name, async move {
154 match tokio::time::timeout(MCP_AUTH_TIMEOUT, task).await {
155 Ok(attempt) => attempt,
156 Err(_) => McpConnectAttempt::failed(
157 server_name,
158 McpError::ConnectionFailed("authentication timed out after 3 minutes".to_string()),
159 false,
160 ),
161 }
162 });
163 }
164 Err(e) => tracing::warn!("Authentication failed for '{name}': {e}"),
165 },
166 }
167}
168
169async fn execute_mcp_call(
172 client: Arc<RunningService<RoleClient, McpClient>>,
173 request: &ToolCallRequest,
174 params: CallToolRequestParams,
175 trace_meta: Option<Meta>,
176 timeout: Duration,
177 tool_call_id: String,
178 event_tx: mpsc::Sender<ToolExecutionEvent>,
179) -> Result<(ToolCallResult, Option<ToolResultMeta>), ToolCallError> {
180 use super::tool_bridge::mcp_result_to_tool_call_result;
181 use rmcp::model::{ClientRequest::CallToolRequest, Request, ServerResult};
182 use rmcp::service::PeerRequestOptions;
183
184 let handle = client
185 .send_cancellable_request(CallToolRequest(Request::new(params)), {
186 let mut opts = PeerRequestOptions::default();
187 opts.timeout = Some(timeout);
188 opts.meta = trace_meta;
189 opts
190 })
191 .await
192 .map_err(|e| ToolCallError::from_request(request, format!("Failed to send tool request: {e}")))?;
193
194 let progress_subscriber = client.service().progress_dispatcher.subscribe(handle.progress_token.clone()).await;
195
196 let progress_stream = progress_subscriber
197 .map(move |progress| Either::Left(ToolExecutionEvent::Progress { tool_id: tool_call_id.clone(), progress }));
198
199 let result_stream = stream::once(handle.await_response()).map(Either::Right);
200 let combined_stream = stream::select(progress_stream, result_stream);
201 tokio::pin!(combined_stream);
202
203 let server_result = loop {
204 match combined_stream.next().await {
205 Some(Either::Left(progress_event)) => {
206 let _ = event_tx.send(progress_event).await;
207 }
208 Some(Either::Right(result)) => {
209 break match result {
210 Ok(server_result) => server_result,
211 Err(e) => {
212 if let rmcp::service::ServiceError::McpError(ref error_data) = e
213 && error_data.code == ErrorCode::URL_ELICITATION_REQUIRED
214 {
215 return Err(handle_url_elicitation_required(&client, request, error_data).await);
216 }
217 return Err(ToolCallError::from_request(request, format!("Tool execution failed: {e}")));
218 }
219 };
220 }
221 None => {
222 return Err(ToolCallError::from_request(request, "Stream ended without result"));
223 }
224 }
225 };
226
227 let ServerResult::CallToolResult(mcp_result) = server_result else {
228 return Err(ToolCallError::from_request(request, "Unexpected response type from MCP server"));
229 };
230
231 mcp_result_to_tool_call_result(request, mcp_result)
232}
233
234#[derive(serde::Deserialize)]
235struct UrlElicitationRequiredData {
236 elicitations: Vec<CreateElicitationRequestParams>,
237}
238
239#[derive(Debug)]
240enum UrlElicitationRequiredParseError {
241 MissingData,
242 InvalidData(serde_json::Error),
243 NoUrlRequests,
244}
245
246impl std::fmt::Display for UrlElicitationRequiredParseError {
247 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
248 match self {
249 Self::MissingData => write!(f, "missing error data"),
250 Self::InvalidData(error) => write!(f, "malformed error data: {error}"),
251 Self::NoUrlRequests => write!(f, "provided no URL elicitation requests"),
252 }
253 }
254}
255
256fn parse_required_url_elicitations(
257 error_data: &rmcp::model::ErrorData,
258) -> Result<Vec<CreateElicitationRequestParams>, UrlElicitationRequiredParseError> {
259 let data = error_data.data.as_ref().ok_or(UrlElicitationRequiredParseError::MissingData)?;
260 let parsed: UrlElicitationRequiredData =
261 serde_json::from_value(data.clone()).map_err(UrlElicitationRequiredParseError::InvalidData)?;
262
263 let url_elicitations = parsed
264 .elicitations
265 .into_iter()
266 .filter(|elicitation| matches!(elicitation, CreateElicitationRequestParams::UrlElicitationParams { .. }))
267 .collect::<Vec<_>>();
268
269 if url_elicitations.is_empty() {
270 return Err(UrlElicitationRequiredParseError::NoUrlRequests);
271 }
272
273 Ok(url_elicitations)
274}
275
276async fn handle_url_elicitation_required(
280 client: &Arc<RunningService<RoleClient, McpClient>>,
281 request: &ToolCallRequest,
282 error_data: &rmcp::model::ErrorData,
283) -> ToolCallError {
284 let server_name = client.service().server_name().to_string();
285 let url_elicitations = match parse_required_url_elicitations(error_data) {
286 Ok(url_elicitations) => url_elicitations,
287 Err(UrlElicitationRequiredParseError::NoUrlRequests) => {
288 return ToolCallError::from_request(
289 request,
290 format!("Server '{server_name}' requires URL elicitation but provided no URL elicitation requests"),
291 );
292 }
293 Err(parse_error) => {
294 return ToolCallError::from_request(
295 request,
296 format!("Server '{server_name}' sent an invalid URL elicitation response: {parse_error}"),
297 );
298 }
299 };
300
301 tracing::info!("Server '{server_name}' requires {} URL elicitation(s)", url_elicitations.len());
302
303 for elicitation in url_elicitations {
304 let result = client.service().dispatch_elicitation(elicitation).await;
305 match result.action {
306 rmcp::model::ElicitationAction::Decline => {
307 return ToolCallError::from_request(
308 request,
309 format!("Required browser interaction for server '{server_name}' was declined"),
310 );
311 }
312 rmcp::model::ElicitationAction::Cancel => {
313 return ToolCallError::from_request(
314 request,
315 format!("Required browser interaction for server '{server_name}' was cancelled"),
316 );
317 }
318 rmcp::model::ElicitationAction::Accept => {
319 tracing::info!("User accepted URL elicitation for server '{server_name}'");
320 }
321 }
322 }
323
324 ToolCallError::from_request(
325 request,
326 format!(
327 "Server '{server_name}' requires a browser flow. The URL has been opened for your approval. Retry the previous request after completing the browser flow."
328 ),
329 )
330}
331
332#[cfg(test)]
333mod tests {
334 use super::*;
335
336 #[test]
337 fn url_elicitation_required_data_parses_url_entries() {
338 let data = serde_json::json!({
339 "elicitations": [
340 {
341 "mode": "url",
342 "message": "Auth",
343 "url": "https://example.com/auth?elicitationId=el-1",
344 "elicitationId": "el-1"
345 }
346 ]
347 });
348
349 let parsed: UrlElicitationRequiredData = serde_json::from_value(data).unwrap();
350 assert_eq!(parsed.elicitations.len(), 1);
351 assert!(matches!(
352 &parsed.elicitations[0],
353 CreateElicitationRequestParams::UrlElicitationParams { elicitation_id, .. } if elicitation_id == "el-1"
354 ));
355 }
356
357 #[test]
358 fn parse_required_url_elicitations_filters_to_url_only() {
359 let error_data = rmcp::model::ErrorData {
360 code: rmcp::model::ErrorCode::URL_ELICITATION_REQUIRED,
361 message: "URL elicitation required".into(),
362 data: Some(serde_json::json!({
363 "elicitations": [
364 {
365 "mode": "url",
366 "message": "Auth",
367 "url": "https://example.com/auth",
368 "elicitationId": "el-1"
369 },
370 {
371 "mode": "form",
372 "message": "Pick a color",
373 "requestedSchema": { "type": "object", "properties": {} }
374 }
375 ]
376 })),
377 };
378
379 let result = parse_required_url_elicitations(&error_data).unwrap();
380 assert_eq!(result.len(), 1);
381 assert!(matches!(
382 &result[0],
383 CreateElicitationRequestParams::UrlElicitationParams { elicitation_id, .. } if elicitation_id == "el-1"
384 ));
385 }
386}