use std::collections::{HashMap, VecDeque};
use std::path::Path;
use std::process::Stdio;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use async_trait::async_trait;
use futures::StreamExt;
use reqwest::header::CONTENT_TYPE;
use serde_json::Value;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::Mutex;
use tracing::{error, warn};
use crate::mcp::protocol::messages as protocol_messages;
use crate::mcp::protocol::JsonRpcRequest as ProtocolJsonRpcRequest;
use crate::mcp::protocol::{tool_content_to_value, tool_error_to_string};
use crate::mcp::protocol::{JsonRpcError, JsonRpcResponse};
use crate::mcp::server::{HttpConfig, McpServerConfig, McpToolInfo, McpTransport};
const MCP_ACCEPT_HEADER: &str = "application/json, text/event-stream";
fn apply_headers(
builder: reqwest::RequestBuilder,
custom_headers: &Option<HashMap<String, String>>,
) -> reqwest::RequestBuilder {
let mut builder = builder.header("Accept", MCP_ACCEPT_HEADER);
if let Some(headers) = custom_headers {
for (key, value) in headers {
builder = builder.header(key.as_str(), value.as_str());
}
}
builder
}
async fn decode_http_rpc_response(response: reqwest::Response) -> Result<JsonRpcResponse, String> {
let content_type = response
.headers()
.get(CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_ascii_lowercase();
if content_type
.split(';')
.any(|part| part.trim() == "text/event-stream")
{
let body = response
.text()
.await
.map_err(|e| format!("Failed to read SSE response body: {}", e))?;
return parse_sse_rpc_response(&body);
}
response
.json()
.await
.map_err(|e| format!("Failed to parse JSON response: {}", e))
}
fn parse_sse_rpc_response(body: &str) -> Result<JsonRpcResponse, String> {
let normalized = body.replace("\r\n", "\n").replace('\r', "\n");
for event_block in normalized.split("\n\n") {
let mut event_type: Option<&str> = None;
let mut data_lines = Vec::new();
for line in event_block.lines() {
if line.starts_with(':') || line.is_empty() {
continue;
}
if let Some(rest) = line.strip_prefix("event:") {
event_type = Some(rest.trim_start());
continue;
}
if let Some(rest) = line.strip_prefix("data:") {
data_lines.push(rest.trim_start());
}
}
if data_lines.is_empty() || matches!(event_type, Some("ping" | "keepalive")) {
continue;
}
let payload = data_lines.join("\n");
if let Ok(response) = serde_json::from_str(&payload) {
return Ok(response);
}
}
Err("Failed to parse SSE response: no JSON-RPC data event found".to_string())
}
type WaiterMap =
Arc<Mutex<HashMap<u64, tokio::sync::oneshot::Sender<Result<JsonRpcResponse, String>>>>>;
#[async_trait]
pub trait McpTransportClient: Send + Sync {
async fn initialize(&self) -> Result<protocol_messages::InitializeResponse, String>;
async fn list_tools(&self) -> Result<Vec<McpToolInfo>, String>;
async fn call_tool(&self, tool_name: &str, arguments: Value) -> Result<Value, String>;
fn is_connected(&self) -> bool;
async fn close(&self);
}
pub struct StdioMcpClient {
server_id: String,
_config: McpServerConfig,
stdin: Arc<Mutex<ChildStdin>>,
process: Arc<Mutex<Child>>,
connected: AtomicBool,
request_counter: Arc<Mutex<u64>>,
waiters: WaiterMap,
stdout_task: Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
stderr_task: Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
stderr_excerpt: Arc<Mutex<VecDeque<String>>>,
bootstrap_mode: Arc<AtomicBool>,
}
impl StdioMcpClient {
pub fn new(server_id: String, config: McpServerConfig) -> Result<Self, String> {
let (command, args, env) = match &config.transport {
McpTransport::Stdio { command, args, env } => (command, args, env),
other => {
return Err(format!(
"StdioMcpClient requires Stdio transport, got {:?}",
other
))
}
};
let mut cmd = Command::new(command);
let mut base_env = sanitized_stdio_env();
for var_name in &config.inherited_env_vars {
if let Ok(value) = std::env::var(var_name) {
base_env.insert(var_name.clone(), value);
}
}
cmd.args(args)
.envs(base_env)
.envs(env)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped());
if let Some(working_dir) = &config.working_dir {
cmd.current_dir(working_dir);
}
let mut process = cmd.spawn().map_err(|e| {
format!(
"Failed to spawn MCP server process '{}' for server '{}': {}",
redact_command(command),
server_id,
e
)
})?;
let stdin = process.stdin.take().ok_or_else(|| {
format!(
"Failed to acquire stdin pipe for MCP server '{}'",
server_id
)
})?;
let stdout = process.stdout.take().ok_or_else(|| {
format!(
"Failed to acquire stdout pipe for MCP server '{}'",
server_id
)
})?;
let stderr = process.stderr.take().ok_or_else(|| {
format!(
"Failed to acquire stderr pipe for MCP server '{}'",
server_id
)
})?;
let waiters = Arc::new(Mutex::new(HashMap::<
u64,
tokio::sync::oneshot::Sender<Result<JsonRpcResponse, String>>,
>::new()));
let stderr_excerpt = Arc::new(Mutex::new(VecDeque::new()));
let bootstrap_mode = Arc::new(AtomicBool::new(true));
let stdout_task = tokio::spawn(start_stdio_stdout_reader(
server_id.clone(),
stdout,
Arc::clone(&waiters),
Arc::clone(&stderr_excerpt),
Arc::clone(&bootstrap_mode),
));
let stderr_task = tokio::spawn(start_stdio_stderr_reader(
server_id.clone(),
stderr,
Arc::clone(&stderr_excerpt),
));
Ok(Self {
server_id,
_config: config,
stdin: Arc::new(Mutex::new(stdin)),
process: Arc::new(Mutex::new(process)),
connected: AtomicBool::new(false),
request_counter: Arc::new(Mutex::new(0)),
waiters,
stdout_task: Arc::new(Mutex::new(Some(stdout_task))),
stderr_task: Arc::new(Mutex::new(Some(stderr_task))),
stderr_excerpt,
bootstrap_mode,
})
}
async fn send_request<T: serde::de::DeserializeOwned>(
&self,
method: &str,
params: Value,
) -> Result<T, String> {
let id = {
let mut counter = self.request_counter.lock().await;
*counter += 1;
*counter
};
let request = ProtocolJsonRpcRequest::new(method, params, id);
let request_json = serde_json::to_string(&request)
.map_err(|e| format!("Failed to serialize request: {}", e))?;
let (tx, rx) = tokio::sync::oneshot::channel::<Result<JsonRpcResponse, String>>();
{
let mut waiters = self.waiters.lock().await;
waiters.insert(id, tx);
}
let mut stdin = self.stdin.lock().await;
if let Err(e) = stdin
.write_all(format!("{}\n", request_json).as_bytes())
.await
{
self.cleanup_waiter(id).await;
return Err(format!("Failed to write to stdin: {}", e));
}
if let Err(e) = stdin.flush().await {
self.cleanup_waiter(id).await;
return Err(format!("Failed to flush stdin: {}", e));
}
drop(stdin);
let response: JsonRpcResponse =
match tokio::time::timeout(tokio::time::Duration::from_secs(30), rx).await {
Ok(Ok(Ok(response))) => response,
Ok(Ok(Err(error_msg))) => return Err(error_msg),
Ok(Err(_)) => {
return Err(self.reader_terminated_error(
"stdout reader dropped before delivering response",
));
}
Err(_) => {
self.cleanup_waiter(id).await;
return Err(format!(
"Timeout waiting for stdio response for request {} on server '{}'",
id, self.server_id
));
}
};
if response.id != Some(id) {
if response.id.is_none() && response.error.is_none() && response.result.is_some() {
tracing::debug!(
"Accepting MCP response with missing id during bootstrap for server '{}'",
self.server_id
);
} else {
return Err(format!(
"Response ID mismatch: expected {}, got {:?}",
id, response.id
));
}
}
if let Some(error) = response.error {
return Err(format_rpc_error(error));
}
response
.result
.ok_or_else(|| "Response missing result".to_string())
.and_then(|v| {
serde_json::from_value(v)
.map_err(|e| format!("Failed to deserialize result: {}", e))
})
}
async fn cleanup_waiter(&self, id: u64) {
self.waiters.lock().await.remove(&id);
}
fn reader_terminated_error(&self, prefix: &str) -> String {
format!(
"{} for server '{}'{}",
prefix,
self.server_id,
stderr_suffix_blocking(&self.stderr_excerpt)
)
}
}
async fn start_stdio_stdout_reader(
server_id: String,
stdout: ChildStdout,
waiters: WaiterMap,
stderr_excerpt: Arc<Mutex<VecDeque<String>>>,
bootstrap_mode: Arc<AtomicBool>,
) {
let mut stdout = BufReader::new(stdout);
let mut line = String::new();
loop {
line.clear();
match stdout.read_line(&mut line).await {
Ok(0) => {
fail_pending_waiters(
&waiters,
format!(
"MCP stdio stdout closed for server '{}'{}",
server_id,
stderr_suffix_async(&stderr_excerpt).await
),
)
.await;
break;
}
Ok(_) => {
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let response: JsonRpcResponse = match serde_json::from_str(trimmed) {
Ok(response) => response,
Err(e) => {
warn!(
"Ignoring non-JSON-RPC stdio message from MCP server '{}': {}",
server_id, e
);
continue;
}
};
let Some(id) = response.id else {
if bootstrap_mode.load(Ordering::SeqCst) {
let mut waiters_guard = waiters.lock().await;
if is_acceptable_bootstrap_response(&response, waiters_guard.len()) {
if let Some((_, sender)) = waiters_guard.drain().next() {
tracing::debug!(
"Accepting MCP response with missing id during bootstrap for server '{}'",
server_id
);
let _ = sender.send(Ok(response));
}
continue;
}
}
tracing::debug!(
"Received notification from MCP server '{}': {}",
server_id,
trimmed
);
continue;
};
let mut waiters_guard = waiters.lock().await;
if let Some(sender) = waiters_guard.remove(&id) {
let _ = sender.send(Ok(response));
} else {
tracing::debug!(
"No waiter found for response id {} from server '{}'",
id,
server_id
);
}
}
Err(e) => {
fail_pending_waiters(
&waiters,
format!(
"Failed reading stdio response from server '{}': {}{}",
server_id,
e,
stderr_suffix_async(&stderr_excerpt).await
),
)
.await;
break;
}
}
}
}
async fn start_stdio_stderr_reader(
server_id: String,
stderr: tokio::process::ChildStderr,
stderr_excerpt: Arc<Mutex<VecDeque<String>>>,
) {
let mut stderr = BufReader::new(stderr);
let mut line = String::new();
loop {
line.clear();
match stderr.read_line(&mut line).await {
Ok(0) => break,
Ok(_) => push_stderr_excerpt(&stderr_excerpt, line.trim_end()).await,
Err(e) => {
warn!(
"Failed reading stderr from MCP server '{}': {}",
server_id, e
);
break;
}
}
}
}
async fn fail_pending_waiters(waiters: &WaiterMap, message: String) {
let mut guard = waiters.lock().await;
for (_, sender) in guard.drain() {
let _ = sender.send(Err(message.clone()));
}
}
const SENSITIVE_SUFFIXES: &[&str] = &[
"_API_KEY",
"_SECRET",
"_SECRET_KEY",
"_TOKEN",
"_PASSWORD",
"_CREDENTIALS",
"_AUTH_TOKEN",
"_ACCESS_KEY",
"_ACCESS_TOKEN",
];
const SENSITIVE_EXACT_NAMES: &[&str] = &[
"AWS_ACCESS_KEY_ID",
"AWS_SECRET_ACCESS_KEY",
"AWS_SESSION_TOKEN",
"AZURE_CLIENT_SECRET",
"GOOGLE_APPLICATION_CREDENTIALS",
"DATABASE_URL",
"GITHUB_TOKEN",
"GH_TOKEN",
"ANTHROPIC_API_KEY",
"OPENAI_API_KEY",
];
fn is_sensitive_env_var(name: &str) -> bool {
let upper = name.to_uppercase();
SENSITIVE_EXACT_NAMES.iter().any(|exact| upper == *exact)
|| SENSITIVE_SUFFIXES
.iter()
.any(|suffix| upper.ends_with(suffix))
}
fn sanitized_stdio_env() -> HashMap<String, String> {
let mut env: HashMap<String, String> = std::env::vars().collect();
let sensitive_keys: Vec<String> = env
.keys()
.filter(|k| is_sensitive_env_var(k))
.cloned()
.collect();
if !sensitive_keys.is_empty() {
tracing::debug!(
"Stripping {} sensitive env var(s) from MCP stdio subprocess: {}",
sensitive_keys.len(),
sensitive_keys.join(", ")
);
for key in &sensitive_keys {
env.remove(key);
}
}
env
}
fn redact_command(command: &str) -> String {
Path::new(command)
.file_name()
.and_then(|name| name.to_str())
.unwrap_or(command)
.to_string()
}
fn format_rpc_error(error: JsonRpcError) -> String {
match error.data {
Some(data) => format!(
"RPC error {}: {} ({})",
error.code,
error.message,
serde_json::to_string(&data)
.unwrap_or_else(|_| "unserializable error data".to_string())
),
None => format!("RPC error {}: {}", error.code, error.message),
}
}
fn is_acceptable_bootstrap_response(response: &JsonRpcResponse, waiter_count: usize) -> bool {
response.id.is_none()
&& response.error.is_none()
&& response.result.is_some()
&& waiter_count == 1
}
async fn push_stderr_excerpt(stderr_excerpt: &Arc<Mutex<VecDeque<String>>>, line: &str) {
let sanitized = sanitize_stderr_line(line);
if sanitized.is_empty() {
return;
}
let mut excerpt = stderr_excerpt.lock().await;
excerpt.push_back(sanitized);
while excerpt.len() > 10 {
excerpt.pop_front();
}
}
async fn stderr_suffix_async(stderr_excerpt: &Arc<Mutex<VecDeque<String>>>) -> String {
let excerpt = stderr_excerpt.lock().await;
format_stderr_suffix(&excerpt)
}
fn stderr_suffix_blocking(stderr_excerpt: &Arc<Mutex<VecDeque<String>>>) -> String {
stderr_excerpt
.try_lock()
.ok()
.map_or_else(String::new, |excerpt| format_stderr_suffix(&excerpt))
}
fn format_stderr_suffix(excerpt: &VecDeque<String>) -> String {
excerpt
.back()
.map(|line| format!(" [stderr excerpt: {}]", line))
.unwrap_or_default()
}
fn sanitize_stderr_line(line: &str) -> String {
let collapsed = line.split_whitespace().collect::<Vec<_>>().join(" ");
let truncated: String = collapsed.chars().take(240).collect();
if collapsed.chars().count() > 240 {
format!("{}…", truncated)
} else {
truncated
}
}
#[async_trait]
impl McpTransportClient for StdioMcpClient {
async fn initialize(&self) -> Result<protocol_messages::InitializeResponse, String> {
let request = protocol_messages::InitializeRequest {
protocol_version: "2024-11-05".to_string(),
capabilities: serde_json::json!({}),
client_info: protocol_messages::ClientInfo {
name: "iron-core".to_string(),
version: env!("CARGO_PKG_VERSION").to_string(),
},
};
let result = self
.send_request("initialize", serde_json::to_value(request).unwrap())
.await?;
self.bootstrap_mode.store(false, Ordering::SeqCst);
self.connected.store(true, Ordering::SeqCst);
Ok(result)
}
async fn list_tools(&self) -> Result<Vec<McpToolInfo>, String> {
let mut cursor = None;
let mut discovered = Vec::new();
loop {
let request = protocol_messages::ListToolsRequest {
cursor: cursor.clone(),
};
let response: protocol_messages::ListToolsResponse = self
.send_request("tools/list", serde_json::to_value(request).unwrap())
.await?;
discovered.extend(response.tools.into_iter().map(|tool| McpToolInfo {
name: tool.name,
description: tool.description,
input_schema: tool.input_schema,
}));
match response.next_cursor {
Some(next_cursor) => cursor = Some(next_cursor),
None => return Ok(discovered),
}
}
}
async fn call_tool(&self, tool_name: &str, arguments: Value) -> Result<Value, String> {
let request = protocol_messages::CallToolRequest {
name: tool_name.to_string(),
arguments,
};
let response: protocol_messages::CallToolResponse = self
.send_request("tools/call", serde_json::to_value(request).unwrap())
.await?;
if response.is_error {
return Err(tool_error_to_string(response.content));
}
Ok(tool_content_to_value(response.content))
}
fn is_connected(&self) -> bool {
self.connected.load(Ordering::SeqCst)
}
async fn close(&self) {
self.connected.store(false, Ordering::SeqCst);
self.waiters.lock().await.clear();
if let Some(handle) = self.stdout_task.lock().await.take() {
handle.abort();
}
if let Some(handle) = self.stderr_task.lock().await.take() {
handle.abort();
}
if let Ok(mut process) = self.process.try_lock() {
let _ = process.kill().await;
}
}
}
pub struct HttpMcpClient {
#[allow(dead_code)]
server_id: String,
url: String,
headers: Option<HashMap<String, String>>,
client: reqwest::Client,
bootstrap_mode: AtomicBool,
connected: AtomicBool,
request_counter: Arc<Mutex<u64>>,
}
impl HttpMcpClient {
pub fn new(server_id: String, config: HttpConfig) -> Self {
Self {
server_id,
url: config.url,
headers: config.headers,
client: reqwest::Client::new(),
bootstrap_mode: AtomicBool::new(true),
connected: AtomicBool::new(false),
request_counter: Arc::new(Mutex::new(0)),
}
}
async fn send_request<T: serde::de::DeserializeOwned>(
&self,
method: &str,
params: Value,
) -> Result<T, String> {
let id = {
let mut counter = self.request_counter.lock().await;
*counter += 1;
*counter
};
let request = ProtocolJsonRpcRequest::new(method, params, id);
let builder = self.client.post(&self.url).json(&request);
let response = apply_headers(builder, &self.headers)
.send()
.await
.map_err(|e| format!("HTTP request failed: {}", e))?;
let rpc_response = decode_http_rpc_response(response).await?;
if rpc_response.id != Some(id) {
if method == "initialize"
&& self.bootstrap_mode.load(Ordering::SeqCst)
&& rpc_response.id.is_none()
&& rpc_response.error.is_none()
&& rpc_response.result.is_some()
{
tracing::debug!(
"Accepting MCP response with missing id during bootstrap for server '{}'",
self.server_id
);
} else {
return Err(format!(
"Response ID mismatch: expected {}, got {:?}",
id, rpc_response.id
));
}
}
if let Some(error) = rpc_response.error {
return Err(format_rpc_error(error));
}
rpc_response
.result
.ok_or_else(|| "Response missing result".to_string())
.and_then(|v| {
serde_json::from_value(v)
.map_err(|e| format!("Failed to deserialize result: {}", e))
})
}
}
#[async_trait]
impl McpTransportClient for HttpMcpClient {
async fn initialize(&self) -> Result<protocol_messages::InitializeResponse, String> {
let request = protocol_messages::InitializeRequest {
protocol_version: "2024-11-05".to_string(),
capabilities: serde_json::json!({}),
client_info: protocol_messages::ClientInfo {
name: "iron-core".to_string(),
version: env!("CARGO_PKG_VERSION").to_string(),
},
};
let result = self
.send_request("initialize", serde_json::to_value(request).unwrap())
.await?;
self.bootstrap_mode.store(false, Ordering::SeqCst);
self.connected.store(true, Ordering::SeqCst);
Ok(result)
}
async fn list_tools(&self) -> Result<Vec<McpToolInfo>, String> {
let mut cursor = None;
let mut discovered = Vec::new();
loop {
let request = protocol_messages::ListToolsRequest {
cursor: cursor.clone(),
};
let response: protocol_messages::ListToolsResponse = self
.send_request("tools/list", serde_json::to_value(request).unwrap())
.await?;
discovered.extend(response.tools.into_iter().map(|tool| McpToolInfo {
name: tool.name,
description: tool.description,
input_schema: tool.input_schema,
}));
match response.next_cursor {
Some(next_cursor) => cursor = Some(next_cursor),
None => return Ok(discovered),
}
}
}
async fn call_tool(&self, tool_name: &str, arguments: Value) -> Result<Value, String> {
let request = protocol_messages::CallToolRequest {
name: tool_name.to_string(),
arguments,
};
let response: protocol_messages::CallToolResponse = self
.send_request("tools/call", serde_json::to_value(request).unwrap())
.await?;
if response.is_error {
return Err(tool_error_to_string(response.content));
}
Ok(tool_content_to_value(response.content))
}
fn is_connected(&self) -> bool {
self.connected.load(Ordering::SeqCst)
}
async fn close(&self) {
self.connected.store(false, Ordering::SeqCst);
}
}
pub struct HttpSseMcpClient {
server_id: String,
url: String,
headers: Option<HashMap<String, String>>,
client: reqwest::Client,
connected: AtomicBool,
request_counter: Arc<Mutex<u64>>,
waiters: Arc<Mutex<HashMap<u64, tokio::sync::oneshot::Sender<JsonRpcResponse>>>>,
sse_task: Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
bootstrap_mode: Arc<AtomicBool>,
}
impl HttpSseMcpClient {
pub fn new(server_id: String, config: HttpConfig) -> Self {
Self {
server_id,
url: config.url,
headers: config.headers,
client: reqwest::Client::new(),
connected: AtomicBool::new(false),
request_counter: Arc::new(Mutex::new(0)),
waiters: Arc::new(Mutex::new(HashMap::new())),
sse_task: Arc::new(Mutex::new(None)),
bootstrap_mode: Arc::new(AtomicBool::new(true)),
}
}
async fn ensure_sse_reader(&self) -> Result<(), String> {
let mut task_guard = self.sse_task.lock().await;
if task_guard.is_some() {
return Ok(());
}
let url = self.url.clone();
let client = self.client.clone();
let headers = self.headers.clone();
let waiters: Arc<Mutex<HashMap<u64, tokio::sync::oneshot::Sender<JsonRpcResponse>>>> =
self.waiters.clone();
let bootstrap_mode = self.bootstrap_mode.clone();
let (ready_tx, ready_rx) = tokio::sync::oneshot::channel::<Result<(), String>>();
let handle = tokio::spawn(async move {
let get_builder = client.get(&url);
let response = match tokio::time::timeout(
tokio::time::Duration::from_secs(3),
apply_headers(get_builder, &headers).send(),
)
.await
{
Ok(Ok(resp)) => resp,
Ok(Err(e)) => {
let _ = ready_tx.send(Err(format!(
"Failed to establish SSE endpoint for MCP server: {}",
e
)));
error!("Failed to connect to SSE endpoint: {}", e);
return;
}
Err(_) => {
let _ = ready_tx.send(Err(
"Timed out establishing SSE endpoint for MCP server".to_string(),
));
error!("Timed out connecting to SSE endpoint");
return;
}
};
if !response.status().is_success() {
let _ = ready_tx.send(Err(format!(
"SSE endpoint returned unsuccessful status: {}",
response.status()
)));
error!("SSE connection failed with status: {}", response.status());
return;
}
if ready_tx.send(Ok(())).is_err() {
return;
}
let mut stream = response.bytes_stream();
let mut buffer = String::new();
while let Some(chunk) = stream.next().await {
match chunk {
Ok(bytes) => {
buffer.push_str(&String::from_utf8_lossy(&bytes));
while let Some(pos) = buffer.find("\n\n") {
let event_block = buffer[..pos].to_string();
buffer.drain(..pos + 2);
let mut event_type: Option<String> = None;
let mut data_lines: Vec<String> = Vec::new();
for line in event_block.lines() {
if line.starts_with(':') || line.is_empty() {
continue;
}
if let Some(rest) = line.strip_prefix("event:") {
event_type = Some(rest.trim_start().to_string());
continue;
}
if let Some(rest) = line.strip_prefix("data:") {
data_lines.push(rest.trim_start().to_string());
}
}
if data_lines.is_empty() {
continue;
}
match event_type.as_deref() {
Some("ping") | Some("keepalive") => continue,
_ => {}
}
let payload = data_lines.join("\n");
let rpc_response: JsonRpcResponse = match serde_json::from_str(&payload)
{
Ok(resp) => resp,
Err(_) => {
continue;
}
};
let Some(id) = rpc_response.id else {
if bootstrap_mode.load(Ordering::SeqCst) {
let mut waiters_guard = waiters.lock().await;
if is_acceptable_bootstrap_response(
&rpc_response,
waiters_guard.len(),
) {
if let Some((_, sender)) = waiters_guard.drain().next() {
tracing::debug!(
"Accepting MCP SSE response with missing id during bootstrap"
);
let _ = sender.send(rpc_response);
}
}
}
continue;
};
let mut waiters_guard = waiters.lock().await;
if let Some(sender) = waiters_guard.remove(&id) {
let _ = sender.send(rpc_response);
}
}
}
Err(e) => {
error!("SSE stream error: {}", e);
break;
}
}
}
waiters.lock().await.clear();
});
*task_guard = Some(handle);
drop(task_guard);
tokio::time::timeout(tokio::time::Duration::from_secs(4), ready_rx)
.await
.map_err(|_| "Timed out waiting for SSE startup confirmation".to_string())?
.map_err(|_| "SSE reader exited before startup confirmation".to_string())?
}
async fn send_request<T: serde::de::DeserializeOwned>(
&self,
method: &str,
params: Value,
) -> Result<T, String> {
let id = {
let mut counter = self.request_counter.lock().await;
*counter += 1;
*counter
};
let request = ProtocolJsonRpcRequest::new(method, params, id);
self.ensure_sse_reader().await?;
let (tx, rx) = tokio::sync::oneshot::channel::<JsonRpcResponse>();
{
let mut waiters = self.waiters.lock().await;
waiters.insert(id, tx);
}
let post_builder = self.client.post(&self.url).json(&request);
let post_response = apply_headers(post_builder, &self.headers)
.send()
.await
.map_err(|e| {
self.cleanup_waiter(id);
format!("HTTP+SSE post failed: {}", e)
})?;
if !post_response.status().is_success() && post_response.status() != 202 {
self.cleanup_waiter(id);
return Err(format!("HTTP+SSE post error: {}", post_response.status()));
}
let rpc_response = tokio::time::timeout(tokio::time::Duration::from_secs(30), rx)
.await
.map_err(|_| {
self.cleanup_waiter(id);
format!(
"Timeout waiting for SSE response for request {} on server '{}'",
id, self.server_id
)
})?
.map_err(|_| {
self.cleanup_waiter(id);
format!(
"SSE reader dropped before delivering response for request {} on server '{}'",
id, self.server_id
)
})?;
if let Some(error) = rpc_response.error {
return Err(format_rpc_error(error));
}
rpc_response
.result
.ok_or_else(|| "Response missing result".to_string())
.and_then(|v| {
serde_json::from_value(v)
.map_err(|e| format!("Failed to deserialize result: {}", e))
})
}
fn cleanup_waiter(&self, id: u64) {
if let Ok(mut waiters) = self.waiters.try_lock() {
waiters.remove(&id);
}
}
}
#[async_trait]
impl McpTransportClient for HttpSseMcpClient {
async fn initialize(&self) -> Result<protocol_messages::InitializeResponse, String> {
let request = protocol_messages::InitializeRequest {
protocol_version: "2024-11-05".to_string(),
capabilities: serde_json::json!({}),
client_info: protocol_messages::ClientInfo {
name: "iron-core".to_string(),
version: env!("CARGO_PKG_VERSION").to_string(),
},
};
let result = self
.send_request("initialize", serde_json::to_value(request).unwrap())
.await?;
self.bootstrap_mode.store(false, Ordering::SeqCst);
self.connected.store(true, Ordering::SeqCst);
Ok(result)
}
async fn list_tools(&self) -> Result<Vec<McpToolInfo>, String> {
let mut cursor = None;
let mut discovered = Vec::new();
loop {
let request = protocol_messages::ListToolsRequest {
cursor: cursor.clone(),
};
let response: protocol_messages::ListToolsResponse = self
.send_request("tools/list", serde_json::to_value(request).unwrap())
.await?;
discovered.extend(response.tools.into_iter().map(|tool| McpToolInfo {
name: tool.name,
description: tool.description,
input_schema: tool.input_schema,
}));
match response.next_cursor {
Some(next_cursor) => cursor = Some(next_cursor),
None => return Ok(discovered),
}
}
}
async fn call_tool(&self, tool_name: &str, arguments: Value) -> Result<Value, String> {
let request = protocol_messages::CallToolRequest {
name: tool_name.to_string(),
arguments,
};
let response: protocol_messages::CallToolResponse = self
.send_request("tools/call", serde_json::to_value(request).unwrap())
.await?;
if response.is_error {
return Err(tool_error_to_string(response.content));
}
Ok(tool_content_to_value(response.content))
}
fn is_connected(&self) -> bool {
self.connected.load(Ordering::SeqCst)
}
async fn close(&self) {
self.connected.store(false, Ordering::SeqCst);
if let Some(handle) = self.sse_task.lock().await.take() {
handle.abort();
}
self.waiters.lock().await.clear();
}
}
pub fn create_transport_client(
server_id: &str,
config: &McpServerConfig,
) -> Result<Box<dyn McpTransportClient>, String> {
match &config.transport {
McpTransport::Stdio { .. } => {
let client = StdioMcpClient::new(server_id.to_string(), config.clone())?;
Ok(Box::new(client))
}
McpTransport::Http {
config: http_config,
} => Ok(Box::new(HttpMcpClient::new(
server_id.to_string(),
http_config.clone(),
))),
McpTransport::HttpSse {
config: http_config,
} => Ok(Box::new(HttpSseMcpClient::new(
server_id.to_string(),
http_config.clone(),
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn suffix_api_key_is_sensitive() {
assert!(is_sensitive_env_var("MY_SERVICE_API_KEY"));
assert!(is_sensitive_env_var("SOMETHING_API_KEY"));
}
#[test]
fn suffix_secret_is_sensitive() {
assert!(is_sensitive_env_var("APP_SECRET"));
assert!(is_sensitive_env_var("CLIENT_SECRET_KEY"));
}
#[test]
fn suffix_token_is_sensitive() {
assert!(is_sensitive_env_var("SESSION_TOKEN"));
assert!(is_sensitive_env_var("AUTH_TOKEN"));
}
#[test]
fn suffix_password_is_sensitive() {
assert!(is_sensitive_env_var("DB_PASSWORD"));
}
#[test]
fn suffix_credentials_is_sensitive() {
assert!(is_sensitive_env_var("GOOGLE_CREDENTIALS"));
}
#[test]
fn suffix_access_key_is_sensitive() {
assert!(is_sensitive_env_var("S3_ACCESS_KEY"));
assert!(is_sensitive_env_var("S3_ACCESS_TOKEN"));
}
#[test]
fn exact_aws_keys_are_sensitive() {
assert!(is_sensitive_env_var("AWS_ACCESS_KEY_ID"));
assert!(is_sensitive_env_var("AWS_SECRET_ACCESS_KEY"));
assert!(is_sensitive_env_var("AWS_SESSION_TOKEN"));
}
#[test]
fn exact_cloud_credentials_are_sensitive() {
assert!(is_sensitive_env_var("AZURE_CLIENT_SECRET"));
assert!(is_sensitive_env_var("GOOGLE_APPLICATION_CREDENTIALS"));
}
#[test]
fn exact_token_names_are_sensitive() {
assert!(is_sensitive_env_var("GITHUB_TOKEN"));
assert!(is_sensitive_env_var("GH_TOKEN"));
}
#[test]
fn exact_api_keys_are_sensitive() {
assert!(is_sensitive_env_var("ANTHROPIC_API_KEY"));
assert!(is_sensitive_env_var("OPENAI_API_KEY"));
}
#[test]
fn database_url_is_sensitive() {
assert!(is_sensitive_env_var("DATABASE_URL"));
}
#[test]
fn common_toolchain_vars_are_not_sensitive() {
assert!(!is_sensitive_env_var("PATH"));
assert!(!is_sensitive_env_var("HOME"));
assert!(!is_sensitive_env_var("APPDATA"));
assert!(!is_sensitive_env_var("LOCALAPPDATA"));
assert!(!is_sensitive_env_var("USERPROFILE"));
assert!(!is_sensitive_env_var("XDG_CONFIG_HOME"));
assert!(!is_sensitive_env_var("XDG_DATA_HOME"));
assert!(!is_sensitive_env_var("CARGO_HOME"));
assert!(!is_sensitive_env_var("GOPATH"));
assert!(!is_sensitive_env_var("NODE_PATH"));
assert!(!is_sensitive_env_var("PYTHONPATH"));
assert!(!is_sensitive_env_var("LANG"));
assert!(!is_sensitive_env_var("TERM"));
assert!(!is_sensitive_env_var("SYSTEMROOT"));
}
#[test]
fn case_insensitive_suffix_matching() {
assert!(is_sensitive_env_var("my_service_api_key"));
assert!(is_sensitive_env_var("My_Service_Api_Key"));
assert!(is_sensitive_env_var("APP_Secret"));
assert!(is_sensitive_env_var("Session_Token"));
}
#[test]
fn case_insensitive_exact_matching() {
assert!(is_sensitive_env_var("github_token"));
assert!(is_sensitive_env_var("Github_Token"));
assert!(is_sensitive_env_var("anthropic_api_key"));
assert!(is_sensitive_env_var("Anthropic_Api_Key"));
assert!(is_sensitive_env_var("database_url"));
assert!(is_sensitive_env_var("Database_Url"));
}
#[test]
fn sanitized_env_preserves_non_sensitive_vars() {
let env = sanitized_stdio_env();
assert!(
env.contains_key("PATH"),
"PATH should be preserved in sanitized env"
);
}
#[test]
fn sanitized_env_strips_sensitive_test_vars() {
let test_vars = vec![
("TEST_IRON_CORE_MY_API_KEY", "secret-key-123"),
("TEST_IRON_CORE_AUTH_TOKEN", "tok-456"),
("TEST_IRON_CORE_DB_PASSWORD", "pw-789"),
("TEST_IRON_CORE_GITHUB_TOKEN", "gh-token"),
];
for (key, value) in &test_vars {
std::env::set_var(key, value);
}
let env = sanitized_stdio_env();
for (key, _) in &test_vars {
assert!(
!env.contains_key(*key),
"sensitive var '{}' should have been stripped",
key
);
std::env::remove_var(key);
}
}
#[test]
fn sanitized_env_preserves_toolchain_test_vars() {
let test_vars = vec![
("TEST_IRON_CORE_CARGO_HOME", "/cargo"),
("TEST_IRON_CORE_GOPATH", "/go"),
("TEST_IRON_CORE_NODE_PATH", "/node"),
("TEST_IRON_CORE_XDG_CONFIG_HOME", "/xdg"),
];
for (key, value) in &test_vars {
std::env::set_var(key, value);
}
let env = sanitized_stdio_env();
for (key, expected) in &test_vars {
assert_eq!(
env.get(*key),
Some(&expected.to_string()),
"non-sensitive var '{}' should be preserved",
key
);
std::env::remove_var(key);
}
}
}