use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::fmt;
use crate::protocol::McpToolDefinition;
#[async_trait]
pub trait McpTransport: Send + Sync {
async fn list_tools(&self) -> Result<Vec<McpToolDefinition>, McpTransportError>;
async fn call_tool(&self, name: &str, args: Value) -> Result<Value, McpTransportError>;
async fn shutdown(&self) -> Result<(), McpTransportError>;
fn is_alive(&self) -> bool;
fn transport_type(&self) -> TransportTypeId;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum TransportTypeId {
Stdio,
Http,
}
impl fmt::Display for TransportTypeId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
TransportTypeId::Stdio => write!(f, "stdio"),
TransportTypeId::Http => write!(f, "http"),
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum McpTransportError {
#[error("Unknown tool: {0}")]
UnknownTool(String),
#[error("Server not found: {0}")]
ServerNotFound(String),
#[error("Server error: {0}")]
ServerError(String),
#[error("Transport error: {0}")]
TransportError(String),
#[error("IO error: {0}")]
IoError(#[from] std::io::Error),
#[error("JSON error: {0}")]
JsonError(#[from] serde_json::Error),
#[error("Timeout: {0}")]
Timeout(String),
#[error("Protocol error: {0}")]
ProtocolError(String),
#[error("Not supported: {0}")]
NotSupported(String),
#[error("Connection closed")]
ConnectionClosed,
#[error("Server '{0}' is restarting")]
ServerRestarting(String),
}
impl From<String> for McpTransportError {
fn from(s: String) -> Self {
McpTransportError::TransportError(s)
}
}
impl From<&str> for McpTransportError {
fn from(s: &str) -> Self {
McpTransportError::TransportError(s.to_string())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RestartPolicy {
pub enabled: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_attempts: Option<u32>,
pub delay_ms: u64,
pub backoff_multiplier: f64,
pub max_delay_ms: u64,
}
impl Default for RestartPolicy {
fn default() -> Self {
Self {
enabled: false,
max_attempts: None,
delay_ms: 1000,
backoff_multiplier: 2.0,
max_delay_ms: 30000,
}
}
}
impl RestartPolicy {
pub fn none() -> Self {
Self::default()
}
pub fn always() -> Self {
Self {
enabled: true,
max_attempts: None,
delay_ms: 1000,
backoff_multiplier: 2.0,
max_delay_ms: 30000,
}
}
pub fn max_retries(attempts: u32) -> Self {
Self {
enabled: true,
max_attempts: Some(attempts),
delay_ms: 1000,
backoff_multiplier: 2.0,
max_delay_ms: 30000,
}
}
pub fn with_delay_ms(mut self, ms: u64) -> Self {
self.delay_ms = ms;
self
}
pub fn with_backoff(mut self, multiplier: f64) -> Self {
self.backoff_multiplier = multiplier;
self
}
pub fn with_max_delay_ms(mut self, ms: u64) -> Self {
self.max_delay_ms = ms;
self
}
pub fn delay_for_attempt(&self, attempt: u32) -> u64 {
if self.backoff_multiplier > 1.0 {
let exp_delay = (self.delay_ms as f64) * self.backoff_multiplier.powi(attempt as i32);
(exp_delay as u64).min(self.max_delay_ms)
} else {
self.delay_ms
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpServerConnectionConfig {
pub name: String,
pub transport: TransportTypeId,
#[serde(skip_serializing_if = "Option::is_none")]
pub command: Option<String>,
#[serde(default)]
pub args: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub url: Option<String>,
#[serde(default)]
pub config: Value,
#[serde(default = "default_timeout")]
pub timeout_secs: u64,
#[serde(default)]
pub env: std::collections::HashMap<String, String>,
#[serde(default)]
pub restart_policy: RestartPolicy,
}
fn default_timeout() -> u64 {
30
}
impl McpServerConnectionConfig {
pub fn stdio(name: impl Into<String>, command: impl Into<String>, args: Vec<String>) -> Self {
Self {
name: name.into(),
transport: TransportTypeId::Stdio,
command: Some(command.into()),
args,
url: None,
config: Value::Object(serde_json::Map::new()),
timeout_secs: default_timeout(),
env: std::collections::HashMap::new(),
restart_policy: RestartPolicy::none(),
}
}
pub fn http(name: impl Into<String>, url: impl Into<String>) -> Self {
Self {
name: name.into(),
transport: TransportTypeId::Http,
command: None,
args: Vec::new(),
url: Some(url.into()),
config: Value::Object(serde_json::Map::new()),
timeout_secs: default_timeout(),
env: std::collections::HashMap::new(),
restart_policy: RestartPolicy::none(),
}
}
pub fn with_config(mut self, config: Value) -> Self {
self.config = config;
self
}
pub fn with_timeout(mut self, timeout_secs: u64) -> Self {
self.timeout_secs = timeout_secs;
self
}
pub fn with_env(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.env.insert(key.into(), value.into());
self
}
pub fn with_restart(mut self, policy: RestartPolicy) -> Self {
self.restart_policy = policy;
self
}
pub fn restart_on_failure(self) -> Self {
self.with_restart(RestartPolicy::always())
}
pub fn restart_max_attempts(self, attempts: u32) -> Self {
self.with_restart(RestartPolicy::max_retries(attempts))
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InitializeParams {
#[serde(rename = "protocolVersion")]
pub protocol_version: String,
pub capabilities: InitializeCapabilities,
#[serde(rename = "clientInfo")]
pub client_info: ClientInfo,
#[serde(skip_serializing_if = "Option::is_none")]
pub config: Option<Value>,
}
impl InitializeParams {
pub fn new(config: Option<Value>) -> Self {
Self {
protocol_version: crate::MCP_PROTOCOL_VERSION.to_string(),
capabilities: InitializeCapabilities::default(),
client_info: ClientInfo::new("mcp-rust", env!("CARGO_PKG_VERSION")),
config,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InitializeResult {
#[serde(rename = "protocolVersion")]
pub protocol_version: String,
pub capabilities: ServerCapabilities,
#[serde(rename = "serverInfo")]
pub server_info: ServerInfo,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct InitializeCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub experimental: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub roots: Option<RootsCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub sampling: Option<SamplingCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub elicitation: Option<ElicitationCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tasks: Option<TasksCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<ToolCapabilities>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct RootsCapabilities {
#[serde(rename = "listChanged", skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct SamplingCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub context: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ElicitationCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub form: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub url: Option<Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct TasksCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub list: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cancel: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub requests: Option<Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ServerCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub experimental: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logging: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub completions: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompts: Option<PromptsCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub resources: Option<ResourcesCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<ServerToolCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tasks: Option<TasksCapabilities>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct PromptsCapabilities {
#[serde(rename = "listChanged", skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ResourcesCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub subscribe: Option<bool>,
#[serde(rename = "listChanged", skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ServerToolCapabilities {
#[serde(rename = "listChanged", skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ToolCapabilities {
#[serde(rename = "listChanged", skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClientInfo {
pub name: String,
pub version: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub icons: Option<Vec<crate::protocol::Icon>>,
#[serde(rename = "websiteUrl", skip_serializing_if = "Option::is_none")]
pub website_url: Option<String>,
}
impl ClientInfo {
pub fn new(name: impl Into<String>, version: impl Into<String>) -> Self {
Self {
name: name.into(),
version: version.into(),
title: None,
description: None,
icons: None,
website_url: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ServerInfo {
pub name: String,
pub version: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub icons: Option<Vec<crate::protocol::Icon>>,
#[serde(rename = "websiteUrl", skip_serializing_if = "Option::is_none")]
pub website_url: Option<String>,
}
impl ServerInfo {
pub fn new(name: impl Into<String>, version: impl Into<String>) -> Self {
Self {
name: name.into(),
version: version.into(),
title: None,
description: None,
icons: None,
website_url: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_transport_type_display() {
assert_eq!(TransportTypeId::Stdio.to_string(), "stdio");
assert_eq!(TransportTypeId::Http.to_string(), "http");
}
#[test]
fn test_connection_config_stdio() {
let config =
McpServerConnectionConfig::stdio("test", "node", vec!["server.js".to_string()])
.with_timeout(60);
assert_eq!(config.name, "test");
assert_eq!(config.transport, TransportTypeId::Stdio);
assert_eq!(config.command, Some("node".to_string()));
assert_eq!(config.timeout_secs, 60);
}
#[test]
fn test_connection_config_http() {
let config = McpServerConnectionConfig::http("api", "http://localhost:8080/mcp");
assert_eq!(config.name, "api");
assert_eq!(config.transport, TransportTypeId::Http);
assert_eq!(config.url, Some("http://localhost:8080/mcp".to_string()));
}
#[test]
fn test_initialize_params() {
let params = InitializeParams::new(None);
assert_eq!(params.protocol_version, "2025-11-25");
assert_eq!(params.client_info.name, "mcp-rust");
}
}