Skip to main content

mcp_utils/client/
manager.rs

1use llm::ToolDefinition;
2
3use super::{
4    McpError, Result,
5    config::{McpHttpConfig, McpServer},
6    connection::{
7        ConnectConfig, McpConnectAttempt, McpConnectOutcome, McpServerConnection, Tool, authenticate_http,
8        connect_server,
9    },
10    mcp_client::McpClient,
11    naming::{create_namespaced_tool_name, split_on_server_name},
12    tool_proxy::ToolProxy,
13};
14use aether_auth::{OAuthCredentialStorage, OAuthHandler};
15use futures::future::join_all;
16use rmcp::{
17    RoleClient,
18    model::{
19        CallToolRequestParams, ClientCapabilities, ClientInfo, CreateElicitationRequestParams, CreateElicitationResult,
20        ElicitationAction, FormElicitationCapability, Implementation, Tool as RmcpTool, UrlElicitationCapability,
21    },
22    service::RunningService,
23};
24use serde::{Deserialize, Serialize};
25use serde_json::Value;
26use std::collections::{BTreeMap, HashMap, HashSet};
27use std::future::Future;
28use std::num::NonZeroU16;
29use std::path::PathBuf;
30use std::sync::Arc;
31use tokio::sync::{mpsc, oneshot};
32use tokio::task::JoinHandle;
33
34pub use crate::status::{McpServerAuthCapability, McpServerStatus, McpServerStatusEntry};
35
36pub const DEFAULT_PROXY_NAME: &str = "proxy";
37
38pub type OAuthHandlerFactory = Arc<dyn Fn(OAuthHandlerContext) -> Result<Arc<dyn OAuthHandler>> + Send + Sync>;
39
40/// Context passed to an `OAuthHandlerFactory` so the constructed handler can
41/// dispatch user-facing prompts back to the host through the MCP event channel.
42#[derive(Clone)]
43pub struct OAuthHandlerContext {
44    pub server_name: String,
45    pub callback_port: Option<NonZeroU16>,
46    pub tx: mpsc::Sender<McpClientEvent>,
47}
48
49#[derive(Debug)]
50pub struct ElicitationRequest {
51    pub server_name: String,
52    pub request: CreateElicitationRequestParams,
53    pub response_sender: oneshot::Sender<CreateElicitationResult>,
54}
55
56#[derive(Debug, Clone)]
57pub struct ElicitationResponse {
58    pub action: ElicitationAction,
59    pub content: Option<Value>,
60}
61
62#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
63pub struct UrlElicitationCompleteParams {
64    pub server_name: String,
65    pub elicitation_id: String,
66}
67
68/// Events emitted by MCP clients that require attention from the host
69/// (e.g. the relay or TUI). Flows through a single channel from `McpManager`
70/// to the consumer.
71#[derive(Debug)]
72pub enum McpClientEvent {
73    Elicitation(ElicitationRequest),
74    UrlElicitationComplete(UrlElicitationCompleteParams),
75    ServerStatusesChanged(Vec<McpServerStatusEntry>),
76    ToolDefinitionsChanged(Vec<ToolDefinition>),
77    ServerInstructionsUpdated { server: String, instructions: Option<String> },
78    AuthenticationFailed { server: String, error: String },
79    ConnectionReady(McpConnectionDetails),
80}
81
82#[derive(Debug, Clone)]
83pub struct McpConnectionDetails {
84    pub instructions: BTreeMap<String, String>,
85    pub tool_definitions: Vec<ToolDefinition>,
86    pub server_statuses: Vec<McpServerStatusEntry>,
87}
88
89/// Manages connections to multiple MCP servers and their tools
90pub struct McpManager {
91    servers: HashMap<String, ServerRecord>,
92    server_order: Vec<String>,
93    proxy: Option<ToolProxy>,
94    aether_home: Option<PathBuf>,
95    client_info: ClientInfo,
96    event_sender: mpsc::Sender<McpClientEvent>,
97    root_dir: PathBuf,
98    oauth_handler_factory: Option<OAuthHandlerFactory>,
99    oauth_credential_store: Option<Arc<dyn OAuthCredentialStorage>>,
100}
101
102impl McpManager {
103    pub fn new(event_sender: mpsc::Sender<McpClientEvent>, oauth_handler_factory: Option<OAuthHandlerFactory>) -> Self {
104        let mut capabilities = ClientCapabilities::builder().enable_elicitation().build();
105        if let Some(elicitation) = capabilities.elicitation.as_mut() {
106            elicitation.form = Some(FormElicitationCapability::default());
107            elicitation.url = Some(UrlElicitationCapability::default());
108        }
109
110        Self {
111            servers: HashMap::new(),
112            server_order: Vec::new(),
113            proxy: None,
114            aether_home: None,
115            client_info: ClientInfo::new(capabilities, Implementation::new("aether", "0.1.0")),
116            event_sender,
117            root_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")),
118            oauth_handler_factory,
119            oauth_credential_store: None,
120        }
121    }
122
123    pub fn with_aether_home(mut self, aether_home: impl Into<PathBuf>) -> Self {
124        self.aether_home = Some(aether_home.into());
125        self
126    }
127
128    pub fn with_oauth_credential_store(mut self, store: Arc<dyn OAuthCredentialStorage>) -> Self {
129        self.oauth_credential_store = Some(store);
130        self
131    }
132
133    pub fn with_root_dir(mut self, root_dir: impl Into<PathBuf>) -> Self {
134        self.root_dir = root_dir.into();
135        self
136    }
137
138    pub async fn register_pending(&mut self, servers: Vec<McpServer>) -> Result<Vec<McpServer>> {
139        let has_proxy = servers.iter().any(|server| server.proxy);
140        if has_proxy && servers.iter().any(|server| server.name == DEFAULT_PROXY_NAME) {
141            return Err(McpError::Other("server name 'proxy' collides with the tool proxy".into()));
142        }
143
144        for server in &servers {
145            self.register_record(&server.name, ServerState::Connecting, None, server.proxy);
146        }
147
148        self.emit_server_statuses_changed().await;
149        Ok(servers)
150    }
151
152    pub async fn bootstrap_proxy_setup(&mut self, servers: &[McpServer]) -> Result<()> {
153        let proxied_members: HashSet<String> =
154            servers.iter().filter(|server| server.proxy).map(|server| server.name.clone()).collect();
155
156        if proxied_members.is_empty() {
157            return Ok(());
158        }
159
160        let dir = self.proxy_tool_dir()?;
161        ToolProxy::clean_dir(&dir).await?;
162        self.register_proxy(dir, proxied_members);
163        Ok(())
164    }
165
166    pub fn connect_pending_task(&self, server: McpServer) -> impl Future<Output = McpConnectAttempt> + Send + 'static {
167        let ctx = self.connect_config();
168        async move { connect_server(server, &ctx).await }
169    }
170
171    pub async fn add_mcps(&mut self, servers: Vec<McpServer>) -> Result<()> {
172        self.bootstrap_proxy_setup(&servers).await?;
173        let pending = self.register_pending(servers).await?;
174        let ctx = self.connect_config();
175        let attempts = join_all(pending.into_iter().map(|server| connect_server(server, &ctx))).await;
176        for attempt in attempts {
177            self.apply_connection_attempt(attempt).await;
178        }
179        Ok(())
180    }
181
182    pub fn get_client_for_tool(
183        &self,
184        namespaced_tool_name: &str,
185        arguments_json: &str,
186    ) -> Result<(Arc<RunningService<RoleClient, McpClient>>, CallToolRequestParams)> {
187        if !self.is_routable_tool(namespaced_tool_name) {
188            return Err(McpError::ToolNotFound(namespaced_tool_name.to_string()));
189        }
190
191        let (server_name, tool_name) = split_on_server_name(namespaced_tool_name)
192            .ok_or_else(|| McpError::InvalidToolNameFormat(namespaced_tool_name.to_string()))?;
193
194        if let Some(proxy) = self.proxy.as_ref().filter(|proxy| proxy.name() == server_name) {
195            let call = proxy.resolve_call(arguments_json)?;
196            let conn = self.connection_for(&call.server).ok_or_else(|| {
197                McpError::ServerNotFound(format!("Proxied server '{}' is not connected", call.server))
198            })?;
199            let params = CallToolRequestParams::new(call.tool).with_arguments(call.arguments.unwrap_or_default());
200            return Ok((conn.client.clone(), params));
201        }
202
203        let client =
204            self.client_for_server(server_name).ok_or_else(|| McpError::ServerNotFound(server_name.to_string()))?;
205
206        let arguments = serde_json::from_str::<serde_json::Value>(arguments_json)?.as_object().cloned();
207        let mut params = CallToolRequestParams::new(tool_name.to_string());
208        if let Some(args) = arguments {
209            params = params.with_arguments(args);
210        }
211
212        Ok((client, params))
213    }
214
215    pub fn tool_definitions(&self) -> Vec<ToolDefinition> {
216        let mut definitions = Vec::new();
217        if let Some(proxy) = self.proxy.as_ref() {
218            definitions.push(ToolProxy::call_tool_definition(proxy.name()));
219        }
220        for name in &self.server_order {
221            let Some(record) = self.servers.get(name) else { continue };
222            if record.proxied {
223                continue;
224            }
225            definitions.extend(record.tools().iter().map(|tool| {
226                ToolDefinition::new(
227                    create_namespaced_tool_name(name, &tool.name),
228                    tool.description.clone(),
229                    tool.parameters.clone(),
230                )
231                .with_server(name.clone())
232                .with_annotations(tool.annotations.clone())
233            }));
234        }
235        definitions
236    }
237
238    pub fn server_instructions(&self) -> BTreeMap<String, String> {
239        let mut instructions: BTreeMap<String, String> = self
240            .servers
241            .iter()
242            .filter(|(name, _)| self.proxy.as_ref().is_none_or(|proxy| !proxy.contains_server(name)))
243            .filter_map(|(name, record)| {
244                record
245                    .connection()
246                    .and_then(|conn| conn.instructions.as_ref())
247                    .map(|instr| (name.clone(), instr.clone()))
248            })
249            .collect();
250
251        if let Some((name, body)) = self.proxy_instructions() {
252            instructions.insert(name, body);
253        }
254
255        instructions
256    }
257
258    pub fn server_statuses(&self) -> Vec<McpServerStatusEntry> {
259        self.server_order
260            .iter()
261            .filter_map(|name| self.servers.get(name).map(|record| record.status_entry(name)))
262            .collect()
263    }
264
265    pub async fn authenticate_server_task(
266        &mut self,
267        name: &str,
268    ) -> Result<impl Future<Output = McpConnectAttempt> + Send + 'static> {
269        let record = self
270            .servers
271            .get(name)
272            .ok_or_else(|| McpError::ConnectionFailed(format!("server '{name}' is not OAuth-authenticatable")))?;
273        if !record.can_authenticate() {
274            return Err(McpError::ConnectionFailed(format!("server '{name}' is not OAuth-authenticatable")));
275        }
276        if self.oauth_handler_factory.is_none() {
277            return Err(McpError::ConnectionFailed(format!("No OAuth handler factory available for '{name}'")));
278        }
279
280        let name = name.to_string();
281        let config = record.reauth_config.clone().expect("checked above");
282        let proxied = record.proxied;
283        let ctx = self.connect_config();
284
285        self.set_state(&name, ServerState::Authenticating);
286        self.emit_server_statuses_changed().await;
287
288        Ok(async move { authenticate_http(name, config, ctx, proxied).await })
289    }
290
291    pub async fn apply_connection_attempt(&mut self, attempt: McpConnectAttempt) {
292        let McpConnectAttempt { name, proxied, outcome } = attempt;
293        match outcome {
294            McpConnectOutcome::Connected { conn, reauth_config } => {
295                match self.register_connection(&name, conn, reauth_config, proxied).await {
296                    Ok(tools) => {
297                        self.refresh_proxy_after_auth(&name, &tools, proxied).await;
298                        self.emit_server_statuses_changed().await;
299                        self.emit_tool_definitions_changed().await;
300                        self.emit_instructions_after_connect(&name, proxied).await;
301                    }
302                    Err(error) => self.apply_authentication_failure(name, error.to_string()).await,
303                }
304            }
305            McpConnectOutcome::NeedsOAuth { config, error } => {
306                tracing::warn!("Server '{name}' needs OAuth: {error}");
307                self.register_record(&name, ServerState::NeedsOAuth, Some(config), proxied);
308                self.emit_server_statuses_changed().await;
309            }
310            McpConnectOutcome::Failed { error } => {
311                self.apply_authentication_failure(name, error.to_string()).await;
312            }
313        }
314    }
315
316    /// List all prompts from all connected MCP servers with namespacing
317    pub async fn list_prompts(&self) -> Result<Vec<rmcp::model::Prompt>> {
318        let futures: Vec<_> = self
319            .servers
320            .iter()
321            .filter_map(|(server_name, record)| {
322                let conn = record.connection()?;
323                conn.client.peer_info()?.capabilities.prompts.as_ref()?;
324                let server_name = server_name.clone();
325                let client = conn.client.clone();
326                Some(async move {
327                    let prompts_response = client.list_prompts(None).await.map_err(|e| {
328                        McpError::PromptListFailed(format!("Failed to list prompts for {server_name}: {e}"))
329                    })?;
330
331                    let namespaced_prompts: Vec<rmcp::model::Prompt> = prompts_response
332                        .prompts
333                        .into_iter()
334                        .map(|prompt| {
335                            let namespaced_name = create_namespaced_tool_name(&server_name, &prompt.name);
336                            rmcp::model::Prompt::new(namespaced_name, prompt.description, prompt.arguments)
337                        })
338                        .collect();
339
340                    Ok::<_, McpError>(namespaced_prompts)
341                })
342            })
343            .collect();
344
345        let results = join_all(futures).await;
346        let mut all_prompts = Vec::new();
347        for result in results {
348            all_prompts.extend(result?);
349        }
350
351        Ok(all_prompts)
352    }
353
354    /// Get a specific prompt by namespaced name
355    pub async fn get_prompt(
356        &self,
357        namespaced_prompt_name: &str,
358        arguments: Option<serde_json::Map<String, serde_json::Value>>,
359    ) -> Result<rmcp::model::GetPromptResult> {
360        let (server_name, prompt_name) = split_on_server_name(namespaced_prompt_name)
361            .ok_or_else(|| McpError::InvalidToolNameFormat(namespaced_prompt_name.to_string()))?;
362
363        let server_conn =
364            self.connection_for(server_name).ok_or_else(|| McpError::ServerNotFound(server_name.to_string()))?;
365
366        let mut request = rmcp::model::GetPromptRequestParams::new(prompt_name);
367        if let Some(args) = arguments {
368            request = request.with_arguments(args);
369        }
370
371        server_conn.client.get_prompt(request).await.map_err(|e| {
372            McpError::PromptGetFailed(format!("Failed to get prompt '{prompt_name}' from {server_name}: {e}"))
373        })
374    }
375
376    /// Shutdown all servers and wait for their tasks to complete
377    pub async fn shutdown(&mut self) {
378        let servers: Vec<(String, ServerRecord)> = self.servers.drain().collect();
379
380        for (server_name, record) in servers {
381            if let Some(conn) = record.into_connection()
382                && let Some(handle) = conn.server_task
383            {
384                drop(conn.client);
385                await_server_shutdown(&server_name, handle).await;
386            }
387        }
388
389        self.server_order.clear();
390        self.proxy = None;
391    }
392
393    /// Shutdown a specific server by name
394    pub async fn shutdown_server(&mut self, server_name: &str) -> Result<()> {
395        if let Some(record) = self.servers.remove(server_name)
396            && let Some(conn) = record.into_connection()
397            && let Some(handle) = conn.server_task
398        {
399            drop(conn.client);
400            await_server_shutdown(server_name, handle).await;
401        }
402
403        Ok(())
404    }
405
406    async fn emit_server_statuses_changed(&self) {
407        self.emit_event(McpClientEvent::ServerStatusesChanged(self.server_statuses())).await;
408    }
409
410    async fn emit_tool_definitions_changed(&self) {
411        self.emit_event(McpClientEvent::ToolDefinitionsChanged(self.tool_definitions())).await;
412    }
413
414    async fn emit_instructions_after_connect(&self, server_name: &str, proxied: bool) {
415        if proxied {
416            if let Some((server, body)) = self.proxy_instructions() {
417                self.emit_event(McpClientEvent::ServerInstructionsUpdated { server, instructions: Some(body) }).await;
418            }
419            return;
420        }
421
422        if let Some(instructions) =
423            self.connection_for(server_name).and_then(|conn| conn.instructions.as_ref()).cloned()
424        {
425            self.emit_event(McpClientEvent::ServerInstructionsUpdated {
426                server: server_name.to_string(),
427                instructions: Some(instructions),
428            })
429            .await;
430        }
431    }
432
433    fn proxy_instructions(&self) -> Option<(String, String)> {
434        let proxy = self.proxy.as_ref()?;
435        let descriptions: Vec<(String, String)> = proxy
436            .members()
437            .iter()
438            .filter_map(|member| {
439                let conn = self.connection_for(member)?;
440                Some((member.clone(), ToolProxy::extract_server_description(&conn.client, member)))
441            })
442            .collect();
443        Some((proxy.name().to_string(), ToolProxy::build_instructions(proxy.tool_dir(), &descriptions)))
444    }
445
446    pub async fn emit_connection_ready(&self) {
447        self.emit_event(McpClientEvent::ConnectionReady(McpConnectionDetails {
448            tool_definitions: self.tool_definitions(),
449            instructions: self.server_instructions(),
450            server_statuses: self.server_statuses(),
451        }))
452        .await;
453    }
454
455    async fn emit_authentication_failed(&self, server: String, error: String) {
456        self.emit_event(McpClientEvent::AuthenticationFailed { server, error }).await;
457    }
458
459    async fn emit_event(&self, event: McpClientEvent) {
460        if let Err(e) = self.event_sender.send(event).await {
461            tracing::warn!("Failed to emit MCP client event: {e}");
462        }
463    }
464
465    fn connect_config(&self) -> Arc<ConnectConfig> {
466        Arc::new(ConnectConfig {
467            client_info: self.client_info.clone(),
468            event_sender: self.event_sender.clone(),
469            root_dir: self.root_dir.clone(),
470            oauth_handler_factory: self.oauth_handler_factory.clone(),
471            oauth_credential_store: self.oauth_credential_store.clone(),
472        })
473    }
474
475    fn proxy_tool_dir(&self) -> Result<PathBuf> {
476        self.aether_home
477            .as_ref()
478            .map(|home| ToolProxy::dir_in_home(home, DEFAULT_PROXY_NAME))
479            .map_or_else(|| ToolProxy::dir(DEFAULT_PROXY_NAME), Ok)
480    }
481
482    async fn register_connection(
483        &mut self,
484        name: &str,
485        conn: McpServerConnection,
486        reauth_config: Option<McpHttpConfig>,
487        proxied: bool,
488    ) -> Result<Vec<RmcpTool>> {
489        let tools = conn
490            .list_tools()
491            .await
492            .map_err(|e| McpError::ToolDiscoveryFailed(format!("Failed to list tools for {name}: {e}")))?;
493        self.apply_connected(name, conn, &tools, reauth_config, proxied);
494        Ok(tools)
495    }
496
497    fn apply_connected(
498        &mut self,
499        name: &str,
500        conn: McpServerConnection,
501        tools: &[RmcpTool],
502        reauth_config: Option<McpHttpConfig>,
503        proxied: bool,
504    ) {
505        let existing_reauth = self.servers.get(name).and_then(|r| r.reauth_config.clone());
506        let final_reauth = reauth_config.or(existing_reauth);
507        let tools = tools.iter().map(Tool::from).collect();
508
509        self.remember_server_order(name);
510        self.servers.insert(name.to_string(), ServerRecord::connected(conn, tools, final_reauth, proxied));
511    }
512
513    fn register_proxy(&mut self, tool_dir: std::path::PathBuf, members: HashSet<String>) {
514        self.proxy = Some(ToolProxy::new(DEFAULT_PROXY_NAME.to_string(), members, tool_dir));
515    }
516
517    async fn refresh_proxy_after_auth(&mut self, name: &str, tools: &[RmcpTool], proxied: bool) {
518        if !proxied {
519            return;
520        }
521
522        if let Some(proxy) = self.proxy.as_mut() {
523            proxy.add_member(name.to_string());
524        }
525
526        if let Some(tool_dir) = self.proxy.as_ref().map(|proxy| proxy.tool_dir().to_path_buf())
527            && let Err(e) = ToolProxy::write_tool_entries_to_dir(name, tools, &tool_dir).await
528        {
529            tracing::warn!("Failed to write tool files for '{name}' after OAuth: {e}");
530        }
531    }
532
533    fn remember_server_order(&mut self, name: &str) {
534        if !self.server_order.iter().any(|n| n == name) {
535            self.server_order.push(name.to_string());
536        }
537    }
538
539    async fn apply_authentication_failure(&mut self, name: String, error: String) {
540        self.set_state(&name, ServerState::Failed { error: error.clone() });
541        self.emit_server_statuses_changed().await;
542        self.emit_authentication_failed(name, error).await;
543    }
544
545    fn set_state(&mut self, name: &str, state: ServerState) {
546        self.remember_server_order(name);
547        match self.servers.get_mut(name) {
548            Some(record) => record.state = state,
549            None => {
550                self.servers.insert(name.to_string(), ServerRecord::new(state, None, false));
551            }
552        }
553    }
554
555    fn register_record(&mut self, name: &str, state: ServerState, reauth_config: Option<McpHttpConfig>, proxied: bool) {
556        self.remember_server_order(name);
557        self.servers.insert(name.to_string(), ServerRecord::new(state, reauth_config, proxied));
558    }
559
560    fn connection_for(&self, server_name: &str) -> Option<&McpServerConnection> {
561        self.servers.get(server_name).and_then(ServerRecord::connection)
562    }
563
564    fn client_for_server(&self, server_name: &str) -> Option<Arc<RunningService<RoleClient, McpClient>>> {
565        self.connection_for(server_name).map(|conn| conn.client.clone())
566    }
567
568    fn is_routable_tool(&self, namespaced_tool_name: &str) -> bool {
569        if self.proxy.as_ref().is_some_and(|proxy| proxy.call_tool_name() == namespaced_tool_name) {
570            return true;
571        }
572        match split_on_server_name(namespaced_tool_name) {
573            Some((server_name, tool_name)) => {
574                self.servers.get(server_name).is_some_and(|record| !record.proxied && record.has_tool(tool_name))
575            }
576            None => false,
577        }
578    }
579}
580
581impl Drop for McpManager {
582    fn drop(&mut self) {
583        let servers: Vec<(String, ServerRecord)> = self.servers.drain().collect();
584        for (server_name, record) in servers {
585            if let Some(conn) = record.into_connection()
586                && let Some(handle) = conn.server_task
587            {
588                handle.abort();
589                tracing::warn!("Server '{server_name}' task aborted during cleanup");
590            }
591        }
592    }
593}
594
595/// Internal record holding all mutable state for a single MCP server.
596struct ServerRecord {
597    state: ServerState,
598    reauth_config: Option<McpHttpConfig>,
599    proxied: bool,
600}
601
602enum ServerState {
603    Connecting,
604    Connected { connection: McpServerConnection, tools: Vec<Tool> },
605    Authenticating,
606    Failed { error: String },
607    NeedsOAuth,
608}
609
610impl From<&ServerState> for McpServerStatus {
611    fn from(state: &ServerState) -> Self {
612        match state {
613            ServerState::Connecting => Self::Connecting,
614            ServerState::Connected { tools, .. } => Self::Connected { tool_count: tools.len() },
615            ServerState::Authenticating => Self::Authenticating,
616            ServerState::Failed { error } => Self::Failed { error: error.clone() },
617            ServerState::NeedsOAuth => Self::NeedsOAuth,
618        }
619    }
620}
621
622impl ServerRecord {
623    fn new(state: ServerState, reauth_config: Option<McpHttpConfig>, proxied: bool) -> Self {
624        Self { state, reauth_config, proxied }
625    }
626
627    fn connected(
628        connection: McpServerConnection,
629        tools: Vec<Tool>,
630        reauth_config: Option<McpHttpConfig>,
631        proxied: bool,
632    ) -> Self {
633        Self::new(ServerState::Connected { connection, tools }, reauth_config, proxied)
634    }
635
636    fn tools(&self) -> &[Tool] {
637        match &self.state {
638            ServerState::Connected { tools, .. } => tools,
639            ServerState::Connecting
640            | ServerState::Authenticating
641            | ServerState::Failed { .. }
642            | ServerState::NeedsOAuth => &[],
643        }
644    }
645
646    fn has_tool(&self, tool_name: &str) -> bool {
647        self.tools().iter().any(|tool| tool.name == tool_name)
648    }
649
650    fn connection(&self) -> Option<&McpServerConnection> {
651        match &self.state {
652            ServerState::Connected { connection, .. } => Some(connection),
653            ServerState::Connecting
654            | ServerState::Authenticating
655            | ServerState::Failed { .. }
656            | ServerState::NeedsOAuth => None,
657        }
658    }
659
660    fn into_connection(self) -> Option<McpServerConnection> {
661        match self.state {
662            ServerState::Connected { connection, .. } => Some(connection),
663            ServerState::Connecting
664            | ServerState::Authenticating
665            | ServerState::Failed { .. }
666            | ServerState::NeedsOAuth => None,
667        }
668    }
669
670    fn auth_capability(&self) -> McpServerAuthCapability {
671        if self.reauth_config.is_some() { McpServerAuthCapability::OAuth } else { McpServerAuthCapability::Unavailable }
672    }
673
674    fn can_authenticate(&self) -> bool {
675        self.reauth_config.is_some()
676    }
677
678    fn status(&self) -> McpServerStatus {
679        (&self.state).into()
680    }
681
682    fn status_entry(&self, name: &str) -> McpServerStatusEntry {
683        McpServerStatusEntry::new(name, self.status())
684            .with_auth_capability(self.auth_capability())
685            .with_proxied(self.proxied)
686    }
687}
688
689/// Awaits `handle` for up to 5 seconds, logging whether the server shut down
690/// gracefully, panicked, or timed out. Used during manager teardown.
691async fn await_server_shutdown(server_name: &str, handle: JoinHandle<()>) {
692    let Ok(task_result) = tokio::time::timeout(std::time::Duration::from_secs(5), handle).await else {
693        tracing::warn!("Server '{server_name}' shutdown timed out");
694        return;
695    };
696    match task_result {
697        Ok(()) => tracing::info!("Server '{server_name}' shut down gracefully"),
698        Err(e) => tracing::warn!("Server '{server_name}' task panicked: {e:?}"),
699    }
700}
701
702#[cfg(test)]
703mod tests {
704    use super::{DEFAULT_PROXY_NAME, McpClientEvent, McpManager, McpServerStatus, ServerState};
705    use crate::client::OAuthHandlerFactory;
706    use crate::client::config::{McpHttpConfig, McpServer, McpTransport};
707    use crate::client::connection::{McpConnectAttempt, McpConnectOutcome};
708    use crate::status::McpServerAuthCapability;
709    use aether_auth::{OAuthCallback, OAuthError, OAuthHandler};
710    use futures::future::BoxFuture;
711    use rmcp::{
712        Json, RoleServer, ServerHandler,
713        handler::server::{router::tool::ToolRouter, wrapper::Parameters},
714        model::{Implementation, ServerCapabilities, ServerInfo},
715        service::DynService,
716        tool, tool_handler, tool_router,
717        transport::streamable_http_client::StreamableHttpClientTransportConfig,
718    };
719    use schemars::JsonSchema;
720    use serde::{Deserialize, Serialize};
721    use std::{
722        io,
723        sync::{Arc, Mutex},
724    };
725    use tokio::sync::mpsc;
726
727    #[derive(Clone)]
728    struct TestServer {
729        tool_router: ToolRouter<Self>,
730    }
731
732    #[tool_handler(router = self.tool_router)]
733    impl ServerHandler for TestServer {
734        fn get_info(&self) -> ServerInfo {
735            ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
736                .with_server_info(Implementation::new("test-server", "0.1.0").with_description("Test MCP server"))
737                .with_instructions("Test server instructions")
738        }
739    }
740
741    impl Default for TestServer {
742        fn default() -> Self {
743            Self { tool_router: Self::tool_router() }
744        }
745    }
746
747    #[derive(Debug, Deserialize, Serialize, JsonSchema)]
748    struct EchoRequest {
749        value: String,
750    }
751
752    #[derive(Debug, Deserialize, Serialize, JsonSchema)]
753    struct EchoResult {
754        value: String,
755    }
756
757    #[tool_router]
758    impl TestServer {
759        fn into_dyn(self) -> Box<dyn DynService<RoleServer>> {
760            Box::new(self)
761        }
762
763        #[tool(description = "Returns the provided value", annotations(read_only_hint = true, open_world_hint = false))]
764        async fn echo(&self, request: Parameters<EchoRequest>) -> Json<EchoResult> {
765            let Parameters(EchoRequest { value }) = request;
766            Json(EchoResult { value })
767        }
768    }
769
770    #[derive(Clone)]
771    struct SharedWriter(Arc<Mutex<Vec<u8>>>);
772
773    impl io::Write for SharedWriter {
774        fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
775            self.0.lock().unwrap().extend_from_slice(buf);
776            Ok(buf.len())
777        }
778
779        fn flush(&mut self) -> io::Result<()> {
780            Ok(())
781        }
782    }
783
784    struct TestOAuthHandler;
785
786    impl OAuthHandler for TestOAuthHandler {
787        fn redirect_uri(&self) -> &'static str {
788            "http://127.0.0.1:0/oauth2callback"
789        }
790
791        fn authorize(&self, _auth_url: &str) -> BoxFuture<'_, Result<OAuthCallback, OAuthError>> {
792            Box::pin(async { Err(OAuthError::UserCancelled) })
793        }
794    }
795
796    fn test_oauth_handler_factory() -> OAuthHandlerFactory {
797        Arc::new(|_ctx| Ok(Arc::new(TestOAuthHandler)))
798    }
799
800    fn http_config(uri: &str) -> McpHttpConfig {
801        StreamableHttpClientTransportConfig::with_uri(uri).into()
802    }
803
804    #[tokio::test]
805    async fn authenticate_server_task_rejects_record_without_reauth_config() {
806        let (event_sender, _event_receiver) = mpsc::channel(1);
807        let mut manager = McpManager::new(event_sender, Some(test_oauth_handler_factory()));
808        manager.register_record("public", ServerState::Connecting, None, false);
809
810        let error = match manager.authenticate_server_task("public").await {
811            Ok(_) => panic!("non-OAuth server should be rejected"),
812            Err(error) => error.to_string(),
813        };
814        assert!(error.contains("not OAuth-authenticatable"));
815    }
816
817    #[tokio::test]
818    async fn authenticate_server_task_marks_server_authenticating_and_emits_status() {
819        let (event_sender, mut event_receiver) = mpsc::channel(2);
820        let mut manager = McpManager::new(event_sender, Some(test_oauth_handler_factory()));
821        manager.register_record(
822            "remote",
823            ServerState::NeedsOAuth,
824            Some(http_config("http://localhost:19999/mcp")),
825            false,
826        );
827
828        let _task = manager.authenticate_server_task("remote").await.expect("auth should start");
829
830        assert!(matches!(manager.server_statuses()[0].status, McpServerStatus::Authenticating));
831        let event = event_receiver.recv().await.expect("status change event");
832        let McpClientEvent::ServerStatusesChanged(servers) = event else {
833            panic!("expected ServerStatusesChanged");
834        };
835        let status = servers.iter().find(|entry| entry.name == "remote").expect("remote status");
836        assert!(matches!(status.status, McpServerStatus::Authenticating));
837        assert_eq!(status.auth_capability, McpServerAuthCapability::OAuth);
838    }
839
840    #[tokio::test]
841    async fn apply_connection_attempt_failure_allows_retry() {
842        let (event_sender, mut event_receiver) = mpsc::channel(2);
843        let mut manager = McpManager::new(event_sender, Some(test_oauth_handler_factory()));
844        manager.register_record(
845            "remote",
846            ServerState::NeedsOAuth,
847            Some(http_config("http://localhost:19999/mcp")),
848            false,
849        );
850        let _task = manager.authenticate_server_task("remote").await.expect("auth should start");
851        let _authenticating_event = event_receiver.recv().await.expect("authenticating status change event");
852
853        manager
854            .apply_connection_attempt(McpConnectAttempt {
855                name: "remote".to_string(),
856                proxied: false,
857                outcome: McpConnectOutcome::Failed {
858                    error: crate::client::McpError::ConnectionFailed("boom".to_string()),
859                },
860            })
861            .await;
862
863        let event = event_receiver.recv().await.expect("status change event");
864        let McpClientEvent::ServerStatusesChanged(servers) = event else {
865            panic!("expected ServerStatusesChanged");
866        };
867        let auth_event = event_receiver.recv().await.expect("authentication failure event");
868        let McpClientEvent::AuthenticationFailed { server, error } = auth_event else {
869            panic!("expected AuthenticationFailed");
870        };
871        assert_eq!(server, "remote");
872        assert!(error.contains("boom"));
873
874        let status = servers.iter().find(|entry| entry.name == "remote").expect("remote status");
875        assert_eq!(status.auth_capability, McpServerAuthCapability::OAuth);
876        assert!(matches!(status.status, McpServerStatus::Failed { ref error } if error.contains("boom")));
877        assert!(manager.authenticate_server_task("remote").await.is_ok());
878    }
879
880    #[test]
881    fn status_entries_are_derived_from_reauth_config() {
882        let (event_sender, _event_receiver) = mpsc::channel(1);
883        let mut manager = McpManager::new(event_sender, Some(test_oauth_handler_factory()));
884
885        manager.register_record(
886            "with-oauth",
887            ServerState::Connecting,
888            Some(http_config("http://localhost/mcp")),
889            false,
890        );
891        manager.register_record("without-oauth", ServerState::Connecting, None, false);
892        manager.register_record(
893            "needs-oauth",
894            ServerState::NeedsOAuth,
895            Some(http_config("http://localhost/mcp2")),
896            false,
897        );
898
899        let statuses = manager.server_statuses();
900        let with_oauth = statuses.iter().find(|s| s.name == "with-oauth").unwrap();
901        let without_oauth = statuses.iter().find(|s| s.name == "without-oauth").unwrap();
902        let needs_oauth = statuses.iter().find(|s| s.name == "needs-oauth").unwrap();
903
904        assert_eq!(with_oauth.auth_capability, McpServerAuthCapability::OAuth);
905        assert_eq!(without_oauth.auth_capability, McpServerAuthCapability::Unavailable);
906        assert_eq!(needs_oauth.auth_capability, McpServerAuthCapability::OAuth);
907    }
908
909    #[tokio::test]
910    async fn register_pending_marks_every_server_connecting_and_emits_status() {
911        let (event_sender, mut event_receiver) = mpsc::channel(32);
912        let mut manager = McpManager::new(event_sender, None);
913
914        let servers = vec![
915            McpServer::new("alpha", McpTransport::InMemory { server: TestServer::default().into_dyn() }, false),
916            McpServer::new("beta", McpTransport::InMemory { server: TestServer::default().into_dyn() }, true),
917        ];
918
919        let returned = manager.register_pending(servers).await.unwrap();
920        assert_eq!(returned.iter().map(|s| s.name.as_str()).collect::<Vec<_>>(), vec!["alpha", "beta"]);
921
922        let statuses = manager.server_statuses();
923        assert_eq!(statuses.len(), 2);
924        assert!(matches!(statuses.iter().find(|s| s.name == "alpha").unwrap().status, McpServerStatus::Connecting));
925        assert!(matches!(statuses.iter().find(|s| s.name == "beta").unwrap().status, McpServerStatus::Connecting));
926        assert!(statuses.iter().find(|s| s.name == "beta").unwrap().proxied);
927
928        let event = event_receiver.try_recv().expect("expected initial ServerStatusesChanged emission");
929        let McpClientEvent::ServerStatusesChanged(emitted) = event else {
930            panic!("expected ServerStatusesChanged, got {event:?}");
931        };
932        assert_eq!(emitted.iter().map(|s| s.name.as_str()).collect::<Vec<_>>(), vec!["alpha", "beta"]);
933    }
934
935    #[tokio::test]
936    async fn apply_connection_attempt_emits_instructions_updated_after_connect() {
937        let (event_sender, mut event_receiver) = mpsc::channel(32);
938        let mut manager = McpManager::new(event_sender, None);
939
940        let servers =
941            vec![McpServer::new("test", McpTransport::InMemory { server: TestServer::default().into_dyn() }, false)];
942        manager.add_mcps(servers).await.unwrap();
943
944        let mut update_for_test = None;
945        while let Ok(event) = event_receiver.try_recv() {
946            if let McpClientEvent::ServerInstructionsUpdated { server, instructions } = event
947                && server == "test"
948            {
949                update_for_test = Some(instructions);
950            }
951        }
952        let instructions = update_for_test.expect("expected ServerInstructionsUpdated for 'test'");
953        assert!(instructions.is_some(), "TestServer publishes instructions, so update should carry Some(_)");
954    }
955
956    #[tokio::test]
957    async fn server_statuses_mark_direct_and_proxied_servers_without_proxy_row() {
958        let (event_sender, _event_receiver) = mpsc::channel(32);
959        let mut manager = McpManager::new(event_sender, None);
960        manager
961            .add_mcps(vec![
962                McpServer::new("direct", McpTransport::InMemory { server: TestServer::default().into_dyn() }, false),
963                McpServer::new("math", McpTransport::InMemory { server: TestServer::default().into_dyn() }, true),
964            ])
965            .await
966            .unwrap();
967
968        let statuses = manager.server_statuses();
969        assert_eq!(statuses.iter().map(|status| status.name.as_str()).collect::<Vec<_>>(), vec!["direct", "math"]);
970        assert!(!statuses.iter().find(|status| status.name == "direct").unwrap().proxied);
971        assert!(statuses.iter().find(|status| status.name == "math").unwrap().proxied);
972        assert!(!statuses.iter().any(|status| status.name == DEFAULT_PROXY_NAME));
973    }
974
975    #[tokio::test]
976    async fn tool_definitions_drop_when_a_server_shuts_down() {
977        let (event_sender, _event_receiver) = mpsc::channel(32);
978        let mut manager = McpManager::new(event_sender, None);
979        manager
980            .add_mcps(vec![
981                McpServer::new("git", McpTransport::InMemory { server: TestServer::default().into_dyn() }, false),
982                McpServer::new("github", McpTransport::InMemory { server: TestServer::default().into_dyn() }, false),
983            ])
984            .await
985            .unwrap();
986
987        let names =
988            |manager: &McpManager| manager.tool_definitions().into_iter().map(|tool| tool.name).collect::<Vec<_>>();
989        assert!(names(&manager).contains(&"git__echo".to_string()));
990        assert!(names(&manager).contains(&"github__echo".to_string()));
991
992        manager.shutdown_server("git").await.unwrap();
993
994        assert!(!names(&manager).iter().any(|name| name.starts_with("git__")));
995        assert!(names(&manager).contains(&"github__echo".to_string()));
996    }
997
998    #[tokio::test]
999    async fn tool_definitions_preserve_annotations() {
1000        let (event_sender, _event_receiver) = mpsc::channel(32);
1001        let mut manager = McpManager::new(event_sender, None);
1002        manager
1003            .add_mcps(vec![McpServer::new(
1004                "test",
1005                McpTransport::InMemory { server: TestServer::default().into_dyn() },
1006                false,
1007            )])
1008            .await
1009            .unwrap();
1010
1011        let tools = manager.tool_definitions();
1012        let echo = tools.iter().find(|tool| tool.name == "test__echo").expect("echo tool");
1013        let annotations = echo.annotations.as_ref().expect("annotations should be preserved");
1014        assert_eq!(annotations.read_only_hint, Some(true));
1015        assert_eq!(annotations.open_world_hint, Some(false));
1016    }
1017
1018    #[tokio::test]
1019    async fn drop_logs_cleanup_abort_with_tracing() {
1020        let (event_sender, _event_receiver) = mpsc::channel(32);
1021        let mut manager = McpManager::new(event_sender, None);
1022        manager
1023            .add_mcps(vec![McpServer::new(
1024                "test",
1025                McpTransport::InMemory { server: TestServer::default().into_dyn() },
1026                false,
1027            )])
1028            .await
1029            .unwrap();
1030
1031        let output = Arc::new(Mutex::new(Vec::new()));
1032        let subscriber = tracing_subscriber::fmt()
1033            .with_ansi(false)
1034            .without_time()
1035            .with_writer({
1036                let output = Arc::clone(&output);
1037                move || SharedWriter(Arc::clone(&output))
1038            })
1039            .finish();
1040
1041        tracing::subscriber::with_default(subscriber, || {
1042            drop(manager);
1043        });
1044
1045        let logs = String::from_utf8(output.lock().unwrap().clone()).unwrap();
1046        assert!(logs.contains("Server 'test' task aborted during cleanup"));
1047    }
1048}