use dashmap::DashMap;
use futures::future::join_all;
use serde_json::Value;
use std::sync::atomic::{AtomicBool, AtomicU32};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{broadcast, Notify};
use crate::circuit_breaker::CircuitBreaker;
use crate::protocol::McpToolDefinition;
use crate::transport::{McpServerConnectionConfig, McpTransport, McpTransportError};
use crate::transport_factory::TransportFactory;
pub struct ManagedConnection {
pub config: McpServerConnectionConfig,
pub transport: tokio::sync::RwLock<Option<Arc<dyn McpTransport>>>,
pub circuit_breaker: CircuitBreaker,
pub restart_count: AtomicU32,
pub shutdown_requested: AtomicBool,
pub restart_notify: Notify,
pub failure_tx: broadcast::Sender<()>,
}
impl ManagedConnection {
pub fn new(config: McpServerConnectionConfig) -> Self {
let (failure_tx, _) = broadcast::channel(16);
Self {
config,
transport: tokio::sync::RwLock::new(None),
circuit_breaker: CircuitBreaker::new(),
restart_count: AtomicU32::new(0),
shutdown_requested: AtomicBool::new(false),
restart_notify: Notify::new(),
failure_tx,
}
}
pub async fn is_alive(&self) -> bool {
if let Some(transport) = self.transport.read().await.as_ref() {
transport.is_alive()
} else {
false
}
}
pub async fn get_transport(&self) -> Option<Arc<dyn McpTransport>> {
self.transport.read().await.clone()
}
pub fn subscribe_failures(&self) -> broadcast::Receiver<()> {
self.failure_tx.subscribe()
}
pub fn notify_failure(&self) {
let _ = self.failure_tx.send(());
}
}
pub struct HubConnections {
connections: DashMap<String, Arc<ManagedConnection>>,
tool_cache: DashMap<String, (String, Option<McpToolDefinition>)>,
}
impl Default for HubConnections {
fn default() -> Self {
Self::new()
}
}
impl HubConnections {
pub fn new() -> Self {
Self {
connections: DashMap::new(),
tool_cache: DashMap::new(),
}
}
pub async fn connect(
&self,
config: McpServerConnectionConfig,
) -> Result<Arc<ManagedConnection>, McpTransportError> {
let server_name = config.name.clone();
let connection = Arc::new(ManagedConnection::new(config));
self.establish_connection(&connection).await?;
self.connections
.insert(server_name, Arc::clone(&connection));
Ok(connection)
}
pub async fn establish_connection(
&self,
conn: &ManagedConnection,
) -> Result<(), McpTransportError> {
let config = &conn.config;
let server_name = config.name.clone();
let transport = TransportFactory::create(config).await?;
let tools = transport.list_tools().await?;
self.tool_cache.retain(|_, (srv, _)| srv != &server_name);
for tool in tools {
self.tool_cache
.insert(tool.name.clone(), (server_name.clone(), Some(tool)));
}
*conn.transport.write().await = Some(transport);
Ok(())
}
pub fn get(&self, server_name: &str) -> Option<Arc<ManagedConnection>> {
self.connections.get(server_name).map(|r| r.value().clone())
}
pub fn remove(&self, server_name: &str) -> Option<Arc<ManagedConnection>> {
self.connections.remove(server_name).map(|(_, v)| v)
}
pub fn server_for_tool(&self, tool_name: &str) -> Option<String> {
self.tool_cache.get(tool_name).map(|r| r.value().0.clone())
}
pub fn get_tool_definition(&self, tool_name: &str) -> Option<McpToolDefinition> {
self.tool_cache
.get(tool_name)
.and_then(|r| r.value().1.clone())
}
pub fn list_servers(&self) -> Vec<String> {
self.connections.iter().map(|r| r.key().clone()).collect()
}
pub fn list_tools(&self) -> Vec<(String, McpToolDefinition)> {
self.tool_cache
.iter()
.filter_map(|r| r.value().1.clone().map(|def| (r.value().0.clone(), def)))
.collect()
}
pub fn list_tool_definitions(&self) -> Vec<McpToolDefinition> {
self.tool_cache
.iter()
.filter_map(|r| r.value().1.clone())
.collect()
}
pub fn is_connected(&self, server_name: &str) -> bool {
self.connections.contains_key(server_name)
}
pub fn clear_tools_for_server(&self, server_name: &str) {
self.tool_cache.retain(|_, (srv, _)| srv != server_name);
}
pub fn clear(&self) {
self.connections.clear();
self.tool_cache.clear();
}
pub fn iter(&self) -> impl Iterator<Item = (String, Arc<ManagedConnection>)> + '_ {
self.connections
.iter()
.map(|r| (r.key().clone(), r.value().clone()))
}
pub async fn call_tool(&self, name: &str, args: Value) -> Result<Value, McpTransportError> {
let server_name = self
.server_for_tool(name)
.ok_or_else(|| McpTransportError::UnknownTool(name.to_string()))?;
let connection = self
.get(&server_name)
.ok_or_else(|| McpTransportError::ServerNotFound(server_name.clone()))?;
if !connection.circuit_breaker.allow_request() {
return Err(McpTransportError::ServerError(format!(
"Server '{}' circuit breaker is open - server is unhealthy",
server_name
)));
}
let mut failure_rx = connection.subscribe_failures();
let transport = connection
.get_transport()
.await
.ok_or(McpTransportError::ConnectionClosed)?;
let result = tokio::select! {
result = transport.call_tool(name, args) => result,
_ = failure_rx.recv() => {
Err(McpTransportError::ServerRestarting(server_name.clone()))
}
};
match &result {
Ok(_) => connection.circuit_breaker.record_success(),
Err(_) => connection.circuit_breaker.record_failure(),
}
result
}
pub async fn discover_tools_parallel(
&self,
timeout: Duration,
) -> Result<Vec<(String, McpToolDefinition)>, McpTransportError> {
let connections: Vec<_> = self.iter().collect();
let futures: Vec<_> = connections
.into_iter()
.map(|(server_name, conn)| {
let server_name = server_name.clone();
async move {
let result = tokio::time::timeout(timeout, async {
if let Some(transport) = conn.get_transport().await {
transport.list_tools().await
} else {
Err(McpTransportError::ConnectionClosed)
}
})
.await;
match result {
Ok(Ok(tools)) => (server_name, Ok(tools)),
Ok(Err(e)) => (server_name, Err(e)),
Err(_) => (
server_name.clone(),
Err(McpTransportError::Timeout(format!(
"Tool discovery for '{}' timed out",
server_name
))),
),
}
}
})
.collect();
let results = join_all(futures).await;
let mut all_tools = Vec::new();
for (server_name, result) in results {
match result {
Ok(tools) => {
for tool in tools {
self.tool_cache
.insert(tool.name.clone(), (server_name.clone(), Some(tool.clone())));
all_tools.push((server_name.clone(), tool));
}
}
Err(e) => {
eprintln!(
"Warning: Failed to discover tools from '{}': {}",
server_name, e
);
}
}
}
Ok(all_tools)
}
pub async fn refresh_tools_parallel(&self, timeout: Duration) -> Result<(), McpTransportError> {
self.tool_cache.clear();
let _ = self.discover_tools_parallel(timeout).await?;
Ok(())
}
pub async fn health_check(&self) -> Vec<(String, bool)> {
let connections: Vec<_> = self.iter().collect();
let mut results = Vec::new();
for (name, conn) in connections {
let transport_alive = conn.is_alive().await;
let circuit_ok = conn.circuit_breaker.allow_request();
results.push((name, transport_alive && circuit_ok));
}
results
}
pub fn circuit_breaker_stats(
&self,
server_name: &str,
) -> Option<crate::circuit_breaker::CircuitBreakerStats> {
self.get(server_name).map(|c| c.circuit_breaker.stats())
}
pub fn reset_circuit_breaker(&self, server_name: &str) {
if let Some(conn) = self.get(server_name) {
conn.circuit_breaker.reset();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transport::RestartPolicy;
#[test]
fn test_restart_policy_delay() {
let policy = RestartPolicy {
enabled: true,
max_attempts: Some(5),
delay_ms: 1000,
max_delay_ms: 30_000,
backoff_multiplier: 2.0,
};
assert_eq!(policy.delay_for_attempt(0), 1000);
assert_eq!(policy.delay_for_attempt(1), 2000);
assert_eq!(policy.delay_for_attempt(2), 4000);
assert_eq!(policy.delay_for_attempt(5), 30_000); }
#[test]
fn test_hub_connections_creation() {
let conns = HubConnections::new();
assert!(conns.list_servers().is_empty());
assert!(conns.list_tool_definitions().is_empty());
}
}