use crate::error::{Error, Result, TransportError};
use crate::shared::http_constants::{
ACCEPT, ACCEPT_STREAMABLE, APPLICATION_JSON, CONTENT_TYPE, MCP_METHOD, MCP_NAME,
MCP_PROTOCOL_VERSION, MCP_SESSION_ID, TEXT_EVENT_STREAM,
};
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
use crate::shared::http_constants::LAST_EVENT_ID;
use crate::shared::sse_parser::SseParser;
use crate::shared::{Transport, TransportMessage};
use crate::types::mrtr::encode_header_value;
use async_trait::async_trait;
use bytes::Bytes;
use http_body_util::{BodyExt, Full, LengthLimitError, Limited};
use hyper::{Method, Request, Response as HyperResponse, StatusCode};
use hyper_util::client::legacy::Client;
use hyper_util::rt::TokioExecutor;
use parking_lot::RwLock;
use std::fmt::Debug;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
#[cfg(not(target_arch = "wasm32"))]
use tokio::sync::mpsc;
use url::Url;
#[cfg_attr(
feature = "v1-compat",
doc = r#"
Resuming an interrupted stream is v1-only (`v1-compat`), so this example is
compiled only when that feature is on:
```rust
use pmcp::shared::streamable_http::SendOptions;
let opts = SendOptions {
related_request_id: None,
resumption_token: Some("event-456".to_string()),
};
assert_eq!(opts.resumption_token.as_deref(), Some("event-456"));
```
"#
)]
#[derive(Debug, Clone, Default)]
pub struct SendOptions {
pub related_request_id: Option<String>,
#[cfg(feature = "v1-compat")]
pub resumption_token: Option<String>,
}
impl SendOptions {
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
fn resumption_cursor(&self) -> Option<String> {
self.resumption_token.clone()
}
#[cfg(not(feature = "v1-compat"))]
#[allow(clippy::unused_self)]
const fn resumption_cursor(&self) -> Option<String> {
None
}
}
#[cfg_attr(
feature = "v1-compat",
doc = r#"
A session-bearing configuration is v1-only (MCP `2025-11-25`), so this example
compiles only when `v1-compat` is on:
```rust
use pmcp::shared::streamable_http::StreamableHttpTransportConfigBuilder;
use url::Url;
let config = StreamableHttpTransportConfigBuilder::new(
Url::parse("http://localhost:8080").unwrap(),
)
.with_session_id("session-123")
.build();
assert_eq!(config.session_id.as_deref(), Some("session-123"));
```
"#
)]
#[derive(Clone)]
pub struct StreamableHttpTransportConfig {
pub url: Url,
pub extra_headers: Vec<(String, String)>,
pub auth_provider: Option<Arc<dyn AuthProvider>>,
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub session_id: Option<String>,
pub enable_json_response: bool,
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub on_resumption_token: Option<Arc<dyn Fn(String) + Send + Sync>>,
pub http_middleware_chain: Option<Arc<crate::client::http_middleware::HttpMiddlewareChain>>,
}
impl StreamableHttpTransportConfig {
#[cfg(feature = "v1-compat")]
fn debug_v1_fields(&self, out: &mut std::fmt::DebugStruct<'_, '_>) {
out.field("session_id", &self.session_id)
.field("on_resumption_token", &self.on_resumption_token.is_some());
}
#[cfg(not(feature = "v1-compat"))]
#[allow(clippy::unused_self)]
const fn debug_v1_fields(&self, _out: &mut std::fmt::DebugStruct<'_, '_>) {}
}
impl Debug for StreamableHttpTransportConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut out = f.debug_struct("StreamableHttpTransportConfig");
out.field("url", &self.url)
.field("extra_headers", &self.extra_headers)
.field("auth_provider", &self.auth_provider.is_some())
.field("enable_json_response", &self.enable_json_response)
.field(
"http_middleware_chain",
&self.http_middleware_chain.is_some(),
);
self.debug_v1_fields(&mut out);
out.finish()
}
}
pub struct StreamableHttpTransportConfigBuilder {
url: Url,
extra_headers: Vec<(String, String)>,
auth_provider: Option<Arc<dyn AuthProvider>>,
#[cfg(feature = "v1-compat")]
session_id: Option<String>,
enable_json_response: bool,
#[cfg(feature = "v1-compat")]
on_resumption_token: Option<Arc<dyn Fn(String) + Send + Sync>>,
http_middleware_chain: Option<Arc<crate::client::http_middleware::HttpMiddlewareChain>>,
}
impl StreamableHttpTransportConfigBuilder {
#[cfg(feature = "v1-compat")]
fn debug_v1_fields(&self, out: &mut std::fmt::DebugStruct<'_, '_>) {
out.field("session_id", &self.session_id)
.field("on_resumption_token", &self.on_resumption_token.is_some());
}
#[cfg(not(feature = "v1-compat"))]
#[allow(clippy::unused_self)]
const fn debug_v1_fields(&self, _out: &mut std::fmt::DebugStruct<'_, '_>) {}
}
impl Debug for StreamableHttpTransportConfigBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut out = f.debug_struct("StreamableHttpTransportConfigBuilder");
out.field("url", &self.url)
.field("extra_headers", &self.extra_headers)
.field("auth_provider", &self.auth_provider.is_some())
.field("enable_json_response", &self.enable_json_response)
.field(
"http_middleware_chain",
&self.http_middleware_chain.is_some(),
);
self.debug_v1_fields(&mut out);
out.finish()
}
}
impl StreamableHttpTransportConfigBuilder {
pub fn new(url: Url) -> Self {
Self {
url,
extra_headers: Vec::new(),
auth_provider: None,
#[cfg(feature = "v1-compat")]
session_id: None,
enable_json_response: false,
#[cfg(feature = "v1-compat")]
on_resumption_token: None,
http_middleware_chain: None,
}
}
pub fn with_header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.extra_headers.push((name.into(), value.into()));
self
}
pub fn with_auth_provider(mut self, provider: Arc<dyn AuthProvider>) -> Self {
self.auth_provider = Some(provider);
self
}
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub fn with_session_id(mut self, session_id: impl Into<String>) -> Self {
self.session_id = Some(session_id.into());
self
}
pub fn enable_json_response(mut self) -> Self {
self.enable_json_response = true;
self
}
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub fn on_resumption_token(mut self, callback: Arc<dyn Fn(String) + Send + Sync>) -> Self {
self.on_resumption_token = Some(callback);
self
}
pub fn with_http_middleware(
mut self,
chain: Arc<crate::client::http_middleware::HttpMiddlewareChain>,
) -> Self {
self.http_middleware_chain = Some(chain);
self
}
pub fn build(self) -> StreamableHttpTransportConfig {
StreamableHttpTransportConfig {
url: self.url,
extra_headers: self.extra_headers,
auth_provider: self.auth_provider,
#[cfg(feature = "v1-compat")]
session_id: self.session_id,
enable_json_response: self.enable_json_response,
#[cfg(feature = "v1-compat")]
on_resumption_token: self.on_resumption_token,
http_middleware_chain: self.http_middleware_chain,
}
}
}
pub const DEFAULT_MAX_COLLECTED_BODY_BYTES: usize = 16 * 1024 * 1024;
fn collected_body_over_cap(max_bytes: usize, declared: Option<usize>) -> Error {
let observed = match declared {
Some(bytes) => format!("declares Content-Length {bytes}"),
None => "delivered more than the cap (Content-Length absent or understated)".to_string(),
};
Error::Transport(TransportError::Request(format!(
"response body {observed}, over this transport's {max_bytes}-byte collected-body cap \
(DEFAULT_MAX_COLLECTED_BODY_BYTES); raise it with \
StreamableHttpTransport::with_max_collected_body_bytes"
)))
}
#[derive(Clone)]
pub struct StreamableHttpTransport {
config: Arc<RwLock<StreamableHttpTransportConfig>>,
client: Client<
hyper_rustls::HttpsConnector<hyper_util::client::legacy::connect::HttpConnector>,
Full<Bytes>,
>,
receiver: Arc<tokio::sync::Mutex<mpsc::UnboundedReceiver<TransportMessage>>>,
sender: mpsc::UnboundedSender<TransportMessage>,
protocol_version: Arc<RwLock<Option<String>>>,
v2_mode: Arc<AtomicBool>,
abort_handle: Arc<RwLock<Option<tokio::task::JoinHandle<()>>>>,
last_event_id: Arc<RwLock<Option<String>>>,
max_collected_body_bytes: usize,
}
impl Debug for StreamableHttpTransport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StreamableHttpTransport")
.field("config", &self.config)
.field("protocol_version", &self.protocol_version)
.field("last_event_id", &self.last_event_id)
.field("max_collected_body_bytes", &self.max_collected_body_bytes)
.finish()
}
}
impl StreamableHttpTransport {
pub fn new(config: StreamableHttpTransportConfig) -> Self {
Self::new_internal(config, false)
}
pub fn new_with_http2(config: StreamableHttpTransportConfig) -> Self {
Self::new_internal(config, true)
}
fn new_internal(config: StreamableHttpTransportConfig, enable_http2: bool) -> Self {
let _ = rustls::crypto::ring::default_provider().install_default();
let https = if enable_http2 {
tracing::debug!("Creating HTTPS connector with HTTP/1.1 and HTTP/2 support");
hyper_rustls::HttpsConnectorBuilder::new()
.with_native_roots()
.expect("Failed to load native root certificates")
.https_or_http()
.enable_http1()
.enable_http2()
.build()
} else {
tracing::debug!("Creating HTTPS connector with HTTP/1.1 only");
hyper_rustls::HttpsConnectorBuilder::new()
.with_native_roots()
.expect("Failed to load native root certificates")
.https_or_http()
.enable_http1()
.build()
};
let client = Client::builder(TokioExecutor::new())
.pool_idle_timeout(std::time::Duration::from_secs(90))
.pool_max_idle_per_host(10)
.build(https);
let (sender, receiver) = mpsc::unbounded_channel();
Self {
config: Arc::new(RwLock::new(config)),
client,
receiver: Arc::new(tokio::sync::Mutex::new(receiver)),
sender,
protocol_version: Arc::new(RwLock::new(None)),
v2_mode: Arc::new(AtomicBool::new(false)),
abort_handle: Arc::new(RwLock::new(None)),
last_event_id: Arc::new(RwLock::new(None)),
max_collected_body_bytes: DEFAULT_MAX_COLLECTED_BODY_BYTES,
}
}
#[must_use]
pub fn with_max_collected_body_bytes(mut self, max_collected_body_bytes: usize) -> Self {
self.max_collected_body_bytes = max_collected_body_bytes;
self
}
async fn collect_body_within_cap(
response: HyperResponse<hyper::body::Incoming>,
max_bytes: usize,
) -> Result<Bytes> {
let declared = response
.headers()
.get(hyper::header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<usize>().ok());
if let Some(declared) = declared {
if declared > max_bytes {
return Err(collected_body_over_cap(max_bytes, Some(declared)));
}
}
match Limited::new(response.into_body(), max_bytes)
.collect()
.await
{
Ok(collected) => Ok(collected.to_bytes()),
Err(error) if error.is::<LengthLimitError>() => {
Err(collected_body_over_cap(max_bytes, None))
},
Err(error) => Err(Error::Transport(TransportError::Request(error.to_string()))),
}
}
pub(crate) async fn collect_capped_body(
&self,
response: HyperResponse<hyper::body::Incoming>,
) -> Result<Bytes> {
Self::collect_body_within_cap(response, self.max_collected_body_bytes).await
}
fn is_v2(&self) -> bool {
self.v2_mode.load(Ordering::Relaxed)
}
#[cfg(feature = "v1-compat")]
fn resumption_callback(&self) -> Option<Arc<dyn Fn(String) + Send + Sync>> {
self.config.read().on_resumption_token.clone()
}
#[cfg(not(feature = "v1-compat"))]
#[allow(clippy::unused_self)]
const fn resumption_callback(&self) -> Option<Arc<dyn Fn(String) + Send + Sync>> {
None
}
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
fn apply_resumption_header(
request: &mut Request<Full<Bytes>>,
resumption_token: Option<&str>,
) -> Result<()> {
if let Some(token) = resumption_token {
request.headers_mut().insert(
LAST_EVENT_ID,
token.parse().map_err(|e| {
Error::Transport(TransportError::InvalidMessage(format!(
"Invalid header: {}",
e
)))
})?,
);
}
Ok(())
}
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub fn session_id(&self) -> Option<String> {
self.config.read().session_id.clone()
}
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub fn set_session_id(&self, session_id: Option<String>) {
self.config.write().session_id = session_id;
}
#[cfg(feature = "v1-compat")]
fn outbound_session_from(config: &StreamableHttpTransportConfig) -> Option<String> {
config.session_id.clone()
}
#[cfg(not(feature = "v1-compat"))]
const fn outbound_session_from(_config: &StreamableHttpTransportConfig) -> Option<String> {
None
}
#[cfg(feature = "v1-compat")]
fn capture_session_header(&self, headers: &hyper::HeaderMap) {
if self.is_v2() {
return;
}
if let Some(value) = headers.get(MCP_SESSION_ID) {
if let Ok(text) = value.to_str() {
if self.config.read().session_id.as_deref() == Some(text) {
return;
}
self.config.write().session_id = Some(text.to_string());
}
}
}
#[cfg(not(feature = "v1-compat"))]
#[allow(clippy::unused_self)]
const fn capture_session_header(&self, _headers: &hyper::HeaderMap) {}
#[cfg(feature = "v1-compat")]
async fn terminate_session(&self) -> Result<()> {
let Some(url) = ({
let config = self.config.read();
config.session_id.is_some().then(|| config.url.clone())
}) else {
return Ok(());
};
let request = self
.build_request_with_middleware(Method::DELETE, url.as_str(), vec![])
.await?;
let response = self.client.request(request).await;
if let Ok(resp) = response {
if !resp.status().is_success() && resp.status() != StatusCode::METHOD_NOT_ALLOWED {
tracing::warn!("Failed to terminate session: {}", resp.status());
}
}
self.config.write().session_id = None;
Ok(())
}
#[cfg(not(feature = "v1-compat"))]
#[allow(clippy::unused_self, clippy::unused_async)]
async fn terminate_session(&self) -> Result<()> {
Ok(())
}
pub fn protocol_version(&self) -> Option<String> {
self.protocol_version.read().clone()
}
pub fn set_protocol_version(&self, version: Option<String>) {
*self.protocol_version.write() = version;
}
pub fn last_event_id(&self) -> Option<String> {
self.last_event_id.read().clone()
}
pub async fn start_sse(
&self,
#[cfg(feature = "v1-compat")] resumption_token: Option<String>,
#[cfg(not(feature = "v1-compat"))] _ignored_cursor: Option<String>,
) -> Result<()> {
let handle = self.abort_handle.write().take();
if let Some(handle) = handle {
handle.abort();
}
let url = self.config.read().url.clone();
let mut request = self
.build_request_with_middleware(
Method::GET,
url.as_str(),
vec![], )
.await?;
request.headers_mut().insert(
ACCEPT,
TEXT_EVENT_STREAM.parse().map_err(|e| {
Error::Transport(TransportError::InvalidMessage(format!(
"Invalid header: {}",
e
)))
})?,
);
#[cfg(feature = "v1-compat")]
Self::apply_resumption_header(&mut request, resumption_token.as_deref())?;
let response = self
.client
.request(request)
.await
.map_err(|e| Error::Transport(TransportError::Request(e.to_string())))?;
if response.status() == StatusCode::METHOD_NOT_ALLOWED {
return Ok(());
}
if !response.status().is_success() {
return Err(Error::Transport(TransportError::Request(format!(
"SSE request failed with status: {}",
response.status()
))));
}
self.process_response_headers(&response);
let body_bytes =
Self::collect_body_within_cap(response, self.max_collected_body_bytes).await?;
let modified_body = if self.config.read().http_middleware_chain.is_some() {
let temp_response = HyperResponse::builder()
.status(200)
.body(Full::new(Bytes::new()))
.unwrap();
self.apply_response_middleware("GET", url.as_str(), &temp_response, body_bytes.to_vec())
.await?
} else {
body_bytes.to_vec()
};
let sender = self.sender.clone();
let on_resumption = self.resumption_callback();
let last_event_id = self.last_event_id.clone();
let handle = tokio::spawn(async move {
let mut sse_parser = SseParser::new();
let body = String::from_utf8_lossy(&modified_body);
let events = sse_parser.feed_complete_body(&body);
for event in events {
if let Some(id) = &event.id {
*last_event_id.write() = Some(id.clone());
if let Some(callback) = &on_resumption {
callback(id.clone());
}
}
if event.event.as_deref() == Some("message") || event.event.is_none() {
if let Ok(msg) =
crate::shared::StdioTransport::parse_message(event.data.as_bytes())
{
let _ = sender.send(msg);
}
}
}
});
*self.abort_handle.write() = Some(handle);
Ok(())
}
fn apply_v2_outbound_headers(
mut builder: hyper::http::request::Builder,
method: &str,
name: &str,
) -> hyper::http::request::Builder {
if let Ok(value) = hyper::header::HeaderValue::from_str(method) {
builder = builder.header(MCP_METHOD, value);
}
if let Ok(value) = hyper::header::HeaderValue::from_str(name) {
builder = builder.header(MCP_NAME, value);
}
builder
}
async fn build_request_with_middleware(
&self,
method: Method,
url: &str,
body: Vec<u8>,
) -> Result<Request<Full<Bytes>>> {
use crate::client::http_middleware::{HttpMiddlewareContext, HttpRequest};
let (extra_headers, auth_provider, middleware_chain, outbound_session) = {
let config = self.config.read();
(
config.extra_headers.clone(),
config.auth_provider.clone(),
config.http_middleware_chain.clone(),
Self::outbound_session_from(&config),
)
};
let mut request_builder = Request::builder().method(method.clone()).uri(url);
for (key, value) in &extra_headers {
request_builder = request_builder.header(key.as_str(), value.as_str());
}
let has_auth = if let Some(auth_provider) = auth_provider {
let token = auth_provider.get_access_token().await?;
request_builder = request_builder.header("Authorization", format!("Bearer {}", token));
true
} else {
false
};
let is_v2 = self.is_v2();
if let Some(session) = &outbound_session {
if !is_v2 {
request_builder = request_builder.header(MCP_SESSION_ID, session.as_str());
}
}
if let Some(protocol_version) = self.protocol_version.read().as_ref() {
request_builder =
request_builder.header(MCP_PROTOCOL_VERSION, protocol_version.as_str());
}
if is_v2 {
if let Some((method, name)) = v2_routing_headers(&body) {
request_builder = Self::apply_v2_outbound_headers(request_builder, &method, &name);
}
}
let temp_req = request_builder
.body(Full::new(Bytes::from(body.clone())))
.map_err(|e| Error::Transport(TransportError::InvalidMessage(e.to_string())))?;
let headers = temp_req.headers();
if let Some(chain) = middleware_chain {
let mut http_req = HttpRequest::new(method.as_str().to_string(), url.to_string(), body);
for (key, value) in headers {
if let Ok(value_str) = value.to_str() {
http_req.add_header(key.as_str(), value_str);
}
}
let context = HttpMiddlewareContext::new(url.to_string(), method.as_str().to_string());
if has_auth {
context.set_metadata("auth_already_set".to_string(), "true".to_string());
}
if let Err(e) = chain.process_request(&mut http_req, &context).await {
chain.handle_transport_error(&e, &context).await;
return Err(e);
}
let mut final_builder = Request::builder().method(method).uri(url);
for (key, value) in &http_req.headers {
final_builder = final_builder.header(key, value);
}
final_builder
.body(Full::new(Bytes::from(http_req.body)))
.map_err(|e| Error::Transport(TransportError::InvalidMessage(e.to_string())))
} else {
Ok(temp_req)
}
}
#[allow(clippy::future_not_send)]
async fn apply_response_middleware(
&self,
method: &str,
url: &str,
response: &HyperResponse<impl hyper::body::Body>,
body: Vec<u8>,
) -> Result<Vec<u8>> {
use crate::client::http_middleware::{HttpMiddlewareContext, HttpResponse};
let middleware_chain = self.config.read().http_middleware_chain.clone();
if let Some(chain) = middleware_chain {
let header_map = response.headers().clone();
let mut http_resp =
HttpResponse::with_headers(response.status().as_u16(), header_map, body);
let context = HttpMiddlewareContext::new(url.to_string(), method.to_string());
if let Err(e) = chain.process_response(&mut http_resp, &context).await {
chain.handle_transport_error(&e, &context).await;
return Err(e);
}
Ok(http_resp.body)
} else {
Ok(body)
}
}
fn process_response_headers(&self, response: &HyperResponse<impl hyper::body::Body>) {
self.capture_session_header(response.headers());
if let Some(protocol_version) = response.headers().get(MCP_PROTOCOL_VERSION) {
if let Ok(protocol_version_str) = protocol_version.to_str() {
*self.protocol_version.write() = Some(protocol_version_str.to_string());
}
}
}
pub async fn send_with_options(
&mut self,
message: TransportMessage,
options: SendOptions,
) -> Result<()> {
if let Some(token) = options.resumption_cursor() {
self.start_sse(Some(token)).await?;
return Ok(());
}
let body_bytes = crate::shared::StdioTransport::serialize_message(&message)?;
let is_notification = matches!(message, TransportMessage::Notification { .. });
self.post_body(body_bytes, is_notification).await
}
async fn jsonrpc_error_envelope(
response: HyperResponse<hyper::body::Incoming>,
max_collected_body_bytes: usize,
) -> Option<TransportMessage> {
let body = Self::collect_body_within_cap(response, max_collected_body_bytes)
.await
.ok()?;
let value = serde_json::from_slice::<serde_json::Value>(&body).ok()?;
if value.get("jsonrpc").and_then(serde_json::Value::as_str) != Some("2.0")
|| value.get("error").is_none()
{
return None;
}
match crate::shared::StdioTransport::parse_message(&body) {
Ok(message @ TransportMessage::Response(_)) => Some(message),
_ => None,
}
}
fn apply_post_headers(request: &mut Request<Full<Bytes>>) -> Result<()> {
request.headers_mut().insert(
CONTENT_TYPE,
APPLICATION_JSON.parse().map_err(|e| {
Error::Transport(TransportError::InvalidMessage(format!(
"Invalid header: {}",
e
)))
})?,
);
request.headers_mut().insert(
ACCEPT,
ACCEPT_STREAMABLE.parse().map_err(|e| {
Error::Transport(TransportError::InvalidMessage(format!(
"Invalid header: {}",
e
)))
})?,
);
Ok(())
}
async fn post_once(&self, body_bytes: Vec<u8>) -> Result<HyperResponse<hyper::body::Incoming>> {
let body_bytes_snapshot = body_bytes.clone();
let url = self.config.read().url.clone();
let mut request = self
.build_request_with_middleware(Method::POST, url.as_str(), body_bytes)
.await?;
Self::apply_post_headers(&mut request)?;
let response = self
.client
.request(request)
.await
.map_err(|e| Error::Transport(TransportError::Request(e.to_string())))?;
if response.status() != StatusCode::UNAUTHORIZED {
return Ok(response);
}
let auth_provider = self.config.read().auth_provider.clone();
let Some(provider) = auth_provider else {
return Ok(response);
};
provider.on_unauthorized().await?;
let mut retry_request = self
.build_request_with_middleware(Method::POST, url.as_str(), body_bytes_snapshot)
.await?;
Self::apply_post_headers(&mut retry_request)?;
self.client
.request(retry_request)
.await
.map_err(|e| Error::Transport(TransportError::Request(e.to_string())))
}
pub(crate) async fn post_streaming(
&self,
body_bytes: Vec<u8>,
) -> Result<HyperResponse<hyper::body::Incoming>> {
let response = self.post_once(body_bytes).await?;
self.process_response_headers(&response);
Ok(response)
}
async fn post_body(&self, body_bytes: Vec<u8>, is_notification: bool) -> Result<()> {
let response = self.post_once(body_bytes).await?;
self.process_response_headers(&response);
if !response.status().is_success() {
if response.status() == StatusCode::ACCEPTED {
if is_notification {
let _ = self.start_sse(None).await;
}
return Ok(());
}
if self.is_v2() {
let status = response.status();
match Self::jsonrpc_error_envelope(response, self.max_collected_body_bytes).await {
Some(message) => {
tracing::debug!(
%status,
"v2 non-2xx carried a JSON-RPC error envelope — surfacing it structurally"
);
self.sender
.send(message)
.map_err(|e| Error::Transport(TransportError::Send(e.to_string())))?;
return Ok(());
},
None => {
return Err(Error::Transport(TransportError::Request(format!(
"Request failed with status: {}",
status
))));
},
}
}
return Err(Error::Transport(TransportError::Request(format!(
"Request failed with status: {}",
response.status()
))));
}
let status_code = response.status();
let content_type = response
.headers()
.get(CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.to_string();
let content_length = response
.headers()
.get("content-length")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<usize>().ok());
let body_bytes =
Self::collect_body_within_cap(response, self.max_collected_body_bytes).await?;
tracing::debug!(
status = %status_code,
content_type = %content_type,
content_length = ?content_length,
body_len = body_bytes.len(),
"HTTP response received"
);
let middleware_url = {
let config = self.config.read();
config
.http_middleware_chain
.is_some()
.then(|| config.url.clone())
};
let modified_body = if let Some(url) = middleware_url {
let temp_response = HyperResponse::builder()
.status(status_code)
.body(Full::new(Bytes::new()))
.unwrap();
self.apply_response_middleware(
"POST",
url.as_str(),
&temp_response,
body_bytes.to_vec(),
)
.await?
} else {
body_bytes.to_vec()
};
if status_code == StatusCode::OK && (content_length == Some(0) || content_type.is_empty()) {
if modified_body.is_empty() {
return Ok(());
}
if content_type.is_empty() {
return Err(Error::Transport(TransportError::Request(
"Response has body but no Content-Type header".to_string(),
)));
}
if let Ok(batch) = serde_json::from_slice::<Vec<serde_json::Value>>(&modified_body) {
for json_msg in batch {
let json_str = serde_json::to_string(&json_msg).map_err(|e| {
Error::Transport(TransportError::Deserialization(e.to_string()))
})?;
let msg = crate::shared::StdioTransport::parse_message(json_str.as_bytes())?;
self.sender
.send(msg)
.map_err(|e| Error::Transport(TransportError::Send(e.to_string())))?;
}
} else {
let msg_parsed = crate::shared::StdioTransport::parse_message(&modified_body)?;
self.sender
.send(msg_parsed)
.map_err(|e| Error::Transport(TransportError::Send(e.to_string())))?;
}
return Ok(());
}
if content_type.contains(APPLICATION_JSON) {
if modified_body.is_empty() {
if status_code == StatusCode::ACCEPTED {
tracing::debug!(
status = %status_code,
"Notification acknowledged with 202 Accepted"
);
return Ok(());
}
tracing::warn!(
status = %status_code,
content_type = %content_type,
"Server returned empty body with application/json content type"
);
return Err(Error::Transport(TransportError::Request(
"Server returned empty response body with Content-Type: application/json. \
This may indicate a server error or network issue."
.to_string(),
)));
}
if let Ok(batch) = serde_json::from_slice::<Vec<serde_json::Value>>(&modified_body) {
for json_msg in batch {
let json_str = serde_json::to_string(&json_msg).map_err(|e| {
Error::Transport(TransportError::Deserialization(e.to_string()))
})?;
let msg = crate::shared::StdioTransport::parse_message(json_str.as_bytes())?;
self.sender
.send(msg)
.map_err(|e| Error::Transport(TransportError::Send(e.to_string())))?;
}
} else {
let msg_parsed = crate::shared::StdioTransport::parse_message(&modified_body)?;
self.sender
.send(msg_parsed)
.map_err(|e| Error::Transport(TransportError::Send(e.to_string())))?;
}
} else if content_type.contains(TEXT_EVENT_STREAM) {
let sender = self.sender.clone();
let on_resumption = self.resumption_callback();
let last_event_id = self.last_event_id.clone();
tokio::spawn(async move {
let mut sse_parser = SseParser::new();
let body = String::from_utf8_lossy(&modified_body);
let events = sse_parser.feed_complete_body(&body);
for event in events {
if let Some(id) = &event.id {
*last_event_id.write() = Some(id.clone());
if let Some(callback) = &on_resumption {
callback(id.clone());
}
}
if event.event.as_deref() == Some("message") || event.event.is_none() {
if let Ok(msg) =
crate::shared::StdioTransport::parse_message(event.data.as_bytes())
{
let _ = sender.send(msg);
}
}
}
});
} else if status_code == StatusCode::ACCEPTED {
return Ok(());
} else {
return Err(Error::Transport(TransportError::Request(format!(
"Unsupported content type: {}",
content_type
))));
}
Ok(())
}
}
#[async_trait]
impl Transport for StreamableHttpTransport {
async fn send(&mut self, message: TransportMessage) -> Result<()> {
self.send_with_options(message, SendOptions::default())
.await
}
async fn receive(&mut self) -> Result<TransportMessage> {
let mut receiver = self.receiver.lock().await;
receiver
.recv()
.await
.ok_or_else(|| Error::Transport(TransportError::ConnectionClosed))
}
async fn close(&mut self) -> Result<()> {
let handle = self.abort_handle.write().take();
if let Some(handle) = handle {
handle.abort();
}
self.terminate_session().await
}
fn is_connected(&self) -> bool {
true
}
fn transport_type(&self) -> &'static str {
"streamable-http"
}
fn set_negotiated_protocol_version(&mut self, version: Option<String>) {
let is_v2 = version.as_deref().map(crate::types::protocol::protocol_era)
== Some(crate::types::protocol::Era::V2);
self.set_protocol_version(version);
self.v2_mode.store(is_v2, Ordering::Relaxed);
}
fn supports_negotiated_protocol_version(&self) -> bool {
true
}
async fn send_raw(&mut self, body: Vec<u8>) -> Result<()> {
self.post_body(body, false).await
}
}
fn v2_routing_headers(body: &[u8]) -> Option<(String, String)> {
let value = serde_json::from_slice::<serde_json::Value>(body).ok()?;
let (method, name) = crate::types::mrtr::frame_routing_pair(&value)?;
Some((
method.to_string(),
encode_header_value(name.as_deref().unwrap_or_default()),
))
}
#[async_trait]
pub trait AuthProvider: Send + Sync + Debug {
async fn get_access_token(&self) -> Result<String>;
async fn on_unauthorized(&self) -> Result<()> {
Ok(())
}
}
#[cfg(all(test, not(target_arch = "wasm32"), feature = "streamable-http"))]
mod tests {
use super::*;
use crate::shared::TransportMessage;
use mockito::Server as MockServer;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex as StdMutex;
use url::Url;
#[derive(Debug)]
struct CountingProvider {
token: String,
get_count: AtomicUsize,
unauthorized_count: AtomicUsize,
call_order: Option<StdMutex<Vec<&'static str>>>,
}
impl CountingProvider {
fn new(token: impl Into<String>) -> Self {
Self {
token: token.into(),
get_count: AtomicUsize::new(0),
unauthorized_count: AtomicUsize::new(0),
call_order: None,
}
}
fn with_order_tracking(token: impl Into<String>) -> Self {
Self {
token: token.into(),
get_count: AtomicUsize::new(0),
unauthorized_count: AtomicUsize::new(0),
call_order: Some(StdMutex::new(Vec::new())),
}
}
}
#[async_trait]
impl AuthProvider for CountingProvider {
async fn get_access_token(&self) -> Result<String> {
self.get_count.fetch_add(1, Ordering::SeqCst);
if let Some(order) = &self.call_order {
order.lock().unwrap().push("get_access_token");
}
Ok(self.token.clone())
}
async fn on_unauthorized(&self) -> Result<()> {
self.unauthorized_count.fetch_add(1, Ordering::SeqCst);
if let Some(order) = &self.call_order {
order.lock().unwrap().push("on_unauthorized");
}
Ok(())
}
}
fn make_transport(
url: Url,
provider: Option<Arc<dyn AuthProvider>>,
) -> StreamableHttpTransport {
let mut builder = StreamableHttpTransportConfigBuilder::new(url);
if let Some(p) = provider {
builder = builder.with_auth_provider(p);
}
let config = builder.build();
StreamableHttpTransport::new(config)
}
fn ping_message() -> TransportMessage {
use crate::types::{ClientNotification, Notification};
TransportMessage::Notification(Notification::Client(ClientNotification::Initialized))
}
fn list_tools_message() -> TransportMessage {
use crate::types::{ClientRequest, ListToolsRequest, Request, RequestId};
TransportMessage::Request {
id: RequestId::from(42i64),
request: Request::Client(Box::new(ClientRequest::ListTools(ListToolsRequest {
cursor: None,
}))),
}
}
#[tokio::test]
async fn test_on_unauthorized_default_noop_compiles_and_succeeds() {
#[derive(Debug)]
struct MinimalProvider;
#[async_trait]
impl AuthProvider for MinimalProvider {
async fn get_access_token(&self) -> Result<String> {
Ok("token".to_string())
}
}
let p = MinimalProvider;
let result = p.on_unauthorized().await;
assert!(
result.is_ok(),
"default on_unauthorized should return Ok(())"
);
}
#[tokio::test]
async fn test_max_one_retry_on_401() {
let mut server = MockServer::new_async().await;
let _m = server
.mock("POST", "/")
.with_status(401)
.with_header("content-type", "application/json")
.with_body(r#"{"error":"unauthorized"}"#)
.expect(2) .create_async()
.await;
let url = Url::parse(&server.url()).unwrap();
let provider = Arc::new(CountingProvider::new("initial-token"));
let mut transport = make_transport(url, Some(provider.clone() as Arc<dyn AuthProvider>));
let _ = transport
.send_with_options(ping_message(), SendOptions::default())
.await;
assert_eq!(
provider.unauthorized_count.load(Ordering::SeqCst),
1,
"on_unauthorized should be called exactly once"
);
assert_eq!(
provider.get_count.load(Ordering::SeqCst),
2,
"get_access_token should be called twice (once per attempt)"
);
}
#[tokio::test]
async fn test_on_unauthorized_not_called_for_non_401() {
let mut server = MockServer::new_async().await;
let _m200 = server
.mock("POST", "/")
.with_status(200)
.with_header("content-type", "application/json")
.with_body(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#)
.create_async()
.await;
let url = Url::parse(&server.url()).unwrap();
let provider = Arc::new(CountingProvider::new("token"));
let mut transport =
make_transport(url.clone(), Some(provider.clone() as Arc<dyn AuthProvider>));
let _ = transport
.send_with_options(ping_message(), SendOptions::default())
.await;
assert_eq!(
provider.unauthorized_count.load(Ordering::SeqCst),
0,
"on_unauthorized must NOT be called on 200"
);
let mut server2 = MockServer::new_async().await;
let _m500 = server2
.mock("POST", "/")
.with_status(500)
.with_header("content-type", "application/json")
.with_body(r#"{"error":"server error"}"#)
.create_async()
.await;
let url2 = Url::parse(&server2.url()).unwrap();
let provider2 = Arc::new(CountingProvider::new("token"));
let mut transport2 = make_transport(url2, Some(provider2.clone() as Arc<dyn AuthProvider>));
let _ = transport2
.send_with_options(ping_message(), SendOptions::default())
.await;
assert_eq!(
provider2.unauthorized_count.load(Ordering::SeqCst),
0,
"on_unauthorized must NOT be called on 500"
);
}
#[tokio::test]
async fn test_retry_body_and_headers_are_byte_identical() {
use hyper::service::service_fn;
use hyper_util::rt::TokioExecutor;
use hyper_util::server::conn::auto::Builder as ServerBuilder;
use std::sync::Mutex as StdMutex;
use tokio::net::TcpListener;
#[derive(Debug, Default)]
struct Captured {
requests: Vec<(String, Vec<u8>, String)>,
}
#[derive(Debug)]
struct DualTokenProvider {
call_count: AtomicUsize,
}
#[async_trait]
impl AuthProvider for DualTokenProvider {
async fn get_access_token(&self) -> Result<String> {
let n = self.call_count.fetch_add(1, Ordering::SeqCst);
if n == 0 {
Ok("token-attempt-1".to_string())
} else {
Ok("token-attempt-2".to_string())
}
}
}
let captured = Arc::new(StdMutex::new(Captured::default()));
let captured_clone = captured.clone();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let cap = captured_clone.clone();
tokio::spawn(async move {
let mut attempt = 0u8;
loop {
let (stream, _) = listener.accept().await.unwrap();
let cap = cap.clone();
let io = hyper_util::rt::TokioIo::new(stream);
tokio::spawn(async move {
let _ = ServerBuilder::new(TokioExecutor::new())
.serve_connection(
io,
service_fn(move |req: Request<hyper::body::Incoming>| {
let cap = cap.clone();
async move {
let method = req.method().to_string();
let auth = req
.headers()
.get("authorization")
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.to_string();
let body_bytes = req
.collect()
.await
.map(|b| b.to_bytes().to_vec())
.unwrap_or_default();
cap.lock()
.unwrap()
.requests
.push((method, body_bytes, auth));
let status = {
let len = cap.lock().unwrap().requests.len();
if len == 1 {
401u16
} else {
200u16
}
};
Ok::<_, hyper::Error>(
HyperResponse::builder()
.status(status)
.header("content-type", "application/json")
.body(Full::new(Bytes::from(if status == 200 {
r#"{"jsonrpc":"2.0","id":1,"result":{}}"#
} else {
r#"{"error":"unauthorized"}"#
})))
.unwrap(),
)
}
}),
)
.await;
});
attempt += 1;
if attempt >= 2 {
break;
}
}
});
let provider = Arc::new(DualTokenProvider {
call_count: AtomicUsize::new(0),
});
let url = Url::parse(&format!("http://127.0.0.1:{}", addr.port())).unwrap();
let mut transport = make_transport(url, Some(provider as Arc<dyn AuthProvider>));
let _ = transport
.send_with_options(list_tools_message(), SendOptions::default())
.await;
let cap = captured.lock().unwrap();
assert_eq!(
cap.requests.len(),
2,
"expected exactly 2 requests (original + retry)"
);
let (method1, body1, auth1) = &cap.requests[0];
let (method2, body2, auth2) = &cap.requests[1];
assert_eq!(
method1, method2,
"method must be byte-identical across retry"
);
assert_eq!(body1, body2, "body must be byte-identical across retry");
assert_ne!(
auth1, auth2,
"Authorization header should differ (new token)"
);
assert!(auth1.contains("token-attempt-1"), "first auth: {}", auth1);
assert!(auth2.contains("token-attempt-2"), "retry auth: {}", auth2);
}
#[tokio::test]
async fn test_on_unauthorized_called_before_get_access_token_on_retry() {
let mut server = MockServer::new_async().await;
let _m = server
.mock("POST", "/")
.with_status(401)
.with_header("content-type", "application/json")
.with_body(r#"{"error":"unauthorized"}"#)
.expect(2)
.create_async()
.await;
let url = Url::parse(&server.url()).unwrap();
let provider = Arc::new(CountingProvider::with_order_tracking("token"));
let mut transport = make_transport(url, Some(provider.clone() as Arc<dyn AuthProvider>));
let _ = transport
.send_with_options(ping_message(), SendOptions::default())
.await;
let order = provider
.call_order
.as_ref()
.unwrap()
.lock()
.unwrap()
.clone();
assert!(
order.len() >= 3,
"expected at least 3 calls, got {:?}",
order
);
let unauth_pos = order
.iter()
.position(|&s| s == "on_unauthorized")
.expect("on_unauthorized must appear in call order");
let retry_get_pos = order
.iter()
.skip(unauth_pos + 1)
.position(|&s| s == "get_access_token");
assert!(
retry_get_pos.is_some(),
"get_access_token must be called AFTER on_unauthorized; order = {:?}",
order
);
}
mod v2_outbound {
use super::*;
use crate::types::protocol::PROTOCOL_VERSION_2026_07_28;
use serde_json::json;
fn body(method: &str, params: &serde_json::Value) -> Vec<u8> {
json!({ "jsonrpc": "2.0", "id": 1, "method": method, "params": params })
.to_string()
.into_bytes()
}
#[cfg(feature = "v1-compat")]
fn plant_session_id(config: &mut StreamableHttpTransportConfig, session_id: Option<&str>) {
config.session_id = session_id.map(str::to_string);
}
#[cfg(not(feature = "v1-compat"))]
#[allow(clippy::missing_const_for_fn)]
fn plant_session_id(
_config: &mut StreamableHttpTransportConfig,
_session_id: Option<&str>,
) {
}
fn v2_transport(session_id: Option<&str>) -> StreamableHttpTransport {
let mut config = StreamableHttpTransportConfigBuilder::new(
Url::parse("http://127.0.0.1:1/").unwrap(),
)
.build();
plant_session_id(&mut config, session_id);
let mut transport = StreamableHttpTransport::new(config);
transport
.set_negotiated_protocol_version(Some(PROTOCOL_VERSION_2026_07_28.to_string()));
transport
}
fn v1_transport(session_id: Option<&str>) -> StreamableHttpTransport {
let mut config = StreamableHttpTransportConfigBuilder::new(
Url::parse("http://127.0.0.1:1/").unwrap(),
)
.build();
plant_session_id(&mut config, session_id);
StreamableHttpTransport::new(config)
}
async fn headers_for(
transport: &StreamableHttpTransport,
body: Vec<u8>,
) -> hyper::HeaderMap {
transport
.build_request_with_middleware(Method::POST, "http://127.0.0.1:1/", body)
.await
.expect("request builds")
.headers()
.clone()
}
fn header(map: &hyper::HeaderMap, name: &str) -> Option<String> {
map.get(name)
.and_then(|v| v.to_str().ok())
.map(str::to_string)
}
#[test]
fn routing_headers_read_name_for_tools_call() {
let derived = v2_routing_headers(&body("tools/call", &json!({ "name": "search" })));
assert_eq!(
derived,
Some(("tools/call".to_string(), "search".to_string()))
);
}
#[test]
fn routing_headers_read_name_for_prompts_get() {
let derived = v2_routing_headers(&body("prompts/get", &json!({ "name": "greeting" })));
assert_eq!(
derived,
Some(("prompts/get".to_string(), "greeting".to_string()))
);
}
#[test]
fn routing_headers_read_uri_for_resources_read() {
let derived =
v2_routing_headers(&body("resources/read", &json!({ "uri": "mem://greeting" })));
assert_eq!(
derived,
Some(("resources/read".to_string(), "mem://greeting".to_string()))
);
}
#[test]
fn routing_headers_read_task_id_for_the_three_tasks_methods() {
for method in ["tasks/get", "tasks/update", "tasks/cancel"] {
let derived = v2_routing_headers(&body(method, &json!({ "taskId": "abc" })));
assert_eq!(
derived,
Some((method.to_string(), "abc".to_string())),
"{method} must route on its taskId"
);
}
}
#[test]
fn routing_headers_are_empty_for_tasks_list_and_tasks_result() {
for method in ["tasks/list", "tasks/result"] {
let derived = v2_routing_headers(&body(method, &json!({ "taskId": "abc" })));
assert_eq!(
derived,
Some((method.to_string(), String::new())),
"{method} is not name-bearing"
);
}
}
#[test]
fn routing_headers_are_none_for_a_body_without_a_method() {
assert_eq!(
v2_routing_headers(br#"{"jsonrpc":"2.0","id":1,"result":{}}"#),
None
);
assert_eq!(v2_routing_headers(b"not json"), None);
assert_eq!(v2_routing_headers(b""), None);
}
#[test]
fn routing_headers_sentinel_encode_a_non_ascii_name() {
let (_, name) = v2_routing_headers(&body("tools/call", &json!({ "name": "поиск" })))
.expect("derived");
assert!(
name.starts_with(crate::types::mrtr::HEADER_SENTINEL_PREFIX),
"a non-header-safe name must travel as a sentinel, got {name}"
);
assert_eq!(
crate::types::mrtr::decode_header_value(&name).as_deref(),
Some("поиск"),
"the shared codec must round-trip"
);
}
#[tokio::test]
async fn v2_tools_call_emits_all_three_headers() {
let transport = v2_transport(None);
let map = headers_for(
&transport,
body("tools/call", &json!({ "name": "search", "arguments": {} })),
)
.await;
assert_eq!(header(&map, MCP_METHOD).as_deref(), Some("tools/call"));
assert_eq!(header(&map, MCP_NAME).as_deref(), Some("search"));
assert_eq!(
header(&map, MCP_PROTOCOL_VERSION).as_deref(),
Some(PROTOCOL_VERSION_2026_07_28)
);
}
#[tokio::test]
async fn v2_nameless_method_emits_an_empty_mcp_name() {
let transport = v2_transport(None);
let map = headers_for(&transport, body("tools/list", &json!({}))).await;
assert_eq!(header(&map, MCP_METHOD).as_deref(), Some("tools/list"));
assert!(
map.contains_key(MCP_NAME),
"this client emits Mcp-Name on every v2 request"
);
assert_eq!(header(&map, MCP_NAME).as_deref(), Some(""));
}
#[tokio::test]
async fn v2_resources_read_puts_the_uri_in_mcp_name() {
let transport = v2_transport(None);
let map = headers_for(
&transport,
body("resources/read", &json!({ "uri": "mem://greeting" })),
)
.await;
assert_eq!(header(&map, MCP_NAME).as_deref(), Some("mem://greeting"));
}
#[tokio::test]
async fn v2_lists_both_accept_content_types() {
assert_eq!(ACCEPT_STREAMABLE, "application/json, text/event-stream");
}
#[tokio::test]
async fn v2_never_emits_a_stored_session_id() {
let transport = v2_transport(Some("left-over-from-v1"));
let map = headers_for(&transport, body("tools/list", &json!({}))).await;
assert!(
!map.contains_key(MCP_SESSION_ID),
"a session id must never reach the v2 wire, even when one is stored"
);
}
#[cfg(feature = "v1-compat")]
#[test]
fn v2_does_not_store_a_session_id_from_a_response() {
let transport = v2_transport(None);
let response = HyperResponse::builder()
.status(StatusCode::OK)
.header(MCP_SESSION_ID, "planted")
.body(Full::new(Bytes::new()))
.unwrap();
transport.process_response_headers(&response);
assert_eq!(
transport.session_id(),
None,
"a v2 response's Mcp-Session-Id must not be stored"
);
}
#[cfg(feature = "v1-compat")]
#[test]
fn v1_still_stores_a_session_id_from_a_response() {
let transport = v1_transport(None);
let response = HyperResponse::builder()
.status(StatusCode::OK)
.header(MCP_SESSION_ID, "kept")
.body(Full::new(Bytes::new()))
.unwrap();
transport.process_response_headers(&response);
assert_eq!(transport.session_id().as_deref(), Some("kept"));
}
#[test]
fn a_server_echo_cannot_flip_a_v1_client_into_v2() {
let transport = v1_transport(Some("s1"));
let response = HyperResponse::builder()
.status(StatusCode::OK)
.header(MCP_PROTOCOL_VERSION, PROTOCOL_VERSION_2026_07_28)
.body(Full::new(Bytes::new()))
.unwrap();
transport.process_response_headers(&response);
assert!(!transport.is_v2(), "only the client selects the era");
}
const V1_SESSION_EXISTS: bool = cfg!(feature = "v1-compat");
#[tokio::test]
async fn v1_emits_no_v2_routing_headers_and_keeps_its_session() {
let transport = v1_transport(Some("session-123"));
let map = headers_for(
&transport,
body("tools/call", &json!({ "name": "search", "arguments": {} })),
)
.await;
assert!(!map.contains_key(MCP_METHOD));
assert!(!map.contains_key(MCP_NAME));
assert_eq!(
header(&map, MCP_SESSION_ID).as_deref(),
V1_SESSION_EXISTS.then_some("session-123"),
"a v1 client emits its stored session id; a severed client has none to emit"
);
}
proptest::proptest! {
#[test]
fn header_emission_never_panics_for_any_method_or_name(
method in ".{0,64}",
name in ".{0,64}",
) {
let frame = body(&method, &json!({ "name": name, "uri": name }));
let derived = v2_routing_headers(&frame);
if let Some((m, n)) = derived {
let builder = Request::builder().method(Method::POST).uri("http://127.0.0.1:1/");
let _ = StreamableHttpTransport::apply_v2_outbound_headers(builder, &m, &n);
}
}
}
}
mod v2_error_envelope {
use super::*;
use crate::types::jsonrpc::ResponsePayload;
use crate::types::protocol::PROTOCOL_VERSION_2026_07_28;
const INVALID_PARAMS_BODY: &str = r#"{"jsonrpc":"2.0","id":"abc","error":{"code":-32602,"message":"requestState could not be accepted"}}"#;
fn transport_for(url: &str, v2: bool) -> StreamableHttpTransport {
let config =
StreamableHttpTransportConfigBuilder::new(Url::parse(url).unwrap()).build();
let mut transport = StreamableHttpTransport::new(config);
if v2 {
transport
.set_negotiated_protocol_version(Some(PROTOCOL_VERSION_2026_07_28.to_string()));
}
transport
}
#[tokio::test]
async fn v2_surfaces_a_jsonrpc_error_carried_on_a_400() {
let mut server = MockServer::new_async().await;
let mock = server
.mock("POST", "/")
.with_status(400)
.with_header("content-type", "application/json")
.with_body(INVALID_PARAMS_BODY)
.create_async()
.await;
let mut transport = transport_for(&server.url(), true);
transport
.send_raw(br#"{"jsonrpc":"2.0","id":"abc","method":"tools/call","params":{"name":"x","arguments":{}}}"#.to_vec())
.await
.expect("a structured error must NOT be a transport failure");
let message = transport
.receive()
.await
.expect("the envelope is delivered");
let TransportMessage::Response(response) = message else {
panic!("expected a response, got {message:?}");
};
let ResponsePayload::Error(error) = response.payload else {
panic!("expected the error payload");
};
assert_eq!(error.code, -32602);
mock.assert_async().await;
}
#[tokio::test]
async fn v2_falls_back_to_the_status_error_for_a_non_envelope_body() {
let mut server = MockServer::new_async().await;
let _mock = server
.mock("POST", "/")
.with_status(502)
.with_header("content-type", "text/html")
.with_body("<html>bad gateway</html>")
.create_async()
.await;
let mut transport = transport_for(&server.url(), true);
let error = transport
.send_raw(br#"{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}"#.to_vec())
.await
.expect_err("a proxy error page is still a transport failure");
assert!(
error.to_string().contains("502"),
"the status must survive: {error}"
);
}
#[tokio::test]
async fn v1_still_errors_on_the_status_alone() {
let mut server = MockServer::new_async().await;
let _mock = server
.mock("POST", "/")
.with_status(400)
.with_header("content-type", "application/json")
.with_body(INVALID_PARAMS_BODY)
.create_async()
.await;
let mut transport = transport_for(&server.url(), false);
let error = transport
.send(list_tools_message())
.await
.expect_err("v1 behavior must be byte-identical to prior releases");
assert!(
error.to_string().contains("400"),
"v1 must still report the status: {error}"
);
}
}
mod collected_body_cap {
use super::*;
use std::time::Duration;
const CAP: usize = 512;
const RESPONSE_JSON: &str = r#"{"jsonrpc":"2.0","id":42,"result":{"tools":[]}}"#;
const QUIET_WINDOW: Duration = Duration::from_millis(250);
fn capped_transport(url: &str, cap: usize) -> StreamableHttpTransport {
let config =
StreamableHttpTransportConfigBuilder::new(Url::parse(url).unwrap()).build();
StreamableHttpTransport::new(config).with_max_collected_body_bytes(cap)
}
fn sse_body_of(len: usize) -> String {
let frame = format!("event: message\ndata: {RESPONSE_JSON}\n\n");
let padding = len
.checked_sub(frame.len())
.expect("requested length must fit one frame");
let comment = match padding {
0 => String::new(),
1 => panic!("a comment line costs at least two bytes"),
n => format!(":{}\n", "p".repeat(n - 2)),
};
let body = format!("{comment}{frame}");
assert_eq!(body.len(), len, "the body must be exactly {len} bytes");
body
}
fn assert_over_cap_refusal(error: &Error, cap: usize) {
let text = error.to_string();
assert!(
text.contains(&cap.to_string()),
"the refusal must NAME the limit: {text}"
);
assert!(
!text.contains("jsonrpc") && !text.contains("pppppppp"),
"the refusal must not echo body content: {text}"
);
}
#[tokio::test]
async fn post_response_one_byte_over_the_cap_is_refused_before_the_parser() {
let mut server = MockServer::new_async().await;
let body = sse_body_of(CAP + 1);
let _mock = server
.mock("POST", "/")
.with_status(200)
.with_header("content-type", TEXT_EVENT_STREAM)
.with_chunked_body(move |w| w.write_all(body.as_bytes()))
.create_async()
.await;
let mut transport = capped_transport(&server.url(), CAP);
let error = transport
.send(list_tools_message())
.await
.expect_err("a body over the cap must be refused");
assert_over_cap_refusal(&error, CAP);
assert!(
tokio::time::timeout(QUIET_WINDOW, transport.receive())
.await
.is_err(),
"an over-cap body must never reach the parser, so nothing can be dispatched"
);
}
#[tokio::test]
async fn post_response_at_the_cap_parses_normally() {
let mut server = MockServer::new_async().await;
let body = sse_body_of(CAP);
let _mock = server
.mock("POST", "/")
.with_status(200)
.with_header("content-type", TEXT_EVENT_STREAM)
.with_chunked_body(move |w| w.write_all(body.as_bytes()))
.create_async()
.await;
let mut transport = capped_transport(&server.url(), CAP);
transport
.send(list_tools_message())
.await
.expect("a body at the cap must be accepted");
let message = tokio::time::timeout(QUIET_WINDOW, transport.receive())
.await
.expect("the parsed event must be dispatched")
.expect("the parsed event must be a message");
assert!(
matches!(message, TransportMessage::Response(_)),
"expected the parsed response, got {message:?}"
);
}
#[tokio::test]
async fn start_sse_one_byte_over_the_cap_is_refused_before_the_parser() {
let mut server = MockServer::new_async().await;
let body = sse_body_of(CAP + 1);
let _mock = server
.mock("GET", "/")
.with_status(200)
.with_header("content-type", TEXT_EVENT_STREAM)
.with_chunked_body(move |w| w.write_all(body.as_bytes()))
.create_async()
.await;
let mut transport = capped_transport(&server.url(), CAP);
let error = transport
.start_sse(None)
.await
.expect_err("a body over the cap must be refused");
assert_over_cap_refusal(&error, CAP);
assert!(
tokio::time::timeout(QUIET_WINDOW, transport.receive())
.await
.is_err(),
"an over-cap body must never reach the parser, so nothing can be dispatched"
);
}
#[tokio::test]
async fn start_sse_at_the_cap_parses_normally() {
let mut server = MockServer::new_async().await;
let body = sse_body_of(CAP);
let _mock = server
.mock("GET", "/")
.with_status(200)
.with_header("content-type", TEXT_EVENT_STREAM)
.with_chunked_body(move |w| w.write_all(body.as_bytes()))
.create_async()
.await;
let mut transport = capped_transport(&server.url(), CAP);
transport
.start_sse(None)
.await
.expect("a body at the cap must be accepted");
let message = tokio::time::timeout(QUIET_WINDOW, transport.receive())
.await
.expect("the parsed event must be dispatched")
.expect("the parsed event must be a message");
assert!(
matches!(message, TransportMessage::Response(_)),
"expected the parsed response, got {message:?}"
);
}
#[tokio::test]
async fn a_declared_content_length_over_the_cap_is_refused_early() {
let mut server = MockServer::new_async().await;
let _mock = server
.mock("POST", "/")
.with_status(200)
.with_header("content-type", TEXT_EVENT_STREAM)
.with_body(sse_body_of(CAP + 1))
.create_async()
.await;
let mut transport = capped_transport(&server.url(), CAP);
let error = transport
.send(list_tools_message())
.await
.expect_err("an over-cap Content-Length must be refused");
assert_over_cap_refusal(&error, CAP);
assert!(
error.to_string().contains(&(CAP + 1).to_string()),
"the early refusal must name the DECLARED size: {error}"
);
}
#[tokio::test]
async fn raising_the_cap_admits_a_body_the_lower_one_refuses() {
let mut server = MockServer::new_async().await;
let body = sse_body_of(CAP + 1);
let _mock = server
.mock("POST", "/")
.with_status(200)
.with_header("content-type", TEXT_EVENT_STREAM)
.with_chunked_body(move |w| w.write_all(body.as_bytes()))
.expect_at_least(1)
.create_async()
.await;
let mut transport = capped_transport(&server.url(), CAP * 4);
transport
.send(list_tools_message())
.await
.expect("the raised cap must admit the body the lower one refused");
let message = tokio::time::timeout(QUIET_WINDOW, transport.receive())
.await
.expect("the parsed event must be dispatched")
.expect("the parsed event must be a message");
assert!(
matches!(message, TransportMessage::Response(_)),
"expected the parsed response, got {message:?}"
);
}
#[test]
fn every_constructor_defaults_the_cap_to_the_named_constant() {
let url = Url::parse("http://127.0.0.1:1/").unwrap();
let config = StreamableHttpTransportConfigBuilder::new(url).build();
assert_eq!(
StreamableHttpTransport::new(config.clone()).max_collected_body_bytes,
DEFAULT_MAX_COLLECTED_BODY_BYTES,
"`new` must default from the named constant"
);
assert_eq!(
StreamableHttpTransport::new_with_http2(config.clone()).max_collected_body_bytes,
DEFAULT_MAX_COLLECTED_BODY_BYTES,
"`new_with_http2` must default from the named constant"
);
assert_eq!(
StreamableHttpTransport::new(config)
.with_max_collected_body_bytes(CAP)
.max_collected_body_bytes,
CAP,
"the builder must override the default"
);
}
#[tokio::test]
async fn an_over_cap_v2_error_envelope_falls_back_to_the_status_error() {
let padding = "z".repeat(CAP);
let body = format!(
r#"{{"jsonrpc":"2.0","id":"abc","error":{{"code":-32602,"message":"{padding}"}}}}"#
);
assert!(body.len() > CAP);
let mut server = MockServer::new_async().await;
let _mock = server
.mock("POST", "/")
.with_status(400)
.with_header("content-type", APPLICATION_JSON)
.with_chunked_body(move |w| w.write_all(body.as_bytes()))
.create_async()
.await;
let mut transport = capped_transport(&server.url(), CAP);
transport.set_negotiated_protocol_version(Some(
crate::types::protocol::PROTOCOL_VERSION_2026_07_28.to_string(),
));
let error = transport
.send_raw(
br#"{"jsonrpc":"2.0","id":"abc","method":"tools/list","params":{}}"#.to_vec(),
)
.await
.expect_err("an over-cap envelope cannot be surfaced structurally");
assert!(
error.to_string().contains("400"),
"the status must survive: {error}"
);
}
}
}