Skip to main content

aether_core/mcp/
run_mcp_task.rs

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/// Events emitted during tool execution lifecycle
24#[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/// Commands that can be sent to the MCP manager task
33#[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
169/// Shared logic for sending an MCP tool call, streaming progress events,
170/// and collecting the result.
171async 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
276/// Handle a `URL_ELICITATION_REQUIRED` (-32042) error by dispatching each
277/// URL elicitation through the same consent channel used by normal
278/// `create_elicitation` requests.
279async 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}