use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::{
atomic::{AtomicU64, Ordering},
Arc,
};
use tracing::{info, warn};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct HttpConfig {
pub url: String,
#[serde(default)]
#[serde(skip_serializing_if = "Option::is_none")]
pub headers: Option<HashMap<String, String>>,
}
impl HttpConfig {
pub fn new(url: String) -> Self {
Self { url, headers: None }
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub enum McpTransport {
Stdio {
command: String,
args: Vec<String>,
env: HashMap<String, String>,
},
Http {
#[serde(flatten)]
config: HttpConfig,
},
HttpSse {
#[serde(flatten)]
config: HttpConfig,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum McpServerHealth {
Configured,
Connecting,
Connected,
Error,
Disabled,
}
impl McpServerHealth {
pub fn is_usable(&self) -> bool {
matches!(self, McpServerHealth::Connected)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpToolInfo {
pub name: String,
pub description: String,
pub input_schema: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpServerConfig {
pub id: String,
pub label: String,
#[serde(default)]
pub description: Option<String>,
pub transport: McpTransport,
pub enabled_by_default: bool,
pub working_dir: Option<PathBuf>,
#[serde(default)]
pub inherited_env_vars: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct McpServerState {
pub config: McpServerConfig,
pub health: McpServerHealth,
pub discovered_tools: Vec<McpToolInfo>,
pub last_error: Option<String>,
}
impl McpServerState {
pub fn new(config: McpServerConfig) -> Self {
Self {
config,
health: McpServerHealth::Configured,
discovered_tools: Vec::new(),
last_error: None,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct McpServerRegistry {
servers: Arc<RwLock<HashMap<String, McpServerState>>>,
version: Arc<AtomicU64>,
}
impl McpServerRegistry {
pub fn new() -> Self {
Self {
servers: Arc::new(RwLock::new(HashMap::new())),
version: Arc::new(AtomicU64::new(0)),
}
}
fn bump_version(&self) {
self.version.fetch_add(1, Ordering::SeqCst);
}
pub fn version(&self) -> u64 {
self.version.load(Ordering::SeqCst)
}
pub fn register_server(&self, config: McpServerConfig) {
let mut servers = self.servers.write();
let id = config.id.clone();
let state = McpServerState::new(config);
servers.insert(id.clone(), state);
drop(servers);
self.bump_version();
info!("Registered MCP server: {}", id);
}
pub fn unregister_server(&self, server_id: &str) -> Option<McpServerState> {
let mut servers = self.servers.write();
let removed = servers.remove(server_id);
drop(servers);
if removed.is_some() {
self.bump_version();
info!("Unregistered MCP server: {}", server_id);
}
removed
}
pub fn list_servers(&self) -> Vec<McpServerState> {
let servers = self.servers.read();
servers.values().cloned().collect()
}
pub fn get_server(&self, server_id: &str) -> Option<McpServerState> {
let servers = self.servers.read();
servers.get(server_id).cloned()
}
pub fn update_health(&self, server_id: &str, health: McpServerHealth) {
let mut servers = self.servers.write();
if let Some(state) = servers.get_mut(server_id) {
if matches!(health, McpServerHealth::Connecting)
&& matches!(
state.health,
McpServerHealth::Connected | McpServerHealth::Error
)
{
return;
}
state.health = health;
if health.is_usable() {
state.last_error = None;
}
drop(servers);
self.bump_version();
}
}
pub fn set_error(&self, server_id: &str, error: String) {
let mut servers = self.servers.write();
if let Some(state) = servers.get_mut(server_id) {
let log_message = error.clone();
state.health = McpServerHealth::Error;
state.last_error = Some(error);
drop(servers);
self.bump_version();
warn!(
"MCP server {} entered error state: {}",
server_id, log_message
);
}
}
pub fn update_discovered_tools(&self, server_id: &str, tools: Vec<McpToolInfo>) {
let mut servers = self.servers.write();
if let Some(state) = servers.get_mut(server_id) {
state.discovered_tools = tools;
let tool_count = state.discovered_tools.len();
drop(servers);
self.bump_version();
info!(
"Updated MCP server {} with {} discovered tools",
server_id, tool_count
);
}
}
pub fn is_server_usable(&self, server_id: &str) -> bool {
let servers = self.servers.read();
servers
.get(server_id)
.map(|s| s.health.is_usable())
.unwrap_or(false)
}
pub fn get_all_usable_tools(&self) -> Vec<(String, McpToolInfo)> {
let servers = self.servers.read();
let mut all_tools = Vec::new();
for (server_id, state) in servers.iter() {
if state.health.is_usable() {
for tool in &state.discovered_tools {
all_tools.push((server_id.clone(), tool.clone()));
}
}
}
all_tools
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn http_config_serializes_without_headers_when_none() {
let config = HttpConfig::new("https://example.com/mcp".to_string());
let json = serde_json::to_string(&config).unwrap();
assert!(
!json.contains("headers"),
"serialized output should not contain 'headers' when None: {}",
json
);
assert!(
json.contains("\"url\""),
"serialized output should contain 'url': {}",
json
);
}
#[test]
fn http_config_round_trip_with_headers() {
let mut headers = HashMap::new();
headers.insert("Authorization".to_string(), "Bearer token123".to_string());
headers.insert("X-Custom".to_string(), "value".to_string());
let config = HttpConfig {
url: "https://example.com/mcp".to_string(),
headers: Some(headers),
};
let json = serde_json::to_string(&config).unwrap();
let deserialized: HttpConfig = serde_json::from_str(&json).unwrap();
assert_eq!(config, deserialized);
assert_eq!(
deserialized.headers.as_ref().unwrap().get("Authorization"),
Some(&"Bearer token123".to_string())
);
}
#[test]
fn http_config_deserializes_without_headers() {
let json = r#"{"url":"https://example.com/mcp"}"#;
let config: HttpConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.url, "https://example.com/mcp");
assert_eq!(config.headers, None);
}
#[test]
fn mcp_transport_http_flattens_http_config() {
let json = r#"{"Http":{"url":"https://example.com/mcp","headers":{"X-Key":"val"}}}"#;
let transport: McpTransport = serde_json::from_str(json).unwrap();
match transport {
McpTransport::Http { config } => {
assert_eq!(config.url, "https://example.com/mcp");
assert_eq!(
config.headers.as_ref().unwrap().get("X-Key"),
Some(&"val".to_string())
);
}
other => panic!("expected Http variant, got {:?}", other),
}
}
#[test]
fn mcp_transport_http_deserializes_without_headers() {
let json = r#"{"Http":{"url":"https://example.com/mcp"}}"#;
let transport: McpTransport = serde_json::from_str(json).unwrap();
match transport {
McpTransport::Http { config } => {
assert_eq!(config.url, "https://example.com/mcp");
assert_eq!(config.headers, None);
}
other => panic!("expected Http variant, got {:?}", other),
}
}
}