use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::RwLock;
use super::config::McpServerConfig;
use super::error::{McpError, McpResult};
use super::protocol::{
ClientInfo, InitializeParams, InitializeResult, JsonRpcRequest, JsonRpcResponse,
McpCapabilities, SUPPORTED_PROTOCOL_VERSION, methods,
};
use super::tools::{Tool, ToolCall, ToolList, ToolResult};
use super::transport::Transport;
use crate::utils::net::http::{
get_client_with_timeout, get_ssrf_safe_client_with_timeout_fallible,
};
use serde_json::Value;
use sha2::{Digest, Sha256};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ServerState {
Disconnected,
Connecting,
Connected,
Failed,
}
#[derive(Debug)]
pub struct McpServer {
config: McpServerConfig,
state: RwLock<ServerState>,
http_client: Arc<reqwest::Client>,
custom_headers: reqwest::header::HeaderMap,
tools_cache: RwLock<Option<Vec<Tool>>>,
tools_baseline_hash: RwLock<Option<String>>,
capabilities: RwLock<Option<McpCapabilities>>,
request_id: std::sync::atomic::AtomicU64,
}
impl McpServer {
pub fn new(config: McpServerConfig) -> McpResult<Self> {
config
.validate()
.map_err(|e| McpError::ConfigurationError { message: e })?;
let timeout_secs = config.timeout_ms / 1000;
let timeout = Duration::from_secs(timeout_secs.max(1));
let http_client = match config.transport {
Transport::Http | Transport::Sse | Transport::WebSocket => {
get_ssrf_safe_client_with_timeout_fallible(timeout).map_err(|error| {
McpError::ConfigurationError {
message: format!("Failed to create SSRF-safe HTTP client: {error}"),
}
})?
}
Transport::Stdio => get_client_with_timeout(timeout),
};
let mut headers = reqwest::header::HeaderMap::new();
for (key, value) in &config.static_headers {
if let (Ok(name), Ok(val)) = (
reqwest::header::HeaderName::from_bytes(key.as_bytes()),
reqwest::header::HeaderValue::from_str(value),
) {
headers.insert(name, val);
}
}
if let Some(auth) = &config.auth
&& let Some(header_value) = auth.get_header_value()
{
let header_name = auth.get_header_name();
if let (Ok(name), Ok(val)) = (
reqwest::header::HeaderName::from_bytes(header_name.as_bytes()),
reqwest::header::HeaderValue::from_str(&header_value),
) {
headers.insert(name, val);
}
}
Ok(Self {
config,
state: RwLock::new(ServerState::Disconnected),
http_client,
custom_headers: headers,
tools_cache: RwLock::new(None),
tools_baseline_hash: RwLock::new(None),
capabilities: RwLock::new(None),
request_id: std::sync::atomic::AtomicU64::new(1),
})
}
pub fn name(&self) -> &str {
&self.config.name
}
pub fn url(&self) -> &str {
&self.config.url
}
pub fn transport(&self) -> Transport {
self.config.transport
}
pub async fn is_connected(&self) -> bool {
*self.state.read().await == ServerState::Connected
}
pub async fn state(&self) -> ServerState {
*self.state.read().await
}
pub async fn connect(&self) -> McpResult<()> {
{
let mut state = self.state.write().await;
if *state == ServerState::Connected {
return Ok(());
}
*state = ServerState::Connecting;
}
*self.capabilities.write().await = None;
match self.initialize().await {
Ok(caps) => {
*self.capabilities.write().await = Some(caps);
*self.state.write().await = ServerState::Connected;
Ok(())
}
Err(e) => {
*self.state.write().await = ServerState::Failed;
Err(e)
}
}
}
async fn initialize(&self) -> McpResult<McpCapabilities> {
let params = InitializeParams {
protocol_version: SUPPORTED_PROTOCOL_VERSION.to_string(),
capabilities: McpCapabilities::default(),
client_info: ClientInfo::default(),
};
let response = self
.send_request(methods::INITIALIZE, Some(serde_json::to_value(params)?))
.await?;
let result = parse_initialize_response(response)?;
self.send_notification(methods::INITIALIZED, None).await?;
Ok(result.capabilities)
}
pub async fn disconnect(&self) {
*self.state.write().await = ServerState::Disconnected;
*self.tools_cache.write().await = None;
*self.capabilities.write().await = None;
}
pub async fn list_tools(&self) -> McpResult<ToolList> {
if let Some(tools) = self.tools_cache.read().await.as_ref() {
return Ok(ToolList {
tools: tools.clone(),
next_cursor: None,
});
}
let response = self.send_request(methods::LIST_TOOLS, None).await?;
if let Some(result) = response.result {
let list: ToolList = serde_json::from_value(result)?;
self.cache_tools_with_baseline(&list).await?;
Ok(list)
} else if let Some(error) = response.error {
Err(McpError::ProtocolError {
message: format!("List tools failed: {}", error.message),
})
} else {
Ok(ToolList::empty())
}
}
pub async fn get_tool(&self, name: &str) -> McpResult<Option<Tool>> {
let list = self.list_tools().await?;
Ok(list.tools.into_iter().find(|t| t.name == name))
}
pub async fn call_tool(&self, call: ToolCall) -> McpResult<ToolResult> {
let params = serde_json::json!({
"name": call.name,
"arguments": call.arguments
});
let response = self.send_request(methods::CALL_TOOL, Some(params)).await?;
if let Some(result) = response.result {
let tool_result: ToolResult = serde_json::from_value(result)?;
Ok(tool_result)
} else if let Some(error) = response.error {
Err(McpError::ToolExecutionError {
server_name: self.config.name.clone(),
tool_name: call.name,
code: error.code,
message: error.message,
})
} else {
Err(McpError::ProtocolError {
message: "Empty response from tool call".to_string(),
})
}
}
async fn send_request(
&self,
method: &str,
params: Option<serde_json::Value>,
) -> McpResult<JsonRpcResponse> {
match self.config.transport {
Transport::Http => self.send_http_request(method, params).await,
Transport::Sse => self.send_sse_request(method, params).await,
Transport::Stdio => Err(McpError::TransportError {
transport: "stdio".to_string(),
message: "Stdio transport not yet implemented".to_string(),
}),
Transport::WebSocket => Err(McpError::TransportError {
transport: "websocket".to_string(),
message: "WebSocket transport not yet implemented".to_string(),
}),
}
}
async fn send_notification(
&self,
method: &str,
params: Option<serde_json::Value>,
) -> McpResult<()> {
match self.config.transport {
Transport::Http | Transport::Sse => self.send_http_notification(method, params).await,
Transport::Stdio => Err(McpError::TransportError {
transport: "stdio".to_string(),
message: "Stdio transport not yet implemented".to_string(),
}),
Transport::WebSocket => Err(McpError::TransportError {
transport: "websocket".to_string(),
message: "WebSocket transport not yet implemented".to_string(),
}),
}
}
async fn send_http_request(
&self,
method: &str,
params: Option<serde_json::Value>,
) -> McpResult<JsonRpcResponse> {
let id = self
.request_id
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let request =
JsonRpcRequest::new(method, params).with_id(serde_json::Value::Number(id.into()));
let response = self.send_http_message(&request).await?;
response.json().await.map_err(|e| McpError::ProtocolError {
message: format!("Failed to parse response: {e}"),
})
}
async fn send_http_notification(
&self,
method: &str,
params: Option<serde_json::Value>,
) -> McpResult<()> {
let notification = JsonRpcRequest::notification(method, params);
self.send_http_message(¬ification)
.await?
.bytes()
.await
.map_err(|error| self.map_http_error(error))?;
Ok(())
}
fn map_http_error(&self, error: reqwest::Error) -> McpError {
if error.is_timeout() {
McpError::Timeout {
server_name: self.config.name.clone(),
timeout_ms: self.config.timeout_ms,
}
} else if error.is_connect() {
McpError::ConnectionError {
server_name: self.config.name.clone(),
message: error.to_string(),
}
} else {
McpError::TransportError {
transport: "http".to_string(),
message: error.to_string(),
}
}
}
async fn send_http_message(&self, request: &JsonRpcRequest) -> McpResult<reqwest::Response> {
let response = self
.http_client
.post(&self.config.url)
.headers(self.custom_headers.clone())
.json(&request)
.send()
.await
.map_err(|error| self.map_http_error(error))?;
let status = response.status();
if status == reqwest::StatusCode::UNAUTHORIZED {
return Err(McpError::AuthenticationError {
server_name: self.config.name.clone(),
message: "Unauthorized".to_string(),
});
}
if status == reqwest::StatusCode::FORBIDDEN {
return Err(McpError::AuthorizationError {
server_name: self.config.name.clone(),
tool_name: None,
message: "Forbidden".to_string(),
});
}
if status == reqwest::StatusCode::TOO_MANY_REQUESTS {
let retry_after = response
.headers()
.get("retry-after")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok())
.map(|s| s * 1000);
return Err(McpError::RateLimitExceeded {
server_name: self.config.name.clone(),
retry_after_ms: retry_after,
});
}
if !status.is_success() {
return Err(McpError::TransportError {
transport: "http".to_string(),
message: format!("MCP server returned HTTP status {status}"),
});
}
Ok(response)
}
async fn send_sse_request(
&self,
method: &str,
params: Option<serde_json::Value>,
) -> McpResult<JsonRpcResponse> {
self.send_http_request(method, params).await
}
pub async fn invalidate_cache(&self) {
*self.tools_cache.write().await = None;
}
async fn cache_tools_with_baseline(&self, list: &ToolList) -> McpResult<()> {
let hash = stable_tools_hash(&list.tools)?;
let mut baseline = self.tools_baseline_hash.write().await;
if let Some(expected_hash) = baseline.as_ref() {
if expected_hash != &hash {
return Err(McpError::ToolDefinitionChanged {
server_name: self.config.name.clone(),
expected_hash: expected_hash.clone(),
actual_hash: hash,
});
}
} else {
*baseline = Some(hash);
}
*self.tools_cache.write().await = Some(list.tools.clone());
Ok(())
}
pub async fn capabilities(&self) -> Option<McpCapabilities> {
self.capabilities.read().await.clone()
}
}
fn parse_initialize_response(response: JsonRpcResponse) -> McpResult<InitializeResult> {
let result = match (response.result, response.error) {
(Some(result), None) => result,
(None, Some(error)) => {
return Err(McpError::ProtocolError {
message: format!("Initialize failed: {}", error.message),
});
}
(Some(_), Some(_)) => {
return Err(McpError::ProtocolError {
message: "Initialize response contained both result and error".to_string(),
});
}
(None, None) => {
return Err(McpError::ProtocolError {
message: "Initialize response contained neither result nor error".to_string(),
});
}
};
let result: InitializeResult =
serde_json::from_value(result).map_err(|error| McpError::ProtocolError {
message: format!("Invalid initialize result: {error}"),
})?;
if result.protocol_version != SUPPORTED_PROTOCOL_VERSION {
return Err(McpError::ProtocolError {
message: format!(
"Unsupported MCP protocol version '{}'; expected '{}'",
result.protocol_version, SUPPORTED_PROTOCOL_VERSION
),
});
}
Ok(result)
}
pub type McpServerHandle = Arc<McpServer>;
fn stable_tools_hash(tools: &[Tool]) -> McpResult<String> {
let mut values = tools
.iter()
.map(|tool| {
serde_json::to_value(tool).map_err(|e| McpError::SerializationError {
message: e.to_string(),
})
})
.collect::<McpResult<Vec<_>>>()?;
for value in &mut values {
canonicalize_json(value);
}
values.sort_by(|left, right| {
tool_sort_key(left)
.cmp(&tool_sort_key(right))
.then_with(|| left.to_string().cmp(&right.to_string()))
});
let canonical = serde_json::to_string(&values).map_err(|e| McpError::SerializationError {
message: e.to_string(),
})?;
Ok(hex::encode(Sha256::digest(canonical.as_bytes())))
}
fn tool_sort_key(value: &Value) -> String {
value
.get("name")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string()
}
fn canonicalize_json(value: &mut Value) {
match value {
Value::Object(map) => {
let mut entries = std::mem::take(map).into_iter().collect::<Vec<_>>();
entries.sort_by(|left, right| left.0.cmp(&right.0));
for (key, mut child) in entries {
canonicalize_json(&mut child);
map.insert(key, child);
}
}
Value::Array(items) => {
for item in items {
canonicalize_json(item);
}
}
_ => {}
}
}
#[derive(Debug, Default)]
pub struct McpServerRegistry {
servers: RwLock<HashMap<String, McpServerHandle>>,
}
impl McpServerRegistry {
pub fn new() -> Self {
Self::default()
}
pub async fn register(&self, config: McpServerConfig) -> McpResult<()> {
let name = config.name.clone();
if self.servers.read().await.contains_key(&name) {
return Err(McpError::ServerAlreadyExists { server_name: name });
}
let server = Arc::new(McpServer::new(config)?);
self.servers.write().await.insert(name, server);
Ok(())
}
pub async fn get(&self, name: &str) -> Option<McpServerHandle> {
self.servers.read().await.get(name).cloned()
}
pub async fn remove(&self, name: &str) -> Option<McpServerHandle> {
self.servers.write().await.remove(name)
}
pub async fn list_names(&self) -> Vec<String> {
self.servers.read().await.keys().cloned().collect()
}
pub async fn count(&self) -> usize {
self.servers.read().await.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::mcp::config::AuthConfig;
use crate::core::mcp::tools::ToolInputSchema;
#[test]
fn test_server_state_variants() {
assert_eq!(ServerState::Disconnected, ServerState::Disconnected);
assert_ne!(ServerState::Connected, ServerState::Disconnected);
}
#[tokio::test]
async fn test_server_creation() {
let config = McpServerConfig::new("test", "https://1.1.1.1/mcp");
let server = McpServer::new(config).unwrap();
assert_eq!(server.name(), "test");
assert_eq!(server.url(), "https://1.1.1.1/mcp");
assert!(!server.is_connected().await);
}
#[tokio::test]
async fn test_server_with_auth() {
let config = McpServerConfig::new("test", "https://1.1.1.1/mcp")
.with_auth(AuthConfig::bearer("token123"));
let server = McpServer::new(config).unwrap();
assert_eq!(server.name(), "test");
}
#[tokio::test]
async fn test_server_registry() {
let registry = McpServerRegistry::new();
registry
.register(McpServerConfig::new("server1", "https://1.1.1.1/mcp1"))
.await
.unwrap();
assert!(registry.get("server1").await.is_some());
assert!(registry.get("nonexistent").await.is_none());
assert_eq!(registry.count().await, 1);
let names = registry.list_names().await;
assert!(names.contains(&"server1".to_string()));
}
#[tokio::test]
async fn test_registry_duplicate_server() {
let registry = McpServerRegistry::new();
registry
.register(McpServerConfig::new("server1", "https://1.1.1.1/mcp1"))
.await
.unwrap();
let result = registry
.register(McpServerConfig::new("server1", "https://1.1.1.1/mcp2"))
.await;
assert!(matches!(result, Err(McpError::ServerAlreadyExists { .. })));
}
#[tokio::test]
async fn test_registry_remove() {
let registry = McpServerRegistry::new();
registry
.register(McpServerConfig::new("server1", "https://1.1.1.1/mcp1"))
.await
.unwrap();
let removed = registry.remove("server1").await;
assert!(removed.is_some());
assert_eq!(registry.count().await, 0);
}
#[tokio::test]
async fn test_server_initial_state() {
let config = McpServerConfig::new("test", "https://1.1.1.1/mcp");
let server = McpServer::new(config).unwrap();
assert_eq!(server.state().await, ServerState::Disconnected);
}
#[test]
fn test_tools_hash_is_stable_across_ordering() {
let schema_a = ToolInputSchema::object()
.with_property("city", super::super::tools::PropertySchema::string(), true)
.with_property(
"units",
super::super::tools::PropertySchema::string(),
false,
);
let schema_b = ToolInputSchema::object()
.with_property(
"units",
super::super::tools::PropertySchema::string(),
false,
)
.with_property("city", super::super::tools::PropertySchema::string(), true);
let tools_a = vec![
Tool::new("search").with_description("Search docs"),
Tool::new("weather")
.with_description("Get weather")
.with_schema(schema_a),
];
let tools_b = vec![
Tool::new("weather")
.with_description("Get weather")
.with_schema(schema_b),
Tool::new("search").with_description("Search docs"),
];
assert_eq!(
stable_tools_hash(&tools_a).unwrap(),
stable_tools_hash(&tools_b).unwrap()
);
}
#[tokio::test]
async fn test_tools_baseline_rejects_definition_change() {
let config = McpServerConfig::new("test", "https://1.1.1.1/mcp");
let server = McpServer::new(config).unwrap();
let initial = ToolList {
tools: vec![Tool::new("search").with_description("Search docs")],
next_cursor: None,
};
let changed = ToolList {
tools: vec![Tool::new("search").with_description("Read private files")],
next_cursor: None,
};
server.cache_tools_with_baseline(&initial).await.unwrap();
server.invalidate_cache().await;
let result = server.cache_tools_with_baseline(&changed).await;
assert!(matches!(
result,
Err(McpError::ToolDefinitionChanged {
server_name,
expected_hash: _,
actual_hash: _
}) if server_name == "test"
));
}
#[test]
fn test_invalid_config() {
let config = McpServerConfig {
name: "".to_string(), url: "https://example.com".to_string(),
..Default::default()
};
let result = McpServer::new(config);
assert!(result.is_err());
}
}
#[cfg(test)]
#[path = "server_lifecycle_tests.rs"]
mod lifecycle_tests;