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#[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#[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
114pub 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 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 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 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 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
575struct 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
636async 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}