Skip to main content

mcp_utils/client/
manager.rs

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