use serde_json::Value;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::Duration;
use crate::hub_common::HubConnections;
use crate::protocol::McpToolDefinition;
use crate::server::McpServerConfig;
use crate::tool::{BoxFuture, DynTool, McpTool, ToolCallResult, ToolProvider};
use crate::transport::{McpServerConnectionConfig, McpTransportError};
pub struct McpServerHub {
name: String,
connections: HubConnections,
timeout: Duration,
}
impl McpServerHub {
pub fn new(name: impl Into<String>) -> Self {
Self {
name: name.into(),
connections: HubConnections::new(),
timeout: Duration::from_secs(30),
}
}
pub fn with_timeout(name: impl Into<String>, timeout: Duration) -> Self {
Self {
name: name.into(),
connections: HubConnections::new(),
timeout,
}
}
pub async fn connect(
self: &Arc<Self>,
config: McpServerConnectionConfig,
) -> Result<(), McpTransportError> {
let server_name = config.name.clone();
let restart_enabled = config.restart_policy.enabled;
let connection = self.connections.connect(config).await?;
if restart_enabled {
let hub = Arc::clone(self);
let conn = Arc::clone(&connection);
let name = server_name.clone();
tokio::spawn(async move {
hub.restart_monitor(name, conn).await;
});
}
Ok(())
}
async fn restart_monitor(&self, name: String, conn: Arc<crate::hub_common::ManagedConnection>) {
let policy = &conn.config.restart_policy;
loop {
tokio::select! {
_ = conn.restart_notify.notified() => {}
_ = tokio::time::sleep(Duration::from_secs(5)) => {
if conn.is_alive().await {
continue;
}
}
}
if conn.shutdown_requested.load(Ordering::SeqCst) {
break;
}
if conn.is_alive().await {
continue;
}
conn.notify_failure();
let attempt = conn.restart_count.fetch_add(1, Ordering::SeqCst);
if let Some(max) = policy.max_attempts {
if attempt >= max {
eprintln!(
"[McpServerHub] Server '{}' exceeded max restart attempts ({})",
name, max
);
break;
}
}
let delay = policy.delay_for_attempt(attempt);
eprintln!(
"[McpServerHub] Server '{}' disconnected. Restarting in {}ms (attempt {}/{})",
name,
delay,
attempt + 1,
policy
.max_attempts
.map(|m| m.to_string())
.unwrap_or_else(|| "∞".into())
);
tokio::time::sleep(Duration::from_millis(delay)).await;
if conn.shutdown_requested.load(Ordering::SeqCst) {
break;
}
match self.connections.establish_connection(&conn).await {
Ok(_) => {
eprintln!("[McpServerHub] Server '{}' reconnected successfully", name);
conn.restart_count.store(0, Ordering::SeqCst);
}
Err(e) => {
eprintln!(
"[McpServerHub] Server '{}' failed to reconnect: {}",
name, e
);
}
}
}
}
pub fn trigger_restart(&self, server_name: &str) {
if let Some(conn) = self.connections.get(server_name) {
conn.restart_notify.notify_one();
}
}
pub async fn call_tool(&self, name: &str, args: Value) -> Result<Value, McpTransportError> {
self.connections.call_tool(name, args).await
}
pub async fn list_tools(&self) -> Result<Vec<(String, McpToolDefinition)>, McpTransportError> {
Ok(self.connections.list_tools())
}
pub async fn list_all_tools(&self) -> Result<Vec<McpToolDefinition>, McpTransportError> {
Ok(self.connections.list_tool_definitions())
}
pub async fn discover_tools_parallel(
&self,
) -> Result<Vec<(String, McpToolDefinition)>, McpTransportError> {
self.connections.discover_tools_parallel(self.timeout).await
}
pub async fn refresh_tools(&self) -> Result<(), McpTransportError> {
self.connections.refresh_tools_parallel(self.timeout).await
}
pub fn list_servers(&self) -> Vec<String> {
self.connections.list_servers()
}
pub fn is_connected(&self, server_name: &str) -> bool {
self.connections.is_connected(server_name)
}
pub async fn is_alive(&self, server_name: &str) -> bool {
if let Some(conn) = self.connections.get(server_name) {
conn.is_alive().await
} else {
false
}
}
pub async fn health_check(&self) -> Vec<(String, bool)> {
self.connections.health_check().await
}
pub fn server_for_tool(&self, tool_name: &str) -> Option<String> {
self.connections.server_for_tool(tool_name)
}
pub async fn disconnect(&self, server_name: &str) -> Result<(), McpTransportError> {
let connection = self
.connections
.remove(server_name)
.ok_or_else(|| McpTransportError::ServerNotFound(server_name.to_string()))?;
connection.shutdown_requested.store(true, Ordering::SeqCst);
connection.restart_notify.notify_one();
self.connections.clear_tools_for_server(server_name);
if let Some(transport) = connection.get_transport().await {
transport.shutdown().await?;
}
Ok(())
}
pub async fn shutdown_all(&self) -> Result<(), McpTransportError> {
let names: Vec<String> = self.list_servers();
let mut errors = Vec::new();
for name in names {
if let Err(e) = self.disconnect(&name).await {
errors.push(format!("{}: {}", name, e));
}
}
if errors.is_empty() {
Ok(())
} else {
Err(McpTransportError::TransportError(errors.join("; ")))
}
}
pub fn into_config(self, version: &str) -> McpServerConfig {
let hub = Arc::new(self);
let provider = HubToolProvider {
hub: Arc::clone(&hub),
};
McpServerConfig::builder()
.name(&hub.name)
.version(version)
.with_tools_from(provider)
.build()
}
pub fn to_config(self: &Arc<Self>, version: &str) -> McpServerConfig {
let provider = HubToolProvider {
hub: Arc::clone(self),
};
McpServerConfig::builder()
.name(&self.name)
.version(version)
.with_tools_from(provider)
.build()
}
pub fn proxy_tools(self: &Arc<Self>) -> Vec<DynTool> {
let provider = HubToolProvider {
hub: Arc::clone(self),
};
provider.tools()
}
pub fn circuit_breaker_stats(
&self,
server_name: &str,
) -> Option<crate::circuit_breaker::CircuitBreakerStats> {
self.connections.circuit_breaker_stats(server_name)
}
pub fn reset_circuit_breaker(&self, server_name: &str) {
self.connections.reset_circuit_breaker(server_name);
}
}
struct HubToolProvider {
hub: Arc<McpServerHub>,
}
impl ToolProvider for HubToolProvider {
fn tools(&self) -> Vec<DynTool> {
self.hub
.connections
.list_tools()
.into_iter()
.map(|(_, def)| {
let tool: DynTool = Arc::new(ProxyTool {
name: def.name.clone(),
definition: def,
hub: Arc::clone(&self.hub),
});
tool
})
.collect()
}
}
struct ProxyTool {
name: String,
definition: McpToolDefinition,
hub: Arc<McpServerHub>,
}
impl McpTool for ProxyTool {
fn definition(&self) -> McpToolDefinition {
self.definition.clone()
}
fn call<'a>(&'a self, args: Value) -> BoxFuture<'a, ToolCallResult> {
let name = self.name.clone();
let hub = Arc::clone(&self.hub);
Box::pin(async move {
match hub.call_tool(&name, args).await {
Ok(value) => {
if let Some(s) = value.as_str() {
Ok(vec![crate::protocol::ToolContent::text(s)])
} else {
Ok(vec![crate::protocol::ToolContent::text(value.to_string())])
}
}
Err(e) => Err(e.to_string()),
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hub_creation() {
let hub = McpServerHub::new("test-hub");
assert_eq!(hub.name, "test-hub");
assert!(hub.list_servers().is_empty());
}
#[tokio::test]
async fn test_hub_into_config() {
let hub = McpServerHub::new("test-hub");
let config = hub.into_config("1.0.0");
assert_eq!(config.name(), "test-hub");
assert_eq!(config.version(), "1.0.0");
}
#[tokio::test]
async fn test_hub_unknown_tool() {
let hub = McpServerHub::new("test");
let result = hub.call_tool("nonexistent", serde_json::json!({})).await;
assert!(matches!(result, Err(McpTransportError::UnknownTool(_))));
}
}