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