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#[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#[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
121pub 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 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 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 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 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
578struct 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
639async 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}