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#[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#[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
89pub 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 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 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 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 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
595struct 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
689async 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}