use anyhow::{Context, Result, anyhow, bail};
use futures::StreamExt;
use reqwest::StatusCode;
use reqwest::header::{ACCEPT, HeaderMap, HeaderName, HeaderValue};
use serde_json::{Value, json};
use std::sync::RwLock;
use std::sync::atomic::{AtomicU64, Ordering};
use tokio::time::{Duration, timeout};
use super::transport::{
REQUEST_TIMEOUT_SECS, extract_jsonrpc_result, is_response, parse_response_id,
};
use crate::app::{McpServerConfig, TransportKind};
use crate::utils::{HostClass, classify_host, drain_sse_events};
const CONNECT_TIMEOUT_SECS: u64 = 10;
const NOTIFICATION_TIMEOUT_SECS: u64 = 10;
const DELETE_TIMEOUT_SECS: u64 = 5;
const MAX_SSE_RECONNECTS: u32 = 5;
const ERROR_BODY_SNIPPET_BYTES: usize = 200;
const SESSION_HEADER: HeaderName = HeaderName::from_static("mcp-session-id");
const PROTOCOL_VERSION_HEADER: HeaderName = HeaderName::from_static("mcp-protocol-version");
const LAST_EVENT_ID_HEADER: HeaderName = HeaderName::from_static("last-event-id");
const ACCEPT_POST: &str = "application/json, text/event-stream";
const ACCEPT_SSE: &str = "text/event-stream";
pub(super) struct HttpTransport {
client: reqwest::Client,
url: reqwest::Url,
static_headers: HeaderMap,
env_headers: Vec<(HeaderName, String)>,
session_id: RwLock<Option<HeaderValue>>,
protocol_version: RwLock<Option<HeaderValue>>,
next_id: AtomicU64,
}
impl HttpTransport {
pub fn new(config: &McpServerConfig) -> Result<Self> {
if config.transport_kind()? != TransportKind::Http {
bail!("MCP server config is not url-shaped");
}
let url_str = config
.url
.as_deref()
.expect("transport_kind checked url presence");
let url = reqwest::Url::parse(url_str)
.map_err(|e| anyhow!("invalid MCP server url '{url_str}': {e}"))?;
let host = url.host_str().unwrap_or_default();
let blocked = match classify_host(host) {
HostClass::Loopback | HostClass::Public => false,
_ => !config.allow_private_network,
};
if blocked {
bail!(
"refusing to connect MCP server '{host}': it is a private/internal \
address (set allow_private_network = true for this server to permit it)"
);
}
let mut static_headers = HeaderMap::new();
for (name, value) in &config.headers {
let n = HeaderName::from_bytes(name.as_bytes())
.map_err(|_| anyhow!("invalid MCP header name '{name}'"))?;
let v = HeaderValue::from_str(value)
.map_err(|_| anyhow!("invalid value for MCP header '{name}'"))?;
static_headers.insert(n, v);
}
let mut env_headers = Vec::new();
for (name, var) in &config.env_headers {
let n = HeaderName::from_bytes(name.as_bytes())
.map_err(|_| anyhow!("invalid MCP header name '{name}'"))?;
env_headers.push((n, var.clone()));
}
let client = reqwest::Client::builder()
.user_agent(format!("mermaid/{}", env!("CARGO_PKG_VERSION")))
.connect_timeout(Duration::from_secs(CONNECT_TIMEOUT_SECS))
.redirect(reqwest::redirect::Policy::none())
.dns_resolver(std::sync::Arc::new(McpVettingResolver {
allow_private: config.allow_private_network,
}))
.build()
.context("failed to build MCP HTTP client")?;
Ok(Self {
client,
url,
static_headers,
env_headers,
session_id: RwLock::new(None),
protocol_version: RwLock::new(None),
next_id: AtomicU64::new(1),
})
}
pub async fn send_request(&self, method: &str, params: Value) -> Result<Value> {
self.send_request_with_timeout(method, params, REQUEST_TIMEOUT_SECS)
.await
}
pub async fn send_request_with_timeout(
&self,
method: &str,
params: Value,
response_timeout_secs: u64,
) -> Result<Value> {
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
let request = json!({
"jsonrpc": "2.0",
"id": id,
"method": method,
"params": params,
});
timeout(
Duration::from_secs(response_timeout_secs),
self.request_roundtrip(method, id, &request),
)
.await
.map_err(|_| {
anyhow!(
"MCP request timed out after {}s: {}",
response_timeout_secs,
method
)
})?
}
pub async fn send_notification(&self, method: &str, params: Value) -> Result<()> {
let notification = json!({
"jsonrpc": "2.0",
"method": method,
"params": params,
});
self.post_message(¬ification)
.await
.with_context(|| format!("MCP notification failed (method: {method})"))
}
pub fn set_protocol_version(&self, version: &str) {
match HeaderValue::from_str(version) {
Ok(v) => {
*self
.protocol_version
.write()
.expect("mcp protocol_version lock poisoned") = Some(v);
},
Err(_) => tracing::warn!(
"MCP: server negotiated protocol version {:?} is not a valid header value; \
omitting MCP-Protocol-Version",
version
),
}
}
pub async fn shutdown(&self) {
if self
.session_id
.read()
.expect("mcp session_id lock poisoned")
.is_none()
{
return;
}
let Ok(headers) = self.request_headers(None) else {
return;
};
let result = timeout(
Duration::from_secs(DELETE_TIMEOUT_SECS),
self.client.delete(self.url.clone()).headers(headers).send(),
)
.await;
match result {
Ok(Ok(resp)) if resp.status() == StatusCode::METHOD_NOT_ALLOWED => {
tracing::debug!("MCP: server does not allow client session termination (405)");
},
Ok(Err(e)) => tracing::debug!("MCP: session DELETE failed: {}", e),
Err(_) => tracing::debug!("MCP: session DELETE timed out"),
Ok(Ok(_)) => {},
}
}
async fn request_roundtrip(&self, method: &str, id: u64, request: &Value) -> Result<Value> {
let response = self
.client
.post(self.url.clone())
.headers(self.request_headers(Some(ACCEPT_POST))?)
.json(request)
.send()
.await
.with_context(|| format!("MCP HTTP request failed (method: {method})"))?;
self.capture_session(response.headers());
let status = response.status();
if status == StatusCode::ACCEPTED {
bail!("MCP server returned 202 Accepted to a request (method: {method})");
}
if !status.is_success() {
return Err(self.status_error(method, status, response).await);
}
let content_type = response
.headers()
.get(reqwest::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.to_ascii_lowercase();
if content_type.starts_with("text/event-stream") {
let msg = self.drain_sse_for_response(response, method, id).await?;
extract_jsonrpc_result(msg)
} else if content_type.starts_with("application/json") {
let msg: Value = response
.json()
.await
.with_context(|| format!("MCP response was not valid JSON (method: {method})"))?;
if !is_response(&msg) || msg.get("id").and_then(parse_response_id) != Some(id) {
bail!("MCP JSON response did not answer request {id} (method: {method})");
}
extract_jsonrpc_result(msg)
} else {
bail!(
"MCP server returned unsupported Content-Type '{content_type}' \
(method: {method}; expected application/json or text/event-stream)"
);
}
}
async fn drain_sse_for_response(
&self,
response: reqwest::Response,
method: &str,
id: u64,
) -> Result<Value> {
let mut meta = SseMeta::default();
let mut response = response;
let mut reconnects = 0u32;
loop {
if let Some(msg) = self.drain_one_stream(response, id, &mut meta).await? {
return Ok(msg);
}
let Some(last_id) = meta.last_event_id.clone() else {
bail!("MCP SSE stream ended without a response (method: {method})");
};
reconnects += 1;
if reconnects > MAX_SSE_RECONNECTS {
bail!(
"MCP SSE stream did not deliver a response after {MAX_SSE_RECONNECTS} \
reconnects (method: {method})"
);
}
if let Some(ms) = meta.retry_ms {
tokio::time::sleep(Duration::from_millis(ms)).await;
}
let mut headers = self.request_headers(Some(ACCEPT_SSE))?;
headers.insert(
LAST_EVENT_ID_HEADER,
HeaderValue::from_str(&last_id)
.map_err(|_| anyhow!("MCP SSE event id is not a valid header value"))?,
);
let resumed = self
.client
.get(self.url.clone())
.headers(headers)
.send()
.await
.with_context(|| format!("MCP SSE resume failed (method: {method})"))?;
self.capture_session(resumed.headers());
if !resumed.status().is_success() {
let status = resumed.status();
return Err(self.status_error(method, status, resumed).await);
}
response = resumed;
}
}
async fn drain_one_stream(
&self,
response: reqwest::Response,
id: u64,
meta: &mut SseMeta,
) -> Result<Option<Value>> {
meta.start_stream();
let mut stream = response.bytes_stream();
let mut buf: Vec<u8> = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk = match chunk {
Ok(c) => c,
Err(e) => {
if meta.last_event_id.is_some() {
tracing::debug!("MCP: SSE stream broke, will resume: {}", e);
return Ok(None);
}
return Err(anyhow!(e)).context("MCP SSE stream read failed");
},
};
meta.feed(&chunk);
buf.extend_from_slice(&chunk);
for payload in drain_sse_events(&mut buf) {
let Ok(msg) = serde_json::from_str::<Value>(&payload) else {
tracing::warn!("MCP: unparseable SSE event payload");
continue;
};
if is_response(&msg) {
if msg.get("id").and_then(parse_response_id) == Some(id) {
return Ok(Some(msg));
}
continue;
}
if msg.get("method").is_some() {
match msg.get("id") {
Some(rid) if !rid.is_null() => {
let rid = rid.clone();
self.answer_server_request(&msg, &rid).await;
},
_ => tracing::trace!(
method = msg.get("method").and_then(|m| m.as_str()).unwrap_or(""),
"MCP: ignoring server notification on SSE stream"
),
}
}
}
}
Ok(None)
}
async fn answer_server_request(&self, msg: &Value, rid: &Value) {
let method = msg.get("method").and_then(|m| m.as_str()).unwrap_or("");
let reply = if method == "ping" {
json!({ "jsonrpc": "2.0", "id": rid, "result": {} })
} else {
json!({
"jsonrpc": "2.0",
"id": rid,
"error": { "code": -32601, "message": "method not supported" },
})
};
if let Err(e) = self.post_message(&reply).await {
tracing::debug!(
"MCP: failed to answer server-initiated request '{}': {}",
method,
e
);
}
}
async fn post_message(&self, message: &Value) -> Result<()> {
let fut = async {
let response = self
.client
.post(self.url.clone())
.headers(self.request_headers(Some(ACCEPT_POST))?)
.json(message)
.send()
.await
.context("MCP HTTP post failed")?;
self.capture_session(response.headers());
let status = response.status();
if !status.is_success() {
return Err(self.status_error("(notification)", status, response).await);
}
Ok(())
};
timeout(Duration::from_secs(NOTIFICATION_TIMEOUT_SECS), fut)
.await
.map_err(|_| anyhow!("MCP notification timed out after {NOTIFICATION_TIMEOUT_SECS}s"))?
}
fn request_headers(&self, accept: Option<&'static str>) -> Result<HeaderMap> {
let mut headers = self.static_headers.clone();
for (name, var) in &self.env_headers {
if let Ok(value) = std::env::var(var) {
let v = HeaderValue::from_str(&value).map_err(|_| {
anyhow!("invalid value in env var '{var}' for MCP header '{name}'")
})?;
headers.insert(name.clone(), v);
}
}
if let Some(accept) = accept {
headers.insert(ACCEPT, HeaderValue::from_static(accept));
}
if let Some(sid) = self
.session_id
.read()
.expect("mcp session_id lock poisoned")
.clone()
{
headers.insert(SESSION_HEADER, sid);
}
if let Some(pv) = self
.protocol_version
.read()
.expect("mcp protocol_version lock poisoned")
.clone()
{
headers.insert(PROTOCOL_VERSION_HEADER, pv);
}
Ok(headers)
}
fn capture_session(&self, headers: &HeaderMap) {
if let Some(v) = headers.get(&SESSION_HEADER) {
*self
.session_id
.write()
.expect("mcp session_id lock poisoned") = Some(v.clone());
}
}
async fn status_error(
&self,
method: &str,
status: StatusCode,
response: reqwest::Response,
) -> anyhow::Error {
if status == StatusCode::NOT_FOUND
&& self
.session_id
.read()
.expect("mcp session_id lock poisoned")
.is_some()
{
return anyhow!(
"MCP session expired (HTTP 404); a new session must be initialized — \
restart the server connection (method: {method})"
);
}
let body = response.text().await.unwrap_or_default();
let snippet = crate::utils::redact_secrets(&body);
let end = snippet.floor_char_boundary(ERROR_BODY_SNIPPET_BYTES.min(snippet.len()));
anyhow!(
"MCP server returned HTTP {} (method: {}): {}",
status,
method,
&snippet[..end]
)
}
}
#[derive(Default)]
struct SseMeta {
line_buf: Vec<u8>,
pending_event_id: Option<String>,
last_event_id: Option<String>,
retry_ms: Option<u64>,
}
impl SseMeta {
fn start_stream(&mut self) {
self.line_buf.clear();
self.pending_event_id = None;
}
fn feed(&mut self, chunk: &[u8]) {
self.line_buf.extend_from_slice(chunk);
while let Some(pos) = self.line_buf.iter().position(|&b| b == b'\n') {
let line: Vec<u8> = self.line_buf.drain(..=pos).collect();
let text = String::from_utf8_lossy(&line);
let text = text.trim_end_matches(['\n', '\r']);
if text.is_empty() {
if let Some(id) = self.pending_event_id.take() {
self.last_event_id = Some(id);
}
} else if let Some(v) = text.strip_prefix("id:") {
let v = v.strip_prefix(' ').unwrap_or(v);
self.pending_event_id = Some(v.to_string());
} else if let Some(v) = text.strip_prefix("retry:")
&& let Ok(ms) = v.trim().parse::<u64>()
{
self.retry_ms = Some(ms);
}
}
}
}
struct McpVettingResolver {
allow_private: bool,
}
impl reqwest::dns::Resolve for McpVettingResolver {
fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving {
let allow_private = self.allow_private;
Box::pin(async move {
let host = name.as_str().to_string();
let addrs: Vec<std::net::SocketAddr> =
tokio::net::lookup_host((host.as_str(), 0)).await?.collect();
for addr in &addrs {
let class = classify_host(&addr.ip().to_string());
let blocked = match class {
HostClass::Loopback | HostClass::Public => false,
_ => !allow_private,
};
if blocked {
return Err(format!(
"refusing to connect MCP server '{host}': it resolves to a \
private/internal address (set allow_private_network = true \
for this server to permit it)"
)
.into());
}
}
Ok(Box::new(addrs.into_iter()) as reqwest::dns::Addrs)
})
}
}
#[cfg(test)]
pub(super) mod test_fixture {
use super::*;
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::sync::Mutex;
pub enum Reply {
Raw(String),
Hang,
}
pub struct Fixture {
pub url: String,
requests: Arc<Mutex<Vec<String>>>,
}
impl Fixture {
pub async fn requests(&self) -> Vec<String> {
self.requests.lock().await.clone()
}
pub fn config(&self) -> McpServerConfig {
McpServerConfig {
url: Some(self.url.clone()),
..Default::default()
}
}
pub fn transport(&self) -> HttpTransport {
HttpTransport::new(&self.config()).expect("transport")
}
}
pub async fn fixture(replies: Vec<Reply>) -> Fixture {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind");
let addr = listener.local_addr().expect("addr");
let requests = Arc::new(Mutex::new(Vec::new()));
let recorded = Arc::clone(&requests);
tokio::spawn(async move {
for reply in replies {
let Ok((mut sock, _)) = listener.accept().await else {
return;
};
let raw = read_http_request(&mut sock).await;
recorded.lock().await.push(raw);
match reply {
Reply::Raw(bytes) => {
let _ = sock.write_all(bytes.as_bytes()).await;
let _ = sock.shutdown().await;
},
Reply::Hang => {
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
},
}
}
});
Fixture {
url: format!("http://127.0.0.1:{}/mcp", addr.port()),
requests,
}
}
async fn read_http_request(sock: &mut tokio::net::TcpStream) -> String {
let mut buf: Vec<u8> = Vec::new();
let mut tmp = [0u8; 4096];
loop {
if let Some(pos) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
let head = String::from_utf8_lossy(&buf[..pos]).to_ascii_lowercase();
let content_length = head
.lines()
.find_map(|l| l.strip_prefix("content-length:"))
.and_then(|v| v.trim().parse::<usize>().ok())
.unwrap_or(0);
let total = pos + 4 + content_length;
while buf.len() < total {
match sock.read(&mut tmp).await {
Ok(0) | Err(_) => break,
Ok(n) => buf.extend_from_slice(&tmp[..n]),
}
}
return String::from_utf8_lossy(&buf[..total.min(buf.len())]).into_owned();
}
match sock.read(&mut tmp).await {
Ok(0) | Err(_) => return String::from_utf8_lossy(&buf).into_owned(),
Ok(n) => buf.extend_from_slice(&tmp[..n]),
}
}
}
pub fn json_reply_with_headers(body: &str, extra_headers: &str) -> Reply {
Reply::Raw(format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n{extra_headers}Content-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
))
}
pub fn json_reply(body: &str) -> Reply {
json_reply_with_headers(body, "")
}
pub fn sse_reply(events: &str) -> Reply {
Reply::Raw(format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{events}",
events.len()
))
}
pub fn status_reply(code: u16, reason: &str) -> Reply {
Reply::Raw(format!(
"HTTP/1.1 {code} {reason}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
))
}
pub fn rpc_response(id: u64, result: &str) -> String {
format!(r#"{{"jsonrpc":"2.0","id":{id},"result":{result}}}"#)
}
}
#[cfg(test)]
mod tests {
use super::super::client::McpClient;
use super::test_fixture::*;
use super::*;
#[tokio::test]
async fn plain_json_response_round_trips() {
let fx = fixture(vec![json_reply(&rpc_response(1, r#"{"ok":true}"#))]).await;
let t = fx.transport();
let result = t.send_request("ping", json!({})).await.expect("result");
assert_eq!(result["ok"], true);
let reqs = fx.requests().await;
assert_eq!(reqs.len(), 1);
assert!(reqs[0].starts_with("POST /mcp"), "{}", reqs[0]);
assert!(
reqs[0]
.to_ascii_lowercase()
.contains("accept: application/json, text/event-stream"),
"{}",
reqs[0]
);
assert!(reqs[0].contains(r#""method":"ping""#), "{}", reqs[0]);
}
#[tokio::test]
async fn sse_response_drains_notifications_until_our_id() {
let events = format!(
"data: {}\n\ndata: {}\n\n",
r#"{"jsonrpc":"2.0","method":"notifications/message","params":{}}"#,
rpc_response(1, r#"{"via":"sse"}"#)
);
let fx = fixture(vec![sse_reply(&events)]).await;
let t = fx.transport();
let result = t
.send_request("tools/list", json!({}))
.await
.expect("result");
assert_eq!(result["via"], "sse");
}
#[tokio::test]
async fn session_id_is_captured_and_echoed_including_delete() {
let fx = fixture(vec![
json_reply_with_headers(&rpc_response(1, "{}"), "Mcp-Session-Id: sess-123\r\n"),
json_reply(&rpc_response(2, "{}")),
status_reply(200, "OK"), ])
.await;
let t = fx.transport();
t.send_request("initialize", json!({})).await.expect("init");
t.send_request("tools/list", json!({})).await.expect("list");
t.shutdown().await;
let reqs = fx.requests().await;
assert_eq!(reqs.len(), 3);
assert!(
!reqs[0].to_ascii_lowercase().contains("mcp-session-id"),
"no session id before the server assigns one: {}",
reqs[0]
);
assert!(
reqs[1]
.to_ascii_lowercase()
.contains("mcp-session-id: sess-123"),
"{}",
reqs[1]
);
assert!(reqs[2].starts_with("DELETE /mcp"), "{}", reqs[2]);
assert!(
reqs[2]
.to_ascii_lowercase()
.contains("mcp-session-id: sess-123"),
"{}",
reqs[2]
);
}
#[tokio::test]
async fn shutdown_without_session_sends_nothing() {
let fx = fixture(vec![]).await;
let t = fx.transport();
t.shutdown().await;
assert!(fx.requests().await.is_empty());
}
#[tokio::test]
async fn delete_405_is_tolerated() {
let fx = fixture(vec![
json_reply_with_headers(&rpc_response(1, "{}"), "MCP-Session-Id: s1\r\n"),
status_reply(405, "Method Not Allowed"),
])
.await;
let t = fx.transport();
t.send_request("initialize", json!({})).await.expect("init");
t.shutdown().await;
assert_eq!(fx.requests().await.len(), 2);
}
#[tokio::test]
async fn http_404_with_active_session_reports_expiry() {
let fx = fixture(vec![
json_reply_with_headers(&rpc_response(1, "{}"), "MCP-Session-Id: s1\r\n"),
status_reply(404, "Not Found"),
])
.await;
let t = fx.transport();
t.send_request("initialize", json!({})).await.expect("init");
let err = t
.send_request("tools/list", json!({}))
.await
.expect_err("expired");
assert!(err.to_string().contains("session expired"), "{err}");
}
#[tokio::test]
async fn static_and_env_headers_reach_the_wire() {
const VAR: &str = "MERMAID_TEST_MCP_HTTP_ENV_HEADER_A";
unsafe { std::env::set_var(VAR, "from-env") };
let fx = fixture(vec![json_reply(&rpc_response(1, "{}"))]).await;
let mut config = fx.config();
config
.headers
.insert("X-Static-Token".to_string(), "static-secret".to_string());
config
.env_headers
.insert("X-Env-Token".to_string(), VAR.to_string());
config.env_headers.insert(
"X-Missing".to_string(),
"MERMAID_TEST_MCP_HTTP_ENV_HEADER_MISSING".to_string(),
);
let t = HttpTransport::new(&config).expect("transport");
t.send_request("ping", json!({})).await.expect("result");
let req = fx.requests().await.remove(0).to_ascii_lowercase();
assert!(req.contains("x-static-token: static-secret"), "{req}");
assert!(req.contains("x-env-token: from-env"), "{req}");
assert!(!req.contains("x-missing"), "{req}");
}
#[tokio::test]
async fn invalid_header_name_errors_at_construction_without_value() {
let mut config = McpServerConfig {
url: Some("https://example.com/mcp".to_string()),
..Default::default()
};
config
.headers
.insert("bad header".to_string(), "secret-value".to_string());
let err = HttpTransport::new(&config)
.map(|_| ())
.expect_err("bad name");
assert!(err.to_string().contains("bad header"), "{err}");
assert!(!err.to_string().contains("secret-value"), "{err}");
}
#[tokio::test]
async fn notification_accepts_202_and_rejects_500() {
let fx = fixture(vec![
status_reply(202, "Accepted"),
Reply::Raw(
"HTTP/1.1 500 Internal Server Error\r\nContent-Type: text/plain\r\nContent-Length: 4\r\nConnection: close\r\n\r\noops"
.to_string(),
),
])
.await;
let t = fx.transport();
t.send_notification("notifications/initialized", json!({}))
.await
.expect("202 ok");
let err = t
.send_notification("notifications/initialized", json!({}))
.await
.expect_err("500 err");
let rendered = format!("{err:#}");
assert!(rendered.contains("500"), "{rendered}");
assert!(
rendered.contains("oops"),
"body snippet expected: {rendered}"
);
}
#[tokio::test]
async fn http_202_on_a_request_is_a_protocol_error() {
let fx = fixture(vec![status_reply(202, "Accepted")]).await;
let t = fx.transport();
let err = t
.send_request("tools/list", json!({}))
.await
.expect_err("202 on request");
assert!(err.to_string().contains("202"), "{err}");
}
#[tokio::test]
async fn sse_eof_without_response_and_without_ids_errors() {
let events =
"data: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/message\",\"params\":{}}\n\n";
let fx = fixture(vec![sse_reply(events)]).await;
let t = fx.transport();
let err = t
.send_request("tools/list", json!({}))
.await
.expect_err("eof");
assert!(
err.to_string().contains("ended without a response"),
"{err}"
);
}
#[tokio::test]
async fn sse_disconnect_with_event_id_resumes_via_get_last_event_id() {
let priming = "id: ev-7\nretry: 10\ndata:\n\n";
let resumed = format!(
"id: ev-8\ndata: {}\n\n",
rpc_response(1, r#"{"resumed":true}"#)
);
let fx = fixture(vec![sse_reply(priming), sse_reply(&resumed)]).await;
let t = fx.transport();
let result = t
.send_request("tools/call", json!({}))
.await
.expect("resumed");
assert_eq!(result["resumed"], true);
let reqs = fx.requests().await;
assert_eq!(reqs.len(), 2);
assert!(reqs[1].starts_with("GET /mcp"), "{}", reqs[1]);
let get = reqs[1].to_ascii_lowercase();
assert!(get.contains("last-event-id: ev-7"), "{}", reqs[1]);
assert!(get.contains("accept: text/event-stream"), "{}", reqs[1]);
}
#[tokio::test]
async fn server_request_on_sse_stream_gets_error_reply_posted_back() {
let events = format!(
"data: {}\n\ndata: {}\n\n",
r#"{"jsonrpc":"2.0","id":99,"method":"sampling/createMessage","params":{}}"#,
rpc_response(1, "{}")
);
let fx = fixture(vec![sse_reply(&events), status_reply(202, "Accepted")]).await;
let t = fx.transport();
t.send_request("tools/list", json!({}))
.await
.expect("result");
let reqs = fx.requests().await;
assert_eq!(reqs.len(), 2, "the -32601 reply must be POSTed back");
assert!(reqs[1].contains("-32601"), "{}", reqs[1]);
assert!(reqs[1].contains(r#""id":99"#), "{}", reqs[1]);
}
#[tokio::test]
async fn server_ping_on_sse_stream_gets_result_not_32601() {
let events = format!(
"data: {}\n\ndata: {}\n\n",
r#"{"jsonrpc":"2.0","id":42,"method":"ping"}"#,
rpc_response(1, "{}")
);
let fx = fixture(vec![sse_reply(&events), status_reply(202, "Accepted")]).await;
let t = fx.transport();
t.send_request("tools/list", json!({}))
.await
.expect("result");
let reqs = fx.requests().await;
assert_eq!(reqs.len(), 2);
assert!(reqs[1].contains(r#""result":{}"#), "{}", reqs[1]);
assert!(!reqs[1].contains("-32601"), "{}", reqs[1]);
}
#[tokio::test]
async fn injected_short_timeout_is_enforced() {
let fx = fixture(vec![Reply::Hang]).await;
let t = fx.transport();
let start = std::time::Instant::now();
let err = t
.send_request_with_timeout("tools/list", json!({}), 1)
.await
.expect_err("must time out");
assert!(err.to_string().contains("timed out"), "{err}");
assert!(start.elapsed() < std::time::Duration::from_secs(5));
}
#[tokio::test]
async fn protocol_version_header_sent_after_initialize() {
let init_result = r#"{"protocolVersion":"2025-11-25","capabilities":{},"serverInfo":{"name":"fx","version":"1.0"}}"#;
let fx = fixture(vec![
json_reply_with_headers(&rpc_response(1, init_result), "MCP-Session-Id: s9\r\n"),
status_reply(202, "Accepted"), json_reply(&rpc_response(2, r#"{"tools":[]}"#)),
])
.await;
let mut client = McpClient::new(fx.transport().into());
client.initialize().await.expect("initialize");
client.list_tools().await.expect("list");
let reqs = fx.requests().await;
assert_eq!(reqs.len(), 3);
assert!(
!reqs[0]
.to_ascii_lowercase()
.contains("mcp-protocol-version"),
"initialize itself carries no version header: {}",
reqs[0]
);
for req in &reqs[1..] {
let low = req.to_ascii_lowercase();
assert!(low.contains("mcp-protocol-version: 2025-11-25"), "{req}");
assert!(low.contains("mcp-session-id: s9"), "{req}");
}
}
#[tokio::test]
async fn resolver_allows_loopback_blocks_private_honors_flag() {
use reqwest::dns::Resolve;
use std::str::FromStr;
let name = |s: &str| reqwest::dns::Name::from_str(s).expect("name");
let default = McpVettingResolver {
allow_private: false,
};
assert!(default.resolve(name("localhost")).await.is_ok());
assert!(default.resolve(name("192.168.1.5")).await.is_err());
let opted_in = McpVettingResolver {
allow_private: true,
};
assert!(opted_in.resolve(name("192.168.1.5")).await.is_ok());
}
#[test]
fn sse_meta_tracks_id_and_retry_across_chunk_splits() {
let mut meta = SseMeta::default();
meta.feed(b"id: ev");
assert_eq!(meta.last_event_id, None);
meta.feed(b"-1\nretry: 250\ndata: x\n\n");
assert_eq!(meta.last_event_id.as_deref(), Some("ev-1"));
assert_eq!(meta.retry_ms, Some(250));
meta.feed(b"id: ev-2\r\ndata: y\r\n\r\n");
assert_eq!(meta.last_event_id.as_deref(), Some("ev-2"));
}
#[test]
fn sse_meta_commits_id_only_at_event_boundary() {
let mut meta = SseMeta::default();
meta.feed(b"id: ev-1\ndata: x\n\n");
assert_eq!(meta.last_event_id.as_deref(), Some("ev-1"));
meta.feed(b"id: ev-2\ndata: {\"jsonr");
assert_eq!(meta.last_event_id.as_deref(), Some("ev-1"));
meta.start_stream();
meta.feed(b"id: ev-2\ndata: y\n\n");
assert_eq!(meta.last_event_id.as_deref(), Some("ev-2"));
}
#[test]
fn new_vets_ip_literal_hosts() {
let private = McpServerConfig {
url: Some("https://192.168.1.5/mcp".to_string()),
..Default::default()
};
let err = match HttpTransport::new(&private) {
Ok(_) => panic!("private literal must be rejected"),
Err(e) => e,
};
assert!(err.to_string().contains("allow_private_network"), "{err}");
let metadata = McpServerConfig {
url: Some("https://[::ffff:169.254.169.254]/mcp".to_string()),
..Default::default()
};
assert!(HttpTransport::new(&metadata).is_err());
let opted_in = McpServerConfig {
url: Some("https://192.168.1.5/mcp".to_string()),
allow_private_network: true,
..Default::default()
};
assert!(HttpTransport::new(&opted_in).is_ok());
let loopback = McpServerConfig {
url: Some("http://127.0.0.1:9099/mcp".to_string()),
..Default::default()
};
assert!(HttpTransport::new(&loopback).is_ok());
}
#[test]
fn new_rejects_config_without_url() {
let config = McpServerConfig {
command: "npx".to_string(),
..Default::default()
};
assert!(HttpTransport::new(&config).is_err());
}
}