fix(mcp): address 14 audit findings across MCP module (#1094)

* fix(mcp): address 14 audit findings across MCP module

- Replace panicking assert! in new_with_config with Result return (Critical)
- Fix initialize() race condition using tokio::sync::OnceCell (High)
- Fix localhost check bypass via proper URL parsing (High)
- Extract shared stream_transport_send() to deduplicate stdio/unix send logic
- Use atomic write (tmp+rename) for config file persistence
- Filter SSE responses by request_id to prevent wrong-response dispatch
- Share a single reqwest::Client for OAuth via fallible OnceLock
- Log notification send errors instead of silently discarding
- Fix unwrap_or(0) that could steal id=0 responses
- Store InitializeResult in OnceCell so callers can access server capabilities
- Add redirect logging in OAuth discovery
- Reuse is_localhost_url() in auth.rs
- Add McpToolWrapper unit tests and regression tests
- URL-encode PKCE challenge for consistency

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* chore: retrigger CI with skip-regression-check label

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 <[email protected]>
This commit is contained in:
Illia Polosukhin
2026-03-13 17:37:51 +00:00
committed by GitHub
co-authored by Claude Opus 4.6
parent bc6725205a
commit f53c1bb10b
8 changed files with 450 additions and 235 deletions
+60 -40
View File
@@ -18,6 +18,44 @@ use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::tools::mcp::config::McpServerConfig;
/// Shared HTTP client for all OAuth/discovery requests.
///
/// Redirects are disabled for security (prevents redirect-based SSRF).
/// Per-request timeouts can override the default via `.timeout()` on
/// the request builder.
fn oauth_http_client() -> Result<&'static reqwest::Client, AuthError> {
static CLIENT: std::sync::OnceLock<Result<reqwest::Client, String>> =
std::sync::OnceLock::new();
CLIENT
.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| e.to_string())
})
.as_ref()
.map_err(|e| AuthError::Http(e.clone()))
}
/// Log a debug message when a discovery/auth response is a redirect.
/// Helps users diagnose configuration issues when legitimate servers
/// redirect and our no-redirect policy causes a failure.
fn log_redirect_if_applicable(url: &str, response: &reqwest::Response) {
if response.status().is_redirection() {
let location = response
.headers()
.get("location")
.and_then(|v| v.to_str().ok());
tracing::debug!(
"OAuth request to '{}' returned redirect {} -> {:?} (redirects disabled for security)",
url,
response.status(),
location
);
}
}
/// OAuth authorization error.
#[derive(Debug, thiserror::Error)]
pub enum AuthError {
@@ -287,10 +325,8 @@ async fn validate_url_safe(url: &str) -> Result<(), AuthError> {
)));
}
if scheme == "http" {
let host = parsed.host_str().unwrap_or("");
let is_localhost =
host == "localhost" || host == "127.0.0.1" || host == "::1" || host == "[::1]";
if !is_localhost {
if !crate::tools::mcp::config::is_localhost_url(url) {
let host = parsed.host_str().unwrap_or("");
return Err(AuthError::DiscoveryFailed(format!(
"HTTP is only allowed for localhost; use HTTPS for '{}'",
host
@@ -382,18 +418,17 @@ fn parse_resource_metadata_url(www_authenticate: &str) -> Option<String> {
async fn fetch_resource_metadata(url: &str) -> Result<ProtectedResourceMetadata, AuthError> {
validate_url_safe(url).await?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let client = oauth_http_client()?;
let response = client
.get(url)
.timeout(Duration::from_secs(10))
.send()
.await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(url, &response);
if !response.status().is_success() {
return Err(AuthError::DiscoveryFailed(format!(
"HTTP {}",
@@ -411,20 +446,19 @@ async fn fetch_resource_metadata(url: &str) -> Result<ProtectedResourceMetadata,
async fn discover_via_401(server_url: &str) -> Result<AuthorizationServerMetadata, AuthError> {
validate_url_safe(server_url).await?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let client = oauth_http_client()?;
let response = client
.post(server_url)
.timeout(Duration::from_secs(10))
.header("Content-Type", "application/json")
.body("{}")
.send()
.await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(server_url, &response);
if response.status().as_u16() != 401 {
return Err(AuthError::DiscoveryFailed(format!(
"Expected 401, got {}",
@@ -472,20 +506,19 @@ pub async fn discover_protected_resource(
) -> Result<ProtectedResourceMetadata, AuthError> {
validate_url_safe(server_url).await?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let client = oauth_http_client()?;
let well_known_url = build_well_known_uri(server_url, "oauth-protected-resource")?;
let response = client
.get(&well_known_url)
.timeout(Duration::from_secs(10))
.send()
.await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(&well_known_url, &response);
if !response.status().is_success() {
return Err(AuthError::NotSupported);
}
@@ -502,20 +535,19 @@ pub async fn discover_authorization_server(
) -> Result<AuthorizationServerMetadata, AuthError> {
validate_url_safe(auth_server_url).await?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let client = oauth_http_client()?;
let well_known_url = build_well_known_uri(auth_server_url, "oauth-authorization-server")?;
let response = client
.get(&well_known_url)
.timeout(Duration::from_secs(10))
.send()
.await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(&well_known_url, &response);
if !response.status().is_success() {
return Err(AuthError::DiscoveryFailed(format!(
"HTTP {}",
@@ -595,11 +627,7 @@ pub async fn register_client(
) -> Result<ClientRegistrationResponse, AuthError> {
validate_url_safe(registration_endpoint).await?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let client = oauth_http_client()?;
let request = ClientRegistrationRequest {
client_name: "IronClaw".to_string(),
@@ -813,7 +841,7 @@ pub fn build_authorization_url(
if let Some(pkce) = pkce {
url.push_str(&format!(
"&code_challenge={}&code_challenge_method=S256",
pkce.challenge
urlencoding::encode(&pkce.challenge)
));
}
@@ -863,11 +891,7 @@ pub async fn exchange_code_for_token(
) -> Result<AccessToken, AuthError> {
validate_url_safe(token_url).await?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let client = oauth_http_client()?;
let mut params = vec![
("grant_type", "authorization_code".to_string()),
@@ -1054,11 +1078,7 @@ pub async fn refresh_access_token(
validate_url_safe(&token_url).await?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let client = oauth_http_client()?;
// Compute canonical resource URI for RFC 8707
let resource = canonical_resource_uri(&server_config.url);
+205 -63
View File
@@ -5,7 +5,7 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::atomic::{AtomicU64, Ordering};
use async_trait::async_trait;
use tokio::sync::RwLock;
@@ -58,9 +58,10 @@ pub struct McpClient {
/// Custom headers to include in every request.
custom_headers: HashMap<String, String>,
/// Whether the MCP initialize handshake has completed.
/// Used as a local idempotency guard when no session_manager is present.
initialized: AtomicBool,
/// Ensures the MCP initialize handshake runs exactly once.
/// Uses `OnceCell` to serialize concurrent callers so only one
/// actually sends the request; subsequent calls return immediately.
initialized: tokio::sync::OnceCell<InitializeResult>,
}
impl McpClient {
@@ -83,7 +84,7 @@ impl McpClient {
user_id: "default".to_string(),
server_config: None,
custom_headers: HashMap::new(),
initialized: AtomicBool::new(false),
initialized: tokio::sync::OnceCell::new(),
}
}
@@ -106,7 +107,7 @@ impl McpClient {
user_id: "default".to_string(),
server_config: None,
custom_headers: HashMap::new(),
initialized: AtomicBool::new(false),
initialized: tokio::sync::OnceCell::new(),
}
}
@@ -114,20 +115,24 @@ impl McpClient {
///
/// Use this when you have an `McpServerConfig` with custom headers but no OAuth.
/// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`.
pub fn new_with_config(config: McpServerConfig) -> Self {
assert!(
matches!(
config.effective_transport(),
crate::tools::mcp::config::EffectiveTransport::Http
),
"new_with_config only supports HTTP transport; use new_with_transport for stdio/UDS"
);
///
/// Returns an error if the config uses a non-HTTP transport.
pub fn new_with_config(config: McpServerConfig) -> Result<Self, ToolError> {
if !matches!(
config.effective_transport(),
crate::tools::mcp::config::EffectiveTransport::Http
) {
return Err(ToolError::InvalidParameters(
"new_with_config only supports HTTP transport; use new_with_transport for stdio/UDS"
.to_string(),
));
}
let transport = Arc::new(HttpMcpTransport::new(
config.url.clone(),
config.name.clone(),
));
Self {
Ok(Self {
transport,
server_url: config.url.clone(),
server_name: config.name.clone(),
@@ -137,9 +142,9 @@ impl McpClient {
secrets: None,
user_id: "default".to_string(),
custom_headers: config.headers.clone(),
initialized: AtomicBool::new(false),
initialized: tokio::sync::OnceCell::new(),
server_config: Some(config),
}
})
}
/// Create a new authenticated MCP client.
@@ -169,7 +174,7 @@ impl McpClient {
user_id: user_id.into(),
server_config: Some(config),
custom_headers,
initialized: AtomicBool::new(false),
initialized: tokio::sync::OnceCell::new(),
}
}
@@ -205,7 +210,7 @@ impl McpClient {
user_id: user_id.into(),
server_config,
custom_headers,
initialized: AtomicBool::new(false),
initialized: tokio::sync::OnceCell::new(),
}
}
@@ -336,53 +341,64 @@ impl McpClient {
}
/// Initialize the connection to the MCP server.
///
/// Uses `OnceCell` to guarantee that exactly one caller performs the
/// handshake, even under concurrent access. Subsequent calls return
/// immediately.
pub async fn initialize(&self) -> Result<InitializeResult, ToolError> {
// Fast path: already initialized (local flag or session manager)
if self.initialized.load(Ordering::Relaxed) {
return Ok(InitializeResult::default());
}
if let Some(ref session_manager) = self.session_manager
&& session_manager.is_initialized(&self.server_name).await
{
self.initialized.store(true, Ordering::Relaxed);
return Ok(InitializeResult::default());
}
if let Some(ref session_manager) = self.session_manager {
session_manager
.get_or_create(&self.server_name, &self.server_url)
.await;
}
let result = self
.initialized
.get_or_try_init(|| async {
if let Some(ref session_manager) = self.session_manager
&& session_manager.is_initialized(&self.server_name).await
{
return Ok(InitializeResult::default());
}
if let Some(ref session_manager) = self.session_manager {
session_manager
.get_or_create(&self.server_name, &self.server_url)
.await;
}
let request = McpRequest::initialize(self.next_request_id());
let response = self.send_request(request).await?;
let request = McpRequest::initialize(self.next_request_id());
let response = self.send_request(request).await?;
if let Some(error) = response.error {
return Err(ToolError::ExternalService(format!(
"MCP initialization error: {} (code {})",
error.message, error.code
)));
}
if let Some(error) = response.error {
return Err(ToolError::ExternalService(format!(
"MCP initialization error: {} (code {})",
error.message, error.code
)));
}
let result: InitializeResult = response
.result
.ok_or_else(|| {
ToolError::ExternalService("No result in initialize response".to_string())
let init_result: InitializeResult = response
.result
.ok_or_else(|| {
ToolError::ExternalService("No result in initialize response".to_string())
})
.and_then(|r| {
serde_json::from_value(r).map_err(|e| {
ToolError::ExternalService(format!("Invalid initialize result: {}", e))
})
})?;
if let Some(ref session_manager) = self.session_manager {
session_manager.mark_initialized(&self.server_name).await;
}
let notification = McpRequest::initialized_notification();
if let Err(e) = self.send_request(notification).await {
tracing::debug!(
"Failed to send initialized notification to '{}': {}",
self.server_name,
e
);
}
Ok(init_result)
})
.and_then(|r| {
serde_json::from_value(r).map_err(|e| {
ToolError::ExternalService(format!("Invalid initialize result: {}", e))
})
})?;
.await?;
if let Some(ref session_manager) = self.session_manager {
session_manager.mark_initialized(&self.server_name).await;
}
self.initialized.store(true, Ordering::Relaxed);
let notification = McpRequest::initialized_notification();
let _ = self.send_request(notification).await;
Ok(result)
Ok(result.clone())
}
/// List available tools from the MCP server.
@@ -471,6 +487,11 @@ impl McpClient {
}
}
/// Clone the client, resetting the tools cache and initialization state.
/// The cloned client shares the same transport and session manager, so
/// re-initialization will short-circuit via the session manager check if
/// the source was already initialized. The `next_id` counter is copied
/// so that cloned clients continue with monotonically increasing IDs.
impl Clone for McpClient {
fn clone(&self) -> Self {
Self {
@@ -484,7 +505,7 @@ impl Clone for McpClient {
user_id: self.user_id.clone(),
server_config: self.server_config.clone(),
custom_headers: self.custom_headers.clone(),
initialized: AtomicBool::new(self.initialized.load(Ordering::Relaxed)),
initialized: tokio::sync::OnceCell::new(),
}
}
}
@@ -707,7 +728,7 @@ mod tests {
headers.insert("X-Custom".to_string(), "value".to_string());
let config = McpServerConfig::new("test", "http://localhost:8080").with_headers(headers);
let client = McpClient::new_with_config(config.clone());
let client = McpClient::new_with_config(config.clone()).expect("HTTP config should work");
assert_eq!(client.server_name(), "test");
assert_eq!(client.server_url(), "http://localhost:8080");
@@ -719,7 +740,7 @@ mod tests {
#[test]
fn test_new_with_config_no_headers() {
let config = McpServerConfig::new("bare", "http://localhost:9090");
let client = McpClient::new_with_config(config);
let client = McpClient::new_with_config(config).expect("HTTP config should work");
assert_eq!(client.server_name(), "bare");
assert!(client.custom_headers.is_empty());
@@ -971,4 +992,125 @@ mod tests {
assert_eq!(obj.len(), 1);
assert!(obj["outer"]["inner"].is_null());
}
// --- Issue 1 regression: new_with_config rejects non-HTTP transport ---
#[test]
fn test_new_with_config_rejects_stdio_transport() {
let config = McpServerConfig::new_stdio(
"stdio-server",
"echo",
vec!["hello".to_string()],
HashMap::new(),
);
let result = McpClient::new_with_config(config);
let err = result
.err()
.expect("stdio config must be rejected")
.to_string();
assert!(
err.contains("new_with_config only supports HTTP"),
"error should explain the restriction: {}",
err
);
}
// --- Issue 13: McpToolWrapper unit tests ---
fn make_test_mcp_tool(destructive: bool) -> McpTool {
use crate::tools::mcp::protocol::McpToolAnnotations;
McpTool {
name: "do_thing".to_string(),
description: "Does a thing".to_string(),
input_schema: serde_json::json!({
"type": "object",
"properties": {
"input": {"type": "string"}
}
}),
annotations: if destructive {
Some(McpToolAnnotations {
destructive_hint: true,
side_effects_hint: false,
read_only_hint: false,
execution_time_hint: None,
})
} else {
None
},
}
}
#[test]
fn test_mcp_tool_wrapper_name_is_prefixed() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__myserver__do_thing".to_string(),
client,
};
assert_eq!(wrapper.name(), "mcp__myserver__do_thing");
}
#[test]
fn test_mcp_tool_wrapper_description() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
assert_eq!(wrapper.description(), "Does a thing");
}
#[test]
fn test_mcp_tool_wrapper_parameters_schema() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
let schema = wrapper.parameters_schema();
assert_eq!(schema["type"], "object");
assert!(schema["properties"]["input"].is_object());
}
#[test]
fn test_mcp_tool_wrapper_requires_sanitization() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
assert!(
wrapper.requires_sanitization(),
"MCP tools should always require sanitization"
);
}
#[test]
fn test_mcp_tool_wrapper_approval_destructive() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(true),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
let approval = wrapper.requires_approval(&serde_json::json!({}));
assert_eq!(approval, ApprovalRequirement::UnlessAutoApproved);
}
#[test]
fn test_mcp_tool_wrapper_approval_non_destructive() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
let approval = wrapper.requires_approval(&serde_json::json!({}));
assert_eq!(approval, ApprovalRequirement::Never);
}
}
+38 -6
View File
@@ -163,10 +163,8 @@ impl McpServerConfig {
}
// Remote servers must use HTTPS (localhost is allowed for development)
let url_lower = self.url.to_lowercase();
let is_localhost =
url_lower.contains("localhost") || url_lower.contains("127.0.0.1");
if !is_localhost && !url_lower.starts_with("https://") {
let is_localhost = is_localhost_url(&self.url);
if !is_localhost && !self.url.to_lowercase().starts_with("https://") {
return Err(ConfigError::InvalidConfig {
reason: "Remote MCP servers must use HTTPS".to_string(),
});
@@ -442,7 +440,12 @@ pub async fn save_mcp_servers_to(
}
let content = serde_json::to_string_pretty(config)?;
fs::write(path, content).await?;
// Write to a temporary file first, then atomically rename to avoid
// corrupting the config if the process crashes during the write.
let tmp_path = path.with_extension("json.tmp");
fs::write(&tmp_path, content).await?;
fs::rename(&tmp_path, path).await?;
Ok(())
}
@@ -570,7 +573,7 @@ pub async fn remove_mcp_server_db(
///
/// Uses `url::Url` for proper parsing so edge cases (IPv6, userinfo, ports)
/// are handled correctly without manual string splitting.
fn is_localhost_url(url: &str) -> bool {
pub(crate) fn is_localhost_url(url: &str) -> bool {
let Ok(parsed) = url::Url::parse(url) else {
return false;
};
@@ -1125,4 +1128,33 @@ mod tests {
assert!(parsed.transport.is_none());
assert_eq!(parsed.headers.get("X-Custom").unwrap(), "value");
}
// --- Issue 3 regression: is_localhost_url rejects attacker subdomains ---
#[test]
fn test_is_localhost_url_rejects_attacker_subdomain() {
// Before the fix, url.contains("localhost") matched this.
assert!(
!is_localhost_url("http://evil.localhost.attacker.com:8080/mcp"),
"attacker subdomain containing 'localhost' must not be treated as local"
);
}
#[test]
fn test_is_localhost_url_accepts_real_localhost() {
assert!(is_localhost_url("http://localhost:8080/mcp"));
assert!(is_localhost_url("https://localhost/path"));
}
#[test]
fn test_is_localhost_url_accepts_loopback_ip() {
assert!(is_localhost_url("http://127.0.0.1:3000"));
assert!(is_localhost_url("http://[::1]:3000"));
}
#[test]
fn test_is_localhost_url_rejects_remote() {
assert!(!is_localhost_url("https://mcp.example.com"));
assert!(!is_localhost_url("http://192.168.1.1:8080"));
}
}
+10
View File
@@ -18,6 +18,8 @@ pub enum McpFactoryError {
UnixConnect { name: String, reason: String },
#[error("Unix socket transport is not supported on this platform (server '{name}')")]
UnixNotSupported { name: String },
#[error("Invalid configuration for MCP server '{name}': {reason}")]
InvalidConfig { name: String, reason: String },
}
/// Create an `McpClient` from a server configuration, dispatching on the
@@ -89,10 +91,18 @@ pub async fn create_client_from_config(
))
} else {
Ok(McpClient::new_with_config(server)
.map_err(|e| McpFactoryError::InvalidConfig {
name: server_name.clone(),
reason: e.to_string(),
})?
.with_session_manager(Arc::clone(session_manager)))
}
} else {
Ok(McpClient::new_with_config(server)
.map_err(|e| McpFactoryError::InvalidConfig {
name: server_name,
reason: e.to_string(),
})?
.with_session_manager(Arc::clone(session_manager)))
}
}
+14 -9
View File
@@ -139,7 +139,7 @@ impl McpTransport for HttpMcpTransport {
.to_string();
if content_type.contains("text/event-stream") {
self.parse_sse_response(response).await
self.parse_sse_response(response, request.id).await
} else {
response.json().await.map_err(|e| {
ToolError::ExternalService(format!(
@@ -161,11 +161,14 @@ impl McpTransport for HttpMcpTransport {
}
impl HttpMcpTransport {
/// Parse a Server-Sent Events response, returning the first valid JSON-RPC
/// `data:` line as an [`McpResponse`].
/// Parse a Server-Sent Events response, returning the JSON-RPC response
/// whose `id` matches `request_id`. Non-matching events (e.g. server
/// notifications or progress updates) are skipped so that the caller
/// receives the actual result for its request.
async fn parse_sse_response(
&self,
response: reqwest::Response,
request_id: Option<u64>,
) -> Result<McpResponse, ToolError> {
use futures::StreamExt;
@@ -202,9 +205,10 @@ impl HttpMcpTransport {
remaining_start = i + 1;
if let Some(json_str) = line.strip_prefix("data: ")
&& let Ok(response) = serde_json::from_str::<McpResponse>(json_str)
&& let Ok(resp) = serde_json::from_str::<McpResponse>(json_str)
&& resp.id == request_id
{
return Ok(response);
return Ok(resp);
}
}
}
@@ -216,14 +220,15 @@ impl HttpMcpTransport {
// Process any remaining data without a trailing newline.
if let Some(json_str) = buffer.strip_prefix("data: ")
&& let Ok(response) = serde_json::from_str::<McpResponse>(json_str.trim())
&& let Ok(resp) = serde_json::from_str::<McpResponse>(json_str.trim())
&& resp.id == request_id
{
return Ok(response);
return Ok(resp);
}
Err(ToolError::ExternalService(format!(
"[{}] No valid data in SSE response: {}",
self.server_name, buffer
"[{}] No matching response (id={:?}) in SSE stream",
self.server_name, request_id
)))
}
}
+9 -58
View File
@@ -14,7 +14,7 @@ use tokio::sync::{Mutex, oneshot};
use tokio::task::JoinHandle;
use crate::tools::mcp::protocol::{McpRequest, McpResponse};
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, write_jsonrpc_line};
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, stream_transport_send};
use crate::tools::tool::ToolError;
/// MCP transport that communicates with a child process over stdin/stdout.
@@ -118,63 +118,14 @@ impl McpTransport for StdioMcpTransport {
request: &McpRequest,
_headers: &HashMap<String, String>,
) -> Result<McpResponse, ToolError> {
// JSON-RPC notifications (no id) are fire-and-forget: the server
// will not send a response, so we must not wait for one.
if request.id.is_none() {
let mut stdin = self.stdin.lock().await;
write_jsonrpc_line(&mut *stdin, request).await?;
return Ok(McpResponse {
jsonrpc: "2.0".to_string(),
id: None,
result: None,
error: None,
});
}
let id = request.id.unwrap_or(0);
let (tx, rx) = oneshot::channel();
// Register the pending response handler before writing the request,
// so we don't miss a fast response from the child.
{
let mut pending = self.pending.lock().await;
pending.insert(id, tx);
}
// Write the request to stdin.
{
let mut stdin = self.stdin.lock().await;
if let Err(e) = write_jsonrpc_line(&mut *stdin, request).await {
// Remove the pending entry on write failure.
let mut pending = self.pending.lock().await;
pending.remove(&id);
return Err(e);
}
}
// Wait for the response with a timeout.
let timeout = Duration::from_secs(30);
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(_)) => {
// Sender was dropped (reader task ended). Clean up pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] MCP server closed connection before responding to request {:?}",
self.server_name, request.id
)))
}
Err(_) => {
// Timeout: remove the pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] Timeout waiting for response to request {:?} after {:?}",
self.server_name, request.id, timeout
)))
}
}
stream_transport_send(
&self.stdin,
&self.pending,
request,
&self.server_name,
Duration::from_secs(30),
)
.await
}
async fn shutdown(&self) -> Result<(), ToolError> {
+105 -1
View File
@@ -97,7 +97,13 @@ pub fn spawn_jsonrpc_reader<R: AsyncBufRead + Unpin + Send + 'static>(
}
};
let id = response.id.unwrap_or(0);
let Some(id) = response.id else {
tracing::debug!(
"[{}] Received JSON-RPC notification (no id), skipping dispatch",
server_name
);
continue;
};
let mut map = pending.lock().await;
if let Some(tx) = map.remove(&id) {
// Ignore send error — the receiver may have been dropped (timeout).
@@ -115,6 +121,76 @@ pub fn spawn_jsonrpc_reader<R: AsyncBufRead + Unpin + Send + 'static>(
})
}
/// Send a JSON-RPC request over a stream-based transport (stdio / unix socket).
///
/// Handles notification fire-and-forget, pending response registration,
/// write, timeout, and cleanup. Used by both [`StdioMcpTransport`] and
/// [`UnixMcpTransport`] to avoid duplicating the send logic.
pub(crate) async fn stream_transport_send<W: AsyncWrite + Unpin>(
writer: &Mutex<W>,
pending: &Mutex<HashMap<u64, oneshot::Sender<McpResponse>>>,
request: &McpRequest,
server_name: &str,
timeout_duration: std::time::Duration,
) -> Result<McpResponse, ToolError> {
// JSON-RPC notifications (no id) are fire-and-forget: the server
// will not send a response, so we must not wait for one.
if request.id.is_none() {
let mut w = writer.lock().await;
write_jsonrpc_line(&mut *w, request).await?;
return Ok(McpResponse {
jsonrpc: "2.0".to_string(),
id: None,
result: None,
error: None,
});
}
let id = request.id.unwrap_or(0);
let (tx, rx) = oneshot::channel();
// Register the pending response handler before writing the request,
// so we don't miss a fast response from the server.
{
let mut map = pending.lock().await;
map.insert(id, tx);
}
// Write the request.
{
let mut w = writer.lock().await;
if let Err(e) = write_jsonrpc_line(&mut *w, request).await {
// Remove the pending entry on write failure.
let mut map = pending.lock().await;
map.remove(&id);
return Err(e);
}
}
// Wait for the response with a timeout.
match tokio::time::timeout(timeout_duration, rx).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(_)) => {
// Sender was dropped (reader task ended). Clean up pending entry.
let mut map = pending.lock().await;
map.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] MCP server closed connection before responding to request {:?}",
server_name, request.id
)))
}
Err(_) => {
// Timeout: remove the pending entry.
let mut map = pending.lock().await;
map.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] Timeout waiting for response to request {:?} after {:?}",
server_name, request.id, timeout_duration
)))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
@@ -193,4 +269,32 @@ mod tests {
handle.await.expect("reader task should finish");
}
/// Issue 9 regression: a JSON-RPC notification (no id) must not resolve
/// a pending request keyed by id 0 (the old `unwrap_or(0)` default).
#[tokio::test]
async fn test_notification_does_not_resolve_pending_id_zero() {
// A notification response (no id), followed by a proper response for id 0.
let notification = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{}}"#;
let real_response = r#"{"jsonrpc":"2.0","id":0,"result":{"ok":true}}"#;
let input = format!("{notification}\n{real_response}\n");
let reader = std::io::Cursor::new(input.into_bytes());
let pending: Arc<Mutex<HashMap<u64, oneshot::Sender<McpResponse>>>> =
Arc::new(Mutex::new(HashMap::new()));
let (tx, rx) = oneshot::channel();
{
let mut map = pending.lock().await;
map.insert(0, tx);
}
let handle = spawn_jsonrpc_reader(reader, pending.clone(), "test".into());
let resp = rx.await.expect("should receive the real id=0 response");
assert_eq!(resp.id, Some(0));
assert!(resp.result.is_some());
handle.await.expect("reader task should finish");
}
}
+9 -58
View File
@@ -15,7 +15,7 @@ use tokio::sync::{Mutex, oneshot};
use tokio::task::JoinHandle;
use crate::tools::mcp::protocol::{McpRequest, McpResponse};
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, write_jsonrpc_line};
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, stream_transport_send};
use crate::tools::tool::ToolError;
/// MCP transport that communicates over a Unix domain socket.
@@ -91,63 +91,14 @@ impl McpTransport for UnixMcpTransport {
request: &McpRequest,
_headers: &HashMap<String, String>,
) -> Result<McpResponse, ToolError> {
// JSON-RPC notifications (no id) are fire-and-forget: the server
// will not send a response, so we must not wait for one.
if request.id.is_none() {
let mut writer = self.writer.lock().await;
write_jsonrpc_line(&mut *writer, request).await?;
return Ok(McpResponse {
jsonrpc: "2.0".to_string(),
id: None,
result: None,
error: None,
});
}
let id = request.id.unwrap_or(0);
let (tx, rx) = oneshot::channel();
// Register the pending response handler before writing the request,
// so we don't miss a fast response from the server.
{
let mut pending = self.pending.lock().await;
pending.insert(id, tx);
}
// Write the request to the socket.
{
let mut writer = self.writer.lock().await;
if let Err(e) = write_jsonrpc_line(&mut *writer, request).await {
// Remove the pending entry on write failure.
let mut pending = self.pending.lock().await;
pending.remove(&id);
return Err(e);
}
}
// Wait for the response with a timeout.
let timeout = Duration::from_secs(30);
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(_)) => {
// Sender was dropped (reader task ended). Clean up pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] MCP server closed connection before responding to request {:?}",
self.server_name, request.id
)))
}
Err(_) => {
// Timeout: remove the pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] Timeout waiting for response to request {:?} after {:?}",
self.server_name, request.id, timeout
)))
}
}
stream_transport_send(
&self.writer,
&self.pending,
request,
&self.server_name,
Duration::from_secs(30),
)
.await
}
async fn shutdown(&self) -> Result<(), ToolError> {