Skip to main content

aether_core/mcp/
mcp_handle.rs

1use super::run_mcp_task::ManagerCommand;
2use futures::{Stream, future::join_all};
3use mcp_utils::client::{CallToolError, CallToolOptions, McpError, McpSnapshot, ToolCallEvent, ToolRoute, call_tool};
4use rmcp::model::{GetPromptRequestParams, GetPromptResult, Prompt};
5use serde_json::{Map, Value};
6use std::{pin::Pin, sync::Arc};
7use thiserror::Error;
8use tokio::sync::{mpsc, oneshot, watch};
9
10pub type ToolCallStream = Pin<Box<dyn Stream<Item = ToolCallEvent> + Send>>;
11
12#[derive(Clone)]
13pub struct McpHandle {
14    control_tx: mpsc::Sender<ManagerCommand>,
15    snapshot_rx: watch::Receiver<Arc<McpSnapshot>>,
16}
17
18#[derive(Debug, Error)]
19pub enum McpHandleError {
20    #[error(transparent)]
21    Route(#[from] McpError),
22    #[error("MCP manager is unavailable")]
23    ManagerUnavailable,
24    #[error("failed to list prompts for {server}: {message}")]
25    PromptList { server: String, message: String },
26    #[error("failed to get prompt '{prompt}' from {server}: {message}")]
27    PromptGet { server: String, prompt: String, message: String },
28}
29
30impl McpHandle {
31    pub(super) fn new(
32        control_tx: mpsc::Sender<ManagerCommand>,
33        snapshot_rx: watch::Receiver<Arc<McpSnapshot>>,
34    ) -> Self {
35        Self { control_tx, snapshot_rx }
36    }
37
38    pub fn snapshot(&self) -> Arc<McpSnapshot> {
39        self.snapshot_rx.borrow().clone()
40    }
41
42    pub fn subscribe(&self) -> watch::Receiver<Arc<McpSnapshot>> {
43        self.snapshot_rx.clone()
44    }
45
46    pub fn call(&self, route: ToolRoute, arguments: Map<String, Value>, options: CallToolOptions) -> ToolCallStream {
47        match self.snapshot().resolve(route, arguments) {
48            Ok((client, params)) => Box::pin(call_tool(client, params, options)),
49            Err(error) => Box::pin(futures::stream::once(async move {
50                ToolCallEvent::Complete(Err(CallToolError::Unavailable {
51                    message: format!("Failed to resolve tool: {error}"),
52                }))
53            })),
54        }
55    }
56
57    pub fn call_model_visible(
58        &self,
59        namespaced_name: String,
60        arguments_json: &str,
61        options: CallToolOptions,
62    ) -> ToolCallStream {
63        let parsed = serde_json::from_str::<Value>(arguments_json)
64            .map_err(McpError::from)
65            .and_then(|value| {
66                value
67                    .as_object()
68                    .cloned()
69                    .ok_or_else(|| McpError::JsonError("tool arguments must be a JSON object".to_string()))
70            })
71            .map(|arguments| (ToolRoute::ModelVisible { namespaced_name }, arguments));
72        match parsed {
73            Ok((route, arguments)) => self.call(route, arguments, options),
74            Err(error) => failed_call(error),
75        }
76    }
77
78    pub async fn list_prompts(&self) -> Result<Vec<Prompt>, McpHandleError> {
79        let futures = self.snapshot().clients_with_prompts().into_iter().map(|(server, client)| async move {
80            let prompts = client
81                .list_all_prompts()
82                .await
83                .map_err(|error| McpHandleError::PromptList { server: server.clone(), message: error.to_string() })?;
84            Ok::<_, McpHandleError>(
85                prompts
86                    .into_iter()
87                    .map(|prompt| {
88                        Prompt::new(format!("{server}__{}", prompt.name), prompt.description, prompt.arguments)
89                    })
90                    .collect::<Vec<_>>(),
91            )
92        });
93        let mut prompts = Vec::new();
94        for result in join_all(futures).await {
95            prompts.extend(result?);
96        }
97        Ok(prompts)
98    }
99
100    pub async fn get_prompt(
101        &self,
102        name: &str,
103        arguments: Option<Map<String, Value>>,
104    ) -> Result<GetPromptResult, McpHandleError> {
105        let (prompt, client) = self.snapshot().client_for_prompt(name)?;
106        let server = client.service().server_name().to_string();
107        let mut request = GetPromptRequestParams::new(prompt.clone());
108        if let Some(arguments) = arguments {
109            request = request.with_arguments(arguments);
110        }
111        client.get_prompt(request).await.map_err(|error| McpHandleError::PromptGet {
112            server,
113            prompt,
114            message: error.to_string(),
115        })
116    }
117
118    pub async fn authenticate_server(&self, name: &str) -> Result<(), McpHandleError> {
119        let (tx, rx) = oneshot::channel();
120        self.control_tx
121            .send(ManagerCommand::AuthenticateServer { name: name.to_string(), tx })
122            .await
123            .map_err(|_| McpHandleError::ManagerUnavailable)?;
124        rx.await.map_err(|_| McpHandleError::ManagerUnavailable)?.map_err(McpHandleError::Route)
125    }
126}
127
128fn failed_call(error: McpError) -> ToolCallStream {
129    Box::pin(futures::stream::once(async move {
130        ToolCallEvent::Complete(Err(CallToolError::Unavailable { message: format!("Failed to resolve tool: {error}") }))
131    }))
132}