use std::collections::HashMap;
use std::time::Duration;
use async_trait::async_trait;
use futures::StreamExt;
use reqwest::header::{ACCEPT, CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue};
use reqwest::{StatusCode, Url};
use serde_json::Value;
use tokio::sync::mpsc;
use std::sync::Arc;
use super::jsonrpc::{self, Inbound, JsonRpcRequest, JsonRpcResponse};
use super::{BearerRefresher, Transport};
const AUTHORIZATION: &str = "authorization";
const SESSION_HEADER: &str = "mcp-session-id";
const PROTOCOL_HEADER: &str = "mcp-protocol-version";
const READ_STALL_TIMEOUT_SECS: u64 = 900;
fn build_http_client() -> reqwest::Client {
reqwest::Client::builder()
.connect_timeout(Duration::from_secs(30))
.read_timeout(Duration::from_secs(READ_STALL_TIMEOUT_SECS))
.tcp_keepalive(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::custom(|attempt| {
let same_origin = attempt.previous().last().is_some_and(|prev| {
prev.scheme() == attempt.url().scheme()
&& prev.host_str() == attempt.url().host_str()
&& prev.port_or_known_default() == attempt.url().port_or_known_default()
});
match same_origin && attempt.previous().len() <= 5 {
true => attempt.follow(),
false => attempt.stop(),
}
}))
.build()
.expect("failed to build reqwest client")
}
#[cfg(test)]
pub(crate) fn expand_env(value: &str) -> String {
expand_env_allowing(value, &[])
}
fn resolve_var(name: &str, allowlist: &[String]) -> Option<String> {
if !leviath_core::script_env_allowed(name, allowlist) {
return None;
}
std::env::var(name).ok()
}
pub(crate) fn expand_env_allowing(value: &str, allowlist: &[String]) -> String {
let mut out = String::with_capacity(value.len());
let mut rest = value;
while let Some((before, after)) = rest.split_once("${") {
out.push_str(before);
match after.split_once('}') {
Some((name, tail)) => {
match resolve_var(name, allowlist) {
Some(v) => out.push_str(&v),
None => tracing::warn!(
var = %name,
"MCP header references an environment variable that is unset \
or refused; add it to `[security] allow_env_vars` if the \
server genuinely needs it"
),
}
rest = tail;
}
None => {
out.push_str("${");
out.push_str(after);
return out;
}
}
}
out.push_str(rest);
out
}
fn build_headers(configured: &HashMap<String, String>, allow_env: &[String]) -> HeaderMap {
let mut headers = HeaderMap::new();
for (name, value) in configured {
let expanded = expand_env_allowing(value, allow_env);
match (
HeaderName::from_bytes(name.as_bytes()),
HeaderValue::from_str(&expanded),
) {
(Ok(n), Ok(v)) => {
headers.insert(n, v);
}
_ => tracing::warn!(header = %name, "Skipping invalid MCP header"),
}
}
headers
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Mode {
Streamable,
Legacy,
}
enum LegacyEvent {
Endpoint(String),
Frame(Value),
}
struct LegacyStream {
frames: mpsc::UnboundedReceiver<LegacyEvent>,
reader: tokio::task::JoinHandle<()>,
post_url: Url,
}
pub(crate) struct HttpTransport {
client: reqwest::Client,
url: Url,
headers: HeaderMap,
mode: Mode,
session_id: Option<String>,
protocol_version: Option<String>,
legacy: Option<LegacyStream>,
refresher: Option<Arc<dyn BearerRefresher>>,
}
impl HttpTransport {
pub(crate) fn new(
url: &str,
headers: &HashMap<String, String>,
allow_env: &[String],
) -> anyhow::Result<Self> {
let url = Url::parse(url)
.map_err(|e| anyhow::anyhow!("Invalid MCP server url '{}': {}", url, e))?;
if !crate::auth::metadata::is_safe_discovery_url(&url) {
tracing::warn!(
url = %url,
"MCP server is not HTTPS, so its credentials travel in cleartext"
);
}
Ok(Self {
client: build_http_client(),
url,
headers: build_headers(headers, allow_env),
mode: Mode::Streamable,
session_id: None,
protocol_version: None,
legacy: None,
refresher: None,
})
}
fn set_auth_header(&mut self, value: &str) {
match HeaderValue::from_str(value) {
Ok(v) => {
self.headers
.insert(HeaderName::from_static(AUTHORIZATION), v);
}
Err(e) => tracing::warn!(error = %e, "Refreshed bearer is not a valid header value"),
}
}
async fn post_maybe_refresh(&mut self, body: &str) -> anyhow::Result<reqwest::Response> {
let response = self.post(body).await?;
if response.status() != StatusCode::UNAUTHORIZED {
return Ok(response);
}
let Some(refresher) = self.refresher.clone() else {
return Ok(response);
};
tracing::info!("MCP request returned 401 - refreshing the token and retrying");
let value = refresher.refresh().await?;
self.set_auth_header(&value);
self.post(body).await
}
fn post_url(&self) -> &Url {
match &self.legacy {
Some(stream) => &stream.post_url,
None => &self.url,
}
}
fn request_headers(&self) -> HeaderMap {
let mut headers = self.headers.clone();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert(
ACCEPT,
HeaderValue::from_static("application/json, text/event-stream"),
);
if let Some(session) = &self.session_id
&& let Ok(value) = HeaderValue::from_str(session)
{
headers.insert(HeaderName::from_static(SESSION_HEADER), value);
}
if let Some(version) = &self.protocol_version
&& let Ok(value) = HeaderValue::from_str(version)
{
headers.insert(HeaderName::from_static(PROTOCOL_HEADER), value);
}
headers
}
fn learn_session(&mut self, response_headers: &HeaderMap, body: &JsonRpcResponse) {
if let Some(session) = response_headers
.get(SESSION_HEADER)
.and_then(|v| v.to_str().ok())
{
self.session_id = Some(session.to_string());
}
if let Some(version) = body
.result
.as_ref()
.and_then(|r| r.get("protocolVersion"))
.and_then(Value::as_str)
{
self.protocol_version = Some(version.to_string());
}
}
async fn post_expecting_success(&self, body: &str) -> anyhow::Result<()> {
let response = self.post(body).await?;
let status = response.status();
if status.is_success() {
return Ok(());
}
Err(error_for_status(status, response).await)
}
async fn post(&self, body: &str) -> anyhow::Result<reqwest::Response> {
self.client
.post(self.post_url().clone())
.headers(self.request_headers())
.body(body.to_string())
.send()
.await
.map_err(|e| anyhow::anyhow!("MCP HTTP request failed: {}", e))
}
async fn read_json_reply(response: reqwest::Response) -> anyhow::Result<JsonRpcResponse> {
let body = response
.text()
.await
.map_err(|e| anyhow::anyhow!("Failed to read MCP response body: {}", e))?;
let frame: Value = serde_json::from_str(&body)
.map_err(|e| anyhow::anyhow!("Failed to parse JSON-RPC response: {}", e))?;
match jsonrpc::classify(frame)? {
Inbound::Response(response) => Ok(*response),
_ => Err(anyhow::anyhow!(
"MCP server answered a request with a non-response frame"
)),
}
}
async fn read_sse_reply(
&self,
response: reqwest::Response,
id: Option<u64>,
) -> anyhow::Result<JsonRpcResponse> {
let mut buffer = String::new();
let mut stream = response.bytes_stream();
loop {
while let Some(event) = super::sse::parse_sse_frame(&mut buffer) {
if event.data.is_empty() {
continue;
}
let frame: Value = serde_json::from_str(&event.data)
.map_err(|e| anyhow::anyhow!("Failed to parse JSON-RPC response: {}", e))?;
if let Some(response) = self.handle_frame(frame, id).await? {
return Ok(response);
}
}
match stream.next().await {
Some(Ok(chunk)) => match std::str::from_utf8(&chunk) {
Ok(text) => buffer.push_str(text),
Err(e) => {
return Err(anyhow::anyhow!("MCP event stream is not UTF-8: {}", e));
}
},
Some(Err(e)) => {
return Err(anyhow::anyhow!("MCP event stream failed: {}", e));
}
None => {
return Err(anyhow::anyhow!(
"MCP event stream ended before answering the request"
));
}
}
}
}
async fn handle_frame(
&self,
frame: Value,
id: Option<u64>,
) -> anyhow::Result<Option<JsonRpcResponse>> {
match jsonrpc::classify(frame)? {
Inbound::Response(response) => {
if response_matches(&response, id) {
Ok(Some(*response))
} else {
tracing::debug!("Ignoring MCP response with a non-matching id");
Ok(None)
}
}
Inbound::ServerRequest {
id: request_id,
method,
} => {
tracing::debug!(method = %method, "Answering server-initiated request");
let reply = jsonrpc::reply_to_server_request(&request_id, &method);
if let Err(e) = self.post_expecting_success(&reply.to_string()).await {
tracing::warn!(error = %e, "Could not answer server-initiated request");
}
Ok(None)
}
Inbound::Notification { method } => {
tracing::debug!(method = %method, "Ignoring server notification");
Ok(None)
}
}
}
async fn start_legacy(&mut self) -> anyhow::Result<()> {
tracing::info!(url = %self.url, "Falling back to the legacy MCP HTTP+SSE transport");
let mut headers = self.headers.clone();
headers.insert(ACCEPT, HeaderValue::from_static("text/event-stream"));
let response = self
.client
.get(self.url.clone())
.headers(headers)
.send()
.await
.map_err(|e| anyhow::anyhow!("MCP event stream request failed: {}", e))?;
let status = response.status();
if !status.is_success() {
return Err(anyhow::anyhow!(
"MCP server rejected the event stream with HTTP {}",
status
));
}
let (tx, mut rx) = mpsc::unbounded_channel();
let base = self.url.clone();
let reader = tokio::spawn(read_event_stream(response, tx));
let endpoint = loop {
match rx.recv().await {
Some(LegacyEvent::Endpoint(path)) => break path,
Some(LegacyEvent::Frame(_)) => {
tracing::debug!("Discarding an MCP frame sent before the endpoint event");
}
None => {
reader.abort();
return Err(anyhow::anyhow!(
"MCP event stream closed before naming a POST endpoint"
));
}
}
};
let post_url = base
.join(&endpoint)
.map_err(|e| anyhow::anyhow!("Invalid MCP endpoint '{}': {}", endpoint, e))?;
if !crate::auth::metadata::same_origin(&post_url, &base) {
reader.abort();
return Err(anyhow::anyhow!(
"MCP server named a POST endpoint at origin '{}', which is not its own '{}' - \
refusing, because the request would carry this server's credentials",
post_url.origin().ascii_serialization(),
base.origin().ascii_serialization(),
));
}
tracing::debug!(post_url = %post_url, "Legacy MCP endpoint established");
self.mode = Mode::Legacy;
self.legacy = Some(LegacyStream {
frames: rx,
reader,
post_url,
});
Ok(())
}
async fn legacy_request(
&mut self,
body: &str,
id: Option<u64>,
) -> anyhow::Result<JsonRpcResponse> {
let response = self.post_maybe_refresh(body).await?;
let status = response.status();
if !status.is_success() {
return Err(error_for_status(status, response).await);
}
loop {
let event = {
let stream = self
.legacy
.as_mut()
.expect("legacy mode always has a stream");
stream.frames.recv().await
};
let frame = match event {
Some(LegacyEvent::Frame(frame)) => frame,
Some(LegacyEvent::Endpoint(_)) => continue,
None => {
return Err(anyhow::anyhow!(
"MCP event stream ended before answering the request"
));
}
};
if let Some(response) = self.handle_frame(frame, id).await? {
return Ok(response);
}
}
}
async fn streamable_request(
&mut self,
body: &str,
id: Option<u64>,
) -> anyhow::Result<JsonRpcResponse> {
let response = self.post_maybe_refresh(body).await?;
let status = response.status();
if status == StatusCode::NOT_FOUND || status == StatusCode::METHOD_NOT_ALLOWED {
self.start_legacy().await?;
return self.legacy_request(body, id).await;
}
if !status.is_success() {
return Err(error_for_status(status, response).await);
}
let response_headers = response.headers().clone();
let content_type = response
.headers()
.get(CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.to_string();
let mut parsed = if content_type.starts_with("text/event-stream") {
self.read_sse_reply(response, id).await?
} else if content_type.starts_with("application/json") {
Self::read_json_reply(response).await?
} else {
return Err(anyhow::anyhow!(
"MCP server replied with unsupported content type '{}'",
content_type
));
};
self.learn_session(&response_headers, &parsed);
parsed.id.take();
Ok(parsed)
}
}
fn response_matches(response: &JsonRpcResponse, id: Option<u64>) -> bool {
match (&response.id, id) {
(Some(Value::Number(n)), Some(expected)) => n.as_u64() == Some(expected),
(Some(Value::Null) | None, _) => true,
_ => false,
}
}
async fn error_for_status(status: StatusCode, response: reqwest::Response) -> anyhow::Error {
let body = response.text().await.unwrap_or_default();
if body.is_empty() {
anyhow::anyhow!("MCP server returned HTTP {}", status)
} else {
anyhow::anyhow!("MCP server returned HTTP {}: {}", status, body.trim())
}
}
async fn read_event_stream(response: reqwest::Response, tx: mpsc::UnboundedSender<LegacyEvent>) {
let mut buffer = String::new();
let mut stream = response.bytes_stream();
loop {
while let Some(event) = super::sse::parse_sse_frame(&mut buffer) {
if event.data.is_empty() {
continue;
}
let decoded = if event.event.as_deref() == Some("endpoint") {
LegacyEvent::Endpoint(event.data.clone())
} else {
match serde_json::from_str(&event.data) {
Ok(frame) => LegacyEvent::Frame(frame),
Err(e) => {
tracing::warn!(error = %e, "Discarding unparseable MCP event");
continue;
}
}
};
if tx.send(decoded).is_err() {
return;
}
}
match stream.next().await {
Some(Ok(chunk)) => match std::str::from_utf8(&chunk) {
Ok(text) => buffer.push_str(text),
Err(e) => {
tracing::warn!(error = %e, "MCP event stream is not UTF-8");
return;
}
},
Some(Err(e)) => {
tracing::warn!(error = %e, "MCP event stream failed");
return;
}
None => return,
}
}
}
#[async_trait]
impl Transport for HttpTransport {
async fn send_request(
&mut self,
req: &JsonRpcRequest,
timeout: Duration,
) -> anyhow::Result<JsonRpcResponse> {
tracing::trace!(method = %req.method, "Sending JSON-RPC request over HTTP");
let body = serde_json::to_string(req).expect("JsonRpcRequest is always serializable");
let id = req.id;
let work = async {
match self.mode {
Mode::Streamable => self.streamable_request(&body, id).await,
Mode::Legacy => self.legacy_request(&body, id).await,
}
};
match tokio::time::timeout(timeout, work).await {
Ok(result) => result,
Err(_) => Err(anyhow::anyhow!(
"MCP server did not respond to '{}' within {}s",
req.method,
timeout.as_secs()
)),
}
}
async fn send_notification(&mut self, req: &JsonRpcRequest) -> anyhow::Result<()> {
tracing::trace!(method = %req.method, "Sending JSON-RPC notification over HTTP");
let body = serde_json::to_string(req).expect("JsonRpcRequest is always serializable");
self.post_expecting_success(&body).await
}
async fn close(&mut self) -> anyhow::Result<()> {
if let Some(stream) = self.legacy.take() {
stream.reader.abort();
}
if let Some(session) = self.session_id.take() {
let mut headers = self.headers.clone();
if let Ok(value) = HeaderValue::from_str(&session) {
headers.insert(HeaderName::from_static(SESSION_HEADER), value);
}
let _ = self
.client
.delete(self.url.clone())
.headers(headers)
.send()
.await;
}
Ok(())
}
fn set_bearer_refresher(&mut self, refresher: Arc<dyn BearerRefresher>) {
self.refresher = Some(refresher);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::always_on_tracing_guard;
use crate::transport::DEFAULT_REQUEST_TIMEOUT;
use axum::Router;
use axum::extract::State;
use axum::http::{HeaderMap as AxumHeaders, StatusCode as AxumStatus};
use axum::response::IntoResponse;
use axum::routing::{get, post};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn expand_env_substitutes_an_ordinary_variable() {
temp_env::with_var("LEV_MCP_TEST_REGION", Some("eu-west-1"), || {
assert_eq!(
expand_env("region=${LEV_MCP_TEST_REGION}"),
"region=eu-west-1"
);
});
}
#[test]
fn expand_env_refuses_a_credential_unless_allowlisted() {
let _guard = always_on_tracing_guard();
temp_env::with_var("LEV_MCP_TEST_TOKEN", Some("s3cret"), || {
assert_eq!(
expand_env("Bearer ${LEV_MCP_TEST_TOKEN}"),
"Bearer ",
"a credential-shaped name is not interpolated by default"
);
assert_eq!(
expand_env_allowing(
"Bearer ${LEV_MCP_TEST_TOKEN}",
&["LEV_MCP_TEST_TOKEN".to_string()]
),
"Bearer s3cret"
);
});
}
#[test]
fn expand_env_drops_an_undefined_variable() {
let _guard = always_on_tracing_guard();
temp_env::with_var_unset("LEV_MCP_TEST_MISSING", || {
assert_eq!(expand_env("Bearer ${LEV_MCP_TEST_MISSING}"), "Bearer ");
});
}
#[test]
fn expand_env_leaves_plain_values_alone() {
assert_eq!(expand_env("Bearer static-token"), "Bearer static-token");
}
#[test]
fn expand_env_handles_several_references() {
temp_env::with_vars(
[
("LEV_MCP_TEST_A", Some("one")),
("LEV_MCP_TEST_B", Some("two")),
],
|| {
assert_eq!(
expand_env("${LEV_MCP_TEST_A}-${LEV_MCP_TEST_B}!"),
"one-two!"
);
},
);
}
#[test]
fn expand_env_keeps_an_unterminated_reference_literal() {
assert_eq!(expand_env("Bearer ${UNCLOSED"), "Bearer ${UNCLOSED");
}
#[test]
fn expand_env_of_an_empty_value_is_empty() {
assert_eq!(expand_env(""), "");
}
#[test]
fn build_headers_expands_values() {
temp_env::with_var("LEV_MCP_TEST_TOKEN", Some("abc"), || {
let configured = HashMap::from([(
"Authorization".to_string(),
"Bearer ${LEV_MCP_TEST_TOKEN}".to_string(),
)]);
let allowed = ["LEV_MCP_TEST_TOKEN".to_string()];
let headers = build_headers(&configured, &allowed);
assert_eq!(headers.get("authorization").unwrap(), "Bearer abc");
let headers = build_headers(&configured, &[]);
assert_eq!(headers.get("authorization").unwrap(), "Bearer ");
});
}
#[test]
fn build_headers_skips_an_unrepresentable_entry() {
let _guard = always_on_tracing_guard();
let configured = HashMap::from([
("Not A Header".to_string(), "x".to_string()),
("X-Good".to_string(), "y".to_string()),
]);
let headers = build_headers(&configured, &[]);
assert_eq!(headers.len(), 1);
assert_eq!(headers.get("x-good").unwrap(), "y");
}
#[test]
fn build_headers_of_nothing_is_empty() {
assert!(build_headers(&HashMap::new(), &[]).is_empty());
}
async fn serve(app: Router) -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
listener, app,
)));
format!("http://{addr}")
}
fn transport(url: &str) -> HttpTransport {
HttpTransport::new(url, &HashMap::new(), &[]).expect("url should parse")
}
fn init() -> JsonRpcRequest {
JsonRpcRequest::request(1, "initialize", serde_json::json!({}))
}
fn ok_frame(id: u64) -> String {
serde_json::json!({"jsonrpc": "2.0", "id": id, "result": {"ok": true}}).to_string()
}
#[tokio::test]
async fn streamable_json_reply_roundtrips() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async { ([(CONTENT_TYPE, "application/json")], ok_frame(1)).into_response() }),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let value = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("request should succeed")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"ok": true}));
}
#[tokio::test]
async fn streamable_sse_reply_roundtrips() {
let _guard = always_on_tracing_guard();
let body = format!("event: message\ndata: {}\n\n", ok_frame(1));
let app = Router::new().route(
"/mcp",
post(move || {
let body = body.clone();
async move { ([(CONTENT_TYPE, "text/event-stream")], body).into_response() }
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let value = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("request should succeed")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"ok": true}));
}
#[tokio::test]
async fn streamable_sse_skips_interleaved_and_stale_frames() {
let _guard = always_on_tracing_guard();
let body = format!(
": keepalive\n\ndata: {}\n\ndata: {}\n\ndata: {}\n\n",
serde_json::json!({"jsonrpc": "2.0", "method": "notifications/progress"}),
serde_json::json!({"jsonrpc": "2.0", "id": 999, "result": {"stale": true}}),
ok_frame(1),
);
let app = Router::new().route(
"/mcp",
post(move || {
let body = body.clone();
async move { ([(CONTENT_TYPE, "text/event-stream")], body).into_response() }
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let value = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("request should succeed")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"ok": true}));
}
#[tokio::test]
async fn streamable_answers_a_server_request_and_keeps_waiting() {
let _guard = always_on_tracing_guard();
let replies = Arc::new(AtomicUsize::new(0));
let body = format!(
"data: {}\n\ndata: {}\n\n",
serde_json::json!({"jsonrpc": "2.0", "id": 7, "method": "ping"}),
ok_frame(1),
);
let app = Router::new().route(
"/mcp",
post({
let replies = replies.clone();
move |body_in: String| {
let (body, replies) = (body.clone(), replies.clone());
async move {
if body_in.contains("\"id\":7") {
replies.fetch_add(1, Ordering::SeqCst);
return AxumStatus::ACCEPTED.into_response();
}
([(CONTENT_TYPE, "text/event-stream")], body).into_response()
}
}
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_ok()
);
assert_eq!(replies.load(Ordering::SeqCst), 1, "ping must be answered");
}
#[tokio::test]
async fn session_id_and_protocol_version_are_echoed_on_later_requests() {
let _guard = always_on_tracing_guard();
let seen = Arc::new(std::sync::Mutex::new(
Vec::<(Option<String>, Option<String>)>::new(),
));
let app = Router::new().route(
"/mcp",
post({
let seen = seen.clone();
move |headers: AxumHeaders| {
let seen = seen.clone();
async move {
let get = |k: &str| {
headers
.get(k)
.and_then(|v| v.to_str().ok())
.map(str::to_string)
};
seen.lock()
.unwrap()
.push((get(SESSION_HEADER), get(PROTOCOL_HEADER)));
(
[
(CONTENT_TYPE, "application/json"),
(HeaderName::from_static(SESSION_HEADER), "sess-42"),
],
serde_json::json!({
"jsonrpc": "2.0", "id": 1,
"result": {"protocolVersion": "2025-06-18"}
})
.to_string(),
)
.into_response()
}
}
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.unwrap();
t.send_request(
&JsonRpcRequest::request(2, "tools/list", serde_json::json!({})),
DEFAULT_REQUEST_TIMEOUT,
)
.await
.unwrap();
let seen = seen.lock().unwrap();
assert_eq!(seen[0], (None, None), "nothing known before initialize");
assert_eq!(
seen[1],
(Some("sess-42".to_string()), Some("2025-06-18".to_string())),
"both must be echoed once learned"
);
}
#[tokio::test]
async fn notification_accepts_a_202() {
let _guard = always_on_tracing_guard();
let app = Router::new().route("/mcp", post(|| async { AxumStatus::ACCEPTED }));
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
t.send_notification(&JsonRpcRequest::notification(
"notifications/initialized",
serde_json::json!({}),
))
.await
.expect("202 is success for a notification");
}
#[tokio::test]
async fn notification_surfaces_a_server_error() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async { (AxumStatus::INTERNAL_SERVER_ERROR, "kaboom") }),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_notification(&JsonRpcRequest::notification("x", serde_json::json!({})))
.await
.expect_err("500 must fail");
assert!(err.to_string().contains("kaboom"), "got: {err}");
}
fn legacy_app(
endpoint_event: impl Into<String>,
) -> (Router, Arc<tokio::sync::broadcast::Sender<String>>) {
let (tx, _) = tokio::sync::broadcast::channel::<String>(16);
let tx = Arc::new(tx);
let endpoint_event = Arc::new(endpoint_event.into());
let app = Router::new()
.route(
"/sse",
get({
let tx = tx.clone();
move || {
let tx = tx.clone();
let endpoint_event = endpoint_event.clone();
async move {
let mut rx = tx.subscribe();
let stream = async_stream::stream! {
yield Ok::<_, std::io::Error>(
format!("event: endpoint\ndata: {endpoint_event}\n\n"));
while let Ok(frame) = rx.recv().await {
yield Ok(format!("data: {frame}\n\n"));
}
};
(
[(CONTENT_TYPE, "text/event-stream")],
axum::body::Body::from_stream(stream),
)
}
}
})
.post(|| async { AxumStatus::METHOD_NOT_ALLOWED }),
)
.route(
"/messages",
post({
let tx = tx.clone();
move |State(_): State<()>, body: String| {
let tx = tx.clone();
async move {
let req: Value = serde_json::from_str(&body).unwrap();
if let Some(id) = req.get("id") {
let _ = tx.send(
serde_json::json!({
"jsonrpc": "2.0", "id": id, "result": {"legacy": true}
})
.to_string(),
);
}
AxumStatus::ACCEPTED
}
}
}),
)
.with_state(());
(app, tx)
}
#[tokio::test]
async fn falls_back_to_legacy_on_405_and_completes_the_request() {
let _guard = always_on_tracing_guard();
let (app, _tx) = legacy_app("/messages");
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let value = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("legacy fallback should complete the request")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"legacy": true}));
assert_eq!(t.mode, Mode::Legacy);
}
#[tokio::test]
async fn legacy_endpoint_event_may_be_an_absolute_url() {
let _guard = always_on_tracing_guard();
let (app, _tx) = legacy_app("/messages");
let base = serve(app).await;
let url = format!("{base}/sse");
let mut t = transport(&url);
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("relative endpoint should resolve");
let post_url = t.post_url().to_string();
assert!(post_url.ends_with("/messages"), "got: {post_url}");
}
#[tokio::test]
async fn a_cross_origin_endpoint_event_is_refused_before_anything_is_sent() {
let _guard = always_on_tracing_guard();
let hits = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let thief = Router::new().fallback({
let hits = hits.clone();
move || {
hits.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
async { AxumStatus::ACCEPTED }
}
});
let thief_base = serve(thief).await;
let (app, _tx) = legacy_app(format!("{thief_base}/steal"));
let url = format!("{}/sse", serve(app).await);
let mut headers = HashMap::new();
headers.insert("x-api-key".to_string(), "super-secret".to_string());
let mut t = HttpTransport::new(&url, &headers, &[]).expect("url should parse");
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("a cross-origin endpoint must be refused");
assert_eq!(
hits.load(std::sync::atomic::Ordering::SeqCst),
0,
"no request carrying this server's credentials may reach another origin"
);
assert!(
err.to_string().contains("not its own"),
"the error should name both origins: {err}"
);
reqwest::Client::new()
.post(format!("{thief_base}/steal"))
.send()
.await
.expect("the thief is listening");
assert_eq!(hits.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_same_origin_redirect_is_followed() {
let app = Router::new()
.route(
"/first",
get(|| async {
(
AxumStatus::TEMPORARY_REDIRECT,
[(axum::http::header::LOCATION, "/second")],
)
.into_response()
}),
)
.route("/second", get(|| async { "ok" }));
let base = serve(app).await;
let response = build_http_client()
.get(format!("{base}/first"))
.send()
.await
.expect("a same-origin redirect should be followed");
assert_eq!(response.status(), 200);
}
#[tokio::test]
async fn a_cross_origin_redirect_does_not_carry_the_headers() {
let hits = Arc::new(AtomicUsize::new(0));
let thief = Router::new().fallback({
let hits = hits.clone();
move || {
hits.fetch_add(1, Ordering::SeqCst);
async { "stolen" }
}
});
let thief_base = serve(thief).await;
let target = format!("{thief_base}/steal");
let app = Router::new().route(
"/first",
get(move || {
let target = target.clone();
async move {
(
AxumStatus::TEMPORARY_REDIRECT,
[(axum::http::header::LOCATION, target)],
)
.into_response()
}
}),
);
let base = serve(app).await;
let response = build_http_client()
.get(format!("{base}/first"))
.header("x-api-key", "super-secret")
.send()
.await
.expect("stopping surfaces the 3xx rather than erroring");
assert_eq!(response.status(), 307, "the redirect is not followed");
assert_eq!(
hits.load(Ordering::SeqCst),
0,
"no request carrying the configured headers may reach another origin"
);
reqwest::Client::new()
.get(format!("{thief_base}/steal"))
.send()
.await
.expect("the thief is listening");
assert_eq!(hits.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_cleartext_remote_server_warns() {
let _guard = always_on_tracing_guard();
transport("http://mcp.example.com/mcp");
transport("http://127.0.0.1:9999/mcp");
transport("https://mcp.example.com/mcp");
}
#[test]
fn an_unparseable_url_is_rejected_up_front() {
let err = HttpTransport::new("not a url", &HashMap::new(), &[])
.err()
.expect("garbage url must not build");
assert!(
err.to_string().contains("Invalid MCP server url"),
"got: {err}"
);
}
#[tokio::test]
async fn a_connection_refusal_is_an_error() {
let _guard = always_on_tracing_guard();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
drop(listener);
let mut t = transport(&format!("http://{addr}/mcp"));
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
}
#[tokio::test]
async fn a_server_error_status_includes_the_body() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async { (AxumStatus::INTERNAL_SERVER_ERROR, "upstream exploded") }),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("500 must fail");
assert!(err.to_string().contains("upstream exploded"), "got: {err}");
}
#[tokio::test]
async fn an_empty_error_body_still_reports_the_status() {
let _guard = always_on_tracing_guard();
let app = Router::new().route("/mcp", post(|| async { AxumStatus::BAD_GATEWAY }));
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("502 must fail");
assert!(err.to_string().contains("502"), "got: {err}");
}
#[tokio::test]
async fn an_unsupported_content_type_is_rejected() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async { ([(CONTENT_TYPE, "text/html")], "<html/>").into_response() }),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("html must fail");
assert!(
err.to_string().contains("unsupported content type"),
"got: {err}"
);
}
#[tokio::test]
async fn a_malformed_json_body_is_a_parse_error() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async { ([(CONTENT_TYPE, "application/json")], "not json").into_response() }),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("garbage must fail");
assert!(err.to_string().contains("parse"), "got: {err}");
}
#[tokio::test]
async fn a_json_reply_that_is_not_a_response_is_rejected() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async {
(
[(CONTENT_TYPE, "application/json")],
serde_json::json!({"jsonrpc": "2.0", "method": "notifications/x"}).to_string(),
)
.into_response()
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("non-response frame must fail");
assert!(err.to_string().contains("non-response frame"), "got: {err}");
}
#[tokio::test]
async fn an_sse_stream_that_ends_early_is_an_error() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async { ([(CONTENT_TYPE, "text/event-stream")], ": bye\n\n").into_response() }),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("truncated stream must fail");
assert!(
err.to_string().contains("ended before answering"),
"got: {err}"
);
}
#[tokio::test]
async fn an_unparseable_sse_frame_is_an_error() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async {
([(CONTENT_TYPE, "text/event-stream")], "data: nonsense\n\n").into_response()
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
}
#[tokio::test]
async fn a_slow_server_times_out() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(std::future::pending::<()>),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), Duration::from_millis(200))
.await
.err()
.expect("must time out");
assert!(err.to_string().contains("did not respond"), "got: {err}");
}
fn resp(id: Value) -> JsonRpcResponse {
serde_json::from_value(serde_json::json!({
"jsonrpc": "2.0", "id": id, "result": {}
}))
.unwrap()
}
#[test]
fn response_matching_accepts_the_expected_id() {
assert!(response_matches(&resp(serde_json::json!(3)), Some(3)));
}
#[test]
fn response_matching_rejects_a_different_id() {
assert!(!response_matches(&resp(serde_json::json!(4)), Some(3)));
}
#[test]
fn response_matching_accepts_a_null_id() {
assert!(response_matches(&resp(Value::Null), Some(3)));
}
#[test]
fn response_matching_rejects_a_non_numeric_id() {
assert!(!response_matches(&resp(serde_json::json!("abc")), Some(3)));
}
#[tokio::test]
async fn close_deletes_a_streamable_session() {
let _guard = always_on_tracing_guard();
let deleted = Arc::new(AtomicUsize::new(0));
let app = Router::new().route(
"/mcp",
post(|| async {
(
[
(CONTENT_TYPE, "application/json"),
(HeaderName::from_static(SESSION_HEADER), "sess-1"),
],
ok_frame(1),
)
.into_response()
})
.delete({
let deleted = deleted.clone();
move || {
let deleted = deleted.clone();
async move {
deleted.fetch_add(1, Ordering::SeqCst);
AxumStatus::NO_CONTENT
}
}
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.unwrap();
t.close().await.unwrap();
assert_eq!(deleted.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn close_without_a_session_is_a_no_op() {
let _guard = always_on_tracing_guard();
let mut t = transport("http://127.0.0.1:1/mcp");
t.close().await.expect("close must always succeed");
}
#[tokio::test]
async fn close_aborts_the_legacy_stream() {
let _guard = always_on_tracing_guard();
let (app, _tx) = legacy_app("/messages");
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.unwrap();
t.close().await.expect("close must always succeed");
assert!(t.legacy.is_none());
}
async fn serve_raw(response: &'static [u8]) -> String {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accept");
let mut buf = [0u8; 8192];
let _ = socket.read(&mut buf).await;
let _ = socket.write_all(response).await;
let _ = socket.flush().await;
let _ = socket.shutdown().await;
});
format!("http://{addr}")
}
#[tokio::test]
async fn a_truncated_json_body_is_a_read_error() {
let _guard = always_on_tracing_guard();
let mut t = transport(
&serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n\
Content-Length: 9000\r\nConnection: close\r\n\r\n{\"jsonrpc\"",
)
.await,
);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("truncated body must fail");
assert!(err.to_string().contains("Failed to read"), "got: {err}");
}
#[tokio::test]
async fn a_truncated_event_stream_is_a_stream_error() {
let _guard = always_on_tracing_guard();
let mut t = transport(
&serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Content-Length: 9000\r\nConnection: close\r\n\r\ndata: partial",
)
.await,
);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("truncated stream must fail");
assert!(
err.to_string().contains("event stream failed"),
"got: {err}"
);
}
#[tokio::test]
async fn a_non_utf8_event_stream_is_an_error() {
let _guard = always_on_tracing_guard();
let mut t = transport(
&serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Connection: close\r\n\r\ndata: \xff\xfe not utf8\n\n",
)
.await,
);
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
}
#[tokio::test]
async fn an_error_body_that_cannot_be_read_still_reports_the_status() {
let _guard = always_on_tracing_guard();
let mut t = transport(
&serve_raw(
b"HTTP/1.1 503 Service Unavailable\r\nContent-Length: 9000\r\n\
Connection: close\r\n\r\nshort",
)
.await,
);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("503 must fail");
assert!(err.to_string().contains("503"), "got: {err}");
}
fn legacy_router(sse: axum::routing::MethodRouter) -> Router {
Router::new().route(
"/sse",
sse.post(|| async { AxumStatus::METHOD_NOT_ALLOWED }),
)
}
#[tokio::test]
async fn a_rejected_event_stream_fails_the_fallback() {
let _guard = always_on_tracing_guard();
let app = legacy_router(get(|| async { AxumStatus::FORBIDDEN }));
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("403 on the stream must fail");
assert!(
err.to_string().contains("rejected the event stream"),
"got: {err}"
);
}
#[tokio::test]
async fn an_event_stream_that_never_names_an_endpoint_fails() {
let _guard = always_on_tracing_guard();
let app = legacy_router(get(|| async {
([(CONTENT_TYPE, "text/event-stream")], ": hello\n\n").into_response()
}));
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("no endpoint means nowhere to POST");
assert!(
err.to_string().contains("before naming a POST endpoint"),
"got: {err}"
);
}
#[tokio::test]
async fn frames_before_the_endpoint_event_are_discarded() {
let _guard = always_on_tracing_guard();
let app = legacy_router(get(|| async {
(
[(CONTENT_TYPE, "text/event-stream")],
concat!(
"data: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/x\"}\n\n",
"data: not json at all\n\n",
"event: endpoint\ndata: /messages\n\n",
),
)
.into_response()
}));
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("stream ends with no answer");
assert!(
!err.to_string().contains("before naming a POST endpoint"),
"the endpoint should have been found: {err}"
);
}
#[tokio::test]
async fn an_unusable_endpoint_url_is_rejected() {
let _guard = always_on_tracing_guard();
let app = legacy_router(get(|| async {
(
[(CONTENT_TYPE, "text/event-stream")],
"event: endpoint\ndata: http://\n\n",
)
.into_response()
}));
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("unusable endpoint must fail");
assert!(
err.to_string().contains("Invalid MCP endpoint"),
"got: {err}"
);
}
#[tokio::test]
async fn a_rejected_legacy_post_fails_the_request() {
let _guard = always_on_tracing_guard();
let app = Router::new()
.route(
"/sse",
get(|| async {
(
[(CONTENT_TYPE, "text/event-stream")],
axum::body::Body::from_stream(async_stream::stream! {
yield Ok::<_, std::io::Error>(
"event: endpoint\ndata: /messages\n\n".to_string());
tokio::time::sleep(Duration::from_secs(5)).await;
}),
)
.into_response()
})
.post(|| async { AxumStatus::METHOD_NOT_ALLOWED }),
)
.route(
"/messages",
post(|| async { (AxumStatus::UNAUTHORIZED, "nope") }),
);
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("401 on the message endpoint must fail");
assert!(err.to_string().contains("nope"), "got: {err}");
}
#[tokio::test]
async fn a_legacy_stream_that_ends_leaves_the_request_unanswered() {
let _guard = always_on_tracing_guard();
let app = Router::new()
.route(
"/sse",
get(|| async {
(
[(CONTENT_TYPE, "text/event-stream")],
"event: endpoint\ndata: /messages\n\n",
)
.into_response()
})
.post(|| async { AxumStatus::METHOD_NOT_ALLOWED }),
)
.route("/messages", post(|| async { AxumStatus::ACCEPTED }));
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("no answer can arrive");
assert!(
err.to_string().contains("ended before answering"),
"got: {err}"
);
}
#[tokio::test]
async fn a_second_request_after_fallback_stays_on_the_legacy_path() {
let _guard = always_on_tracing_guard();
let (app, _tx) = legacy_app("/messages");
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.unwrap();
assert_eq!(t.mode, Mode::Legacy);
let value = t
.send_request(
&JsonRpcRequest::request(2, "tools/list", serde_json::json!({})),
DEFAULT_REQUEST_TIMEOUT,
)
.await
.expect("second request should succeed")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"legacy": true}));
}
#[tokio::test]
async fn a_failed_reply_to_a_server_request_does_not_fail_the_call() {
let _guard = always_on_tracing_guard();
let calls = Arc::new(AtomicUsize::new(0));
let body = format!(
"data: {}\n\ndata: {}\n\n",
serde_json::json!({"jsonrpc": "2.0", "id": 7, "method": "ping"}),
ok_frame(1),
);
let app = Router::new().route(
"/mcp",
post({
let calls = calls.clone();
move || {
let (body, calls) = (body.clone(), calls.clone());
async move {
if calls.fetch_add(1, Ordering::SeqCst) == 0 {
([(CONTENT_TYPE, "text/event-stream")], body).into_response()
} else {
AxumStatus::INTERNAL_SERVER_ERROR.into_response()
}
}
}
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
let value = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("a rejected courtesy reply must not fail the request")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"ok": true}));
}
async fn stream_response(url: &str) -> reqwest::Response {
build_http_client()
.get(url)
.send()
.await
.expect("GET should connect")
}
#[tokio::test]
async fn read_event_stream_stops_when_the_receiver_is_gone() {
let _guard = always_on_tracing_guard();
let url = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Connection: close\r\n\r\ndata: {\"jsonrpc\":\"2.0\",\"id\":1}\n\n",
)
.await;
let (tx, rx) = mpsc::unbounded_channel();
drop(rx);
read_event_stream(stream_response(&url).await, tx).await;
}
#[tokio::test]
async fn read_event_stream_stops_on_invalid_utf8() {
let _guard = always_on_tracing_guard();
let url = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Connection: close\r\n\r\ndata: \xff\xfe\n\n",
)
.await;
let (tx, mut rx) = mpsc::unbounded_channel();
read_event_stream(stream_response(&url).await, tx).await;
assert!(rx.recv().await.is_none(), "nothing decodable was sent");
}
#[tokio::test]
async fn read_event_stream_stops_on_a_stream_error() {
let _guard = always_on_tracing_guard();
let url = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Content-Length: 9000\r\nConnection: close\r\n\r\ndata: partial",
)
.await;
let (tx, mut rx) = mpsc::unbounded_channel();
read_event_stream(stream_response(&url).await, tx).await;
assert!(rx.recv().await.is_none());
}
#[tokio::test]
async fn read_event_stream_ends_cleanly_at_eof() {
let _guard = always_on_tracing_guard();
let url = serve_raw(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Connection: close\r\n\r\nevent: endpoint\ndata: /messages\n\n",
)
.await;
let (tx, mut rx) = mpsc::unbounded_channel();
read_event_stream(stream_response(&url).await, tx).await;
assert!(matches!(
rx.recv().await,
Some(LegacyEvent::Endpoint(path)) if path == "/messages"
));
assert!(rx.recv().await.is_none(), "stream is finished");
}
#[tokio::test]
async fn legacy_skips_non_answers_while_waiting() {
let _guard = always_on_tracing_guard();
let app = Router::new()
.route(
"/sse",
get(|| async {
(
[(CONTENT_TYPE, "text/event-stream")],
axum::body::Body::from_stream(async_stream::stream! {
yield Ok::<_, std::io::Error>(
"event: endpoint\ndata: /messages\n\n".to_string());
tokio::time::sleep(Duration::from_millis(50)).await;
yield Ok("event: endpoint\ndata: /messages\n\n".to_string());
yield Ok(format!(
"data: {}\n\n",
serde_json::json!({
"jsonrpc": "2.0", "method": "notifications/progress"
})));
yield Ok(format!("data: {}\n\n", ok_frame(1)));
}),
)
.into_response()
})
.post(|| async { AxumStatus::METHOD_NOT_ALLOWED }),
)
.route("/messages", post(|| async { AxumStatus::ACCEPTED }));
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
let value = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("the real answer should still arrive")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"ok": true}));
}
#[tokio::test]
async fn a_streamable_json_array_reply_fails_classification() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async { ([(CONTENT_TYPE, "application/json")], "[1,2,3]").into_response() }),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
}
#[tokio::test]
async fn a_streamable_sse_array_frame_fails_classification() {
let _guard = always_on_tracing_guard();
let app = Router::new().route(
"/mcp",
post(|| async {
([(CONTENT_TYPE, "text/event-stream")], "data: [1,2,3]\n\n").into_response()
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
}
#[tokio::test]
async fn a_notification_to_a_dead_server_errors() {
let _guard = always_on_tracing_guard();
let mut t = transport("http://127.0.0.1:1/mcp");
assert!(
t.send_notification(&JsonRpcRequest::notification("x", serde_json::json!({})))
.await
.is_err()
);
}
#[tokio::test]
async fn start_legacy_errors_when_the_stream_cannot_be_opened() {
let _guard = always_on_tracing_guard();
let mut t = transport("http://127.0.0.1:1/sse");
assert!(t.start_legacy().await.is_err());
}
#[tokio::test]
async fn a_legacy_post_to_a_dead_endpoint_errors() {
let _guard = always_on_tracing_guard();
let (app, _tx) = legacy_app("/messages");
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
listener, app,
)));
let mut t = transport(&format!("http://{addr}/sse"));
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("the first request establishes legacy mode");
assert_eq!(t.mode, Mode::Legacy);
server.abort();
let _ = server.await;
assert!(
t.send_request(
&JsonRpcRequest::request(2, "tools/list", serde_json::json!({})),
DEFAULT_REQUEST_TIMEOUT
)
.await
.is_err(),
"a POST to a server that has gone away must fail"
);
}
#[tokio::test]
async fn a_legacy_array_frame_fails_classification() {
let _guard = always_on_tracing_guard();
let app = Router::new()
.route(
"/sse",
get(|| async {
(
[(CONTENT_TYPE, "text/event-stream")],
axum::body::Body::from_stream(async_stream::stream! {
yield Ok::<_, std::io::Error>(
"event: endpoint\ndata: /messages\n\n".to_string());
tokio::time::sleep(Duration::from_millis(50)).await;
yield Ok("data: [1,2,3]\n\n".to_string());
}),
)
.into_response()
})
.post(|| async { AxumStatus::METHOD_NOT_ALLOWED }),
)
.route("/messages", post(|| async { AxumStatus::ACCEPTED }));
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
}
#[tokio::test]
async fn a_notification_over_legacy_hits_the_id_less_post_path() {
let _guard = always_on_tracing_guard();
let (app, _tx) = legacy_app("/messages");
let url = format!("{}/sse", serve(app).await);
let mut t = transport(&url);
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.unwrap();
assert_eq!(t.mode, Mode::Legacy);
t.send_notification(&JsonRpcRequest::notification("x", serde_json::json!({})))
.await
.expect("a 202 to a legacy notification is success");
}
struct CountingRefresher {
calls: Arc<AtomicUsize>,
value: String,
fail: bool,
}
#[async_trait::async_trait]
impl BearerRefresher for CountingRefresher {
async fn refresh(&self) -> anyhow::Result<String> {
self.calls.fetch_add(1, Ordering::SeqCst);
if self.fail {
anyhow::bail!("refresh boom");
}
Ok(self.value.clone())
}
}
#[tokio::test]
async fn a_401_triggers_a_refresh_and_a_successful_retry() {
let _guard = always_on_tracing_guard();
let calls = Arc::new(AtomicUsize::new(0));
let app = Router::new().route(
"/mcp",
post(|headers: AxumHeaders| async move {
let auth = headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if auth == "Bearer fresh" {
([(CONTENT_TYPE, "application/json")], ok_frame(1)).into_response()
} else {
AxumStatus::UNAUTHORIZED.into_response()
}
}),
);
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
t.set_bearer_refresher(Arc::new(CountingRefresher {
calls: calls.clone(),
value: "Bearer fresh".to_string(),
fail: false,
}));
let value = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.expect("refresh + retry should succeed")
.into_result()
.unwrap();
assert_eq!(value, serde_json::json!({"ok": true}));
assert_eq!(calls.load(Ordering::SeqCst), 1, "refreshed exactly once");
}
#[tokio::test]
async fn a_401_without_a_refresher_surfaces_the_error() {
let _guard = always_on_tracing_guard();
let app = Router::new().route("/mcp", post(|| async { AxumStatus::UNAUTHORIZED }));
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
}
#[tokio::test]
async fn a_failed_refresh_propagates() {
let _guard = always_on_tracing_guard();
let calls = Arc::new(AtomicUsize::new(0));
let app = Router::new().route("/mcp", post(|| async { AxumStatus::UNAUTHORIZED }));
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
t.set_bearer_refresher(Arc::new(CountingRefresher {
calls: calls.clone(),
value: String::new(),
fail: true,
}));
let err = t
.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.err()
.expect("a failed refresh must surface");
assert!(err.to_string().contains("refresh boom"), "got: {err}");
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn the_retry_happens_at_most_once() {
let _guard = always_on_tracing_guard();
let calls = Arc::new(AtomicUsize::new(0));
let app = Router::new().route("/mcp", post(|| async { AxumStatus::UNAUTHORIZED }));
let url = format!("{}/mcp", serve(app).await);
let mut t = transport(&url);
t.set_bearer_refresher(Arc::new(CountingRefresher {
calls: calls.clone(),
value: "Bearer still-bad".to_string(),
fail: false,
}));
assert!(
t.send_request(&init(), DEFAULT_REQUEST_TIMEOUT)
.await
.is_err()
);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"refresh tried once, not in a loop"
);
}
#[test]
fn set_auth_header_rejects_an_invalid_value() {
let _guard = always_on_tracing_guard();
let mut t = transport("http://127.0.0.1:1/mcp");
t.set_auth_header("Bearer with\nnewline");
assert!(t.headers.get("authorization").is_none());
t.set_auth_header("Bearer good");
assert_eq!(t.headers.get("authorization").unwrap(), "Bearer good");
}
}