use std::{
borrow::Cow,
collections::{HashMap, HashSet, VecDeque},
sync::Arc,
time::Duration,
};
use futures::{
Stream, StreamExt,
future::BoxFuture,
stream::{BoxStream, FuturesUnordered},
};
use http::{HeaderName, HeaderValue, StatusCode};
pub use sse_stream::Error as SseError;
use sse_stream::Sse;
use thiserror::Error;
use tokio_util::sync::CancellationToken;
use tracing::debug;
use super::common::client_side_sse::{
DEFAULT_MAX_SSE_EVENT_SIZE, ExponentialBackoff, SseRetryPolicy, SseStreamReconnect,
};
use crate::{
RoleClient,
model::{
ClientJsonRpcMessage, ClientNotification, ClientRequest, ErrorData, GetExtensions, GetMeta,
InitializedNotification, JsonObject, ProtocolVersion, RequestId, ServerJsonRpcMessage,
ServerResult,
},
service::InboundStreamOrigin,
transport::{
common::{client_side_sse::SseAutoReconnectStream, mcp_headers},
worker::{
RequestCancellationRegistration, Worker, WorkerQuitReason, WorkerSendRequest,
WorkerTransport,
},
},
};
type BoxedSseStream = BoxStream<'static, Result<Sse, SseError>>;
type SseTaskResult<E> = (Option<RequestId>, Result<(), StreamableHttpError<E>>);
const SESSION_CLEANUP_TIMEOUT: Duration = Duration::from_secs(5);
fn build_request_headers(
base: &HashMap<HeaderName, HeaderValue>,
message: &ClientJsonRpcMessage,
tool_cache: &HashMap<String, Arc<JsonObject>>,
version: &ProtocolVersion,
) -> HashMap<HeaderName, HeaderValue> {
use serde_json::Value;
let mut headers = base.clone();
if *version >= ProtocolVersion::STANDARD_HEADERS
&& let Ok(value) = serde_json::to_value(message)
{
let schema = value
.get("method")
.and_then(Value::as_str)
.filter(|method| *method == "tools/call")
.and_then(|_| value.get("params"))
.and_then(|params| params.get("name"))
.and_then(Value::as_str)
.and_then(|name| tool_cache.get(name))
.map(Arc::as_ref);
for (name, val) in mcp_headers::standard_request_headers(&value, schema) {
headers.insert(name, val);
}
}
headers
}
fn request_version_headers(
base: &HashMap<HeaderName, HeaderValue>,
message: &ClientJsonRpcMessage,
fallback: &ProtocolVersion,
tool_cache: &HashMap<String, Arc<JsonObject>>,
) -> (ProtocolVersion, HashMap<HeaderName, HeaderValue>) {
let version = match message {
ClientJsonRpcMessage::Request(request) => request
.request
.get_meta()
.protocol_version()
.unwrap_or_else(|| fallback.clone()),
_ => fallback.clone(),
};
let mut headers = build_request_headers(base, message, tool_cache, &version);
if let Ok(value) = HeaderValue::from_str(version.as_str()) {
headers.insert(HeaderName::from_static("mcp-protocol-version"), value);
}
(version, headers)
}
fn cache_tools_from_response(
cache: &mut HashMap<String, Arc<JsonObject>>,
message: &mut ServerJsonRpcMessage,
protocol_version: &ProtocolVersion,
) {
if protocol_version < &ProtocolVersion::STANDARD_HEADERS {
return;
}
if let ServerJsonRpcMessage::Response(response) = message
&& let ServerResult::ListToolsResult(list) = &mut response.result
{
list.tools.retain(|tool| {
let Err(reason) =
mcp_headers::validate_param_header_annotations(&tool.input_schema)
else {
cache.insert(tool.name.to_string(), tool.input_schema.clone());
return true;
};
tracing::warn!(tool = %tool.name, "rejecting invalid x-mcp-header annotations: {reason}");
false
});
}
}
fn negotiate_version_headers(
init_response: &ServerJsonRpcMessage,
base: HashMap<HeaderName, HeaderValue>,
) -> (ProtocolVersion, HashMap<HeaderName, HeaderValue>) {
let mut version = ProtocolVersion::default();
let mut headers = base;
if let ServerJsonRpcMessage::Response(response) = init_response
&& let ServerResult::InitializeResult(init_result) = &response.result
{
version = init_result.protocol_version.clone();
if let Ok(hv) = HeaderValue::from_str(init_result.protocol_version.as_str()) {
headers.insert(HeaderName::from_static("mcp-protocol-version"), hv);
}
}
(version, headers)
}
#[derive(Debug, Error)]
#[error("authorization required: {www_authenticate_header}")]
#[non_exhaustive]
pub struct AuthRequiredError {
pub www_authenticate_header: String,
}
impl AuthRequiredError {
pub fn new(www_authenticate_header: String) -> Self {
Self {
www_authenticate_header,
}
}
}
#[derive(Debug, Error)]
#[error("insufficient scope: {www_authenticate_header}")]
#[non_exhaustive]
pub struct InsufficientScopeError {
pub www_authenticate_header: String,
pub required_scope: Option<String>,
}
impl InsufficientScopeError {
pub fn new(www_authenticate_header: String, required_scope: Option<String>) -> Self {
Self {
www_authenticate_header,
required_scope,
}
}
pub fn can_upgrade(&self) -> bool {
self.required_scope.is_some()
}
pub fn get_required_scope(&self) -> Option<&str> {
self.required_scope.as_deref()
}
}
#[derive(Error, Debug)]
#[non_exhaustive]
pub enum StreamableHttpError<E: std::error::Error + Send + Sync + 'static> {
#[error("SSE error: {0}")]
Sse(#[from] SseError),
#[error("Io error: {0}")]
Io(#[from] std::io::Error),
#[error("Client error: {0}")]
Client(E),
#[error("unexpected end of stream")]
UnexpectedEndOfStream,
#[error("unexpected server response: {0}")]
UnexpectedServerResponse(Cow<'static, str>),
#[error("Unexpected content type: {0:?}")]
UnexpectedContentType(Option<String>),
#[error("Server does not support SSE")]
ServerDoesNotSupportSse,
#[error("Server does not support delete session")]
ServerDoesNotSupportDeleteSession,
#[error("Tokio join error: {0}")]
TokioJoinError(#[from] tokio::task::JoinError),
#[error("Deserialize error: {0}")]
Deserialize(#[from] serde_json::Error),
#[error("Transport channel closed")]
TransportChannelClosed,
#[error("Missing session id in HTTP response")]
MissingSessionIdInResponse,
#[cfg(feature = "auth")]
#[error("Auth error: {0}")]
Auth(#[from] crate::transport::auth::AuthError),
#[error("Auth required")]
AuthRequired(#[source] AuthRequiredError),
#[error("Insufficient scope")]
InsufficientScope(#[source] InsufficientScopeError),
#[error("Header name '{0}' is reserved and conflicts with default headers")]
ReservedHeaderConflict(String),
#[error("Session expired (HTTP 404)")]
SessionExpired,
#[error("Session recovery timed out; the server may have processed the POST")]
SessionRecoveryTimeout,
#[error("Control POST timed out")]
ControlRequestTimeout,
}
impl<E: std::error::Error + Send + Sync + 'static> StreamableHttpError<E> {
#[cfg(feature = "auth")]
pub fn auth_challenge(&self) -> Option<&str> {
match self {
Self::AuthRequired(error) => Some(&error.www_authenticate_header),
Self::InsufficientScope(error) => Some(&error.www_authenticate_header),
_ => None,
}
}
}
#[derive(Debug, Clone, Error)]
#[non_exhaustive]
pub enum StreamableHttpProtocolError {
#[error("Missing session id in response")]
MissingSessionIdInResponse,
}
#[expect(
clippy::large_enum_variant,
reason = "boxing the streaming response would add an allocation to the common response path"
)]
#[non_exhaustive]
pub enum StreamableHttpPostResponse {
Accepted,
Json(ServerJsonRpcMessage, Option<String>),
Sse(BoxedSseStream, Option<String>),
}
impl std::fmt::Debug for StreamableHttpPostResponse {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Accepted => write!(f, "Accepted"),
Self::Json(arg0, arg1) => f.debug_tuple("Json").field(arg0).field(arg1).finish(),
Self::Sse(_, arg1) => f.debug_tuple("Sse").field(arg1).finish(),
}
}
}
impl StreamableHttpPostResponse {
pub async fn expect_initialized<E>(
self,
) -> Result<(ServerJsonRpcMessage, Option<String>), StreamableHttpError<E>>
where
E: std::error::Error + Send + Sync + 'static,
{
match self {
Self::Json(message, session_id) => Ok((message, session_id)),
Self::Sse(mut stream, session_id) => {
while let Some(event) = stream.next().await {
let event = event?;
let payload = event.data.unwrap_or_default();
if payload.trim().is_empty() {
continue;
}
let message: ServerJsonRpcMessage = serde_json::from_str(&payload)?;
if matches!(message, ServerJsonRpcMessage::Response(_)) {
return Ok((message, session_id));
}
debug!(
?message,
"received message before initialize response; continuing to drain stream"
);
}
Err(StreamableHttpError::UnexpectedServerResponse(
"empty sse stream".into(),
))
}
_ => Err(StreamableHttpError::UnexpectedServerResponse(
"expect initialized, accepted".into(),
)),
}
}
pub fn expect_json<E>(self) -> Result<ServerJsonRpcMessage, StreamableHttpError<E>>
where
E: std::error::Error + Send + Sync + 'static,
{
match self {
Self::Json(message, ..) => Ok(message),
got => Err(StreamableHttpError::UnexpectedServerResponse(
format!("expect json, got {got:?}").into(),
)),
}
}
pub fn expect_accepted_or_json<E>(self) -> Result<(), StreamableHttpError<E>>
where
E: std::error::Error + Send + Sync + 'static,
{
match self {
Self::Accepted => Ok(()),
Self::Json(..) => Ok(()),
got => Err(StreamableHttpError::UnexpectedServerResponse(
format!("expect accepted or json, got {got:?}").into(),
)),
}
}
}
pub(super) fn legacy_discover_response(
message: &ClientJsonRpcMessage,
session_was_attached: bool,
status: StatusCode,
body: &str,
) -> Option<StreamableHttpPostResponse> {
if session_was_attached
|| !status.is_client_error()
|| matches!(status, StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN)
{
return None;
}
let ClientJsonRpcMessage::Request(request) = message else {
return None;
};
if !matches!(request.request, ClientRequest::DiscoverRequest(_)) {
return None;
}
let error = ErrorData::invalid_request(
format!("server/discover rejected with HTTP {status}: {body}"),
None,
);
Some(StreamableHttpPostResponse::Json(
ServerJsonRpcMessage::error(error, Some(request.id.clone())),
None,
))
}
pub trait StreamableHttpClient: Clone + Send + 'static {
type Error: std::error::Error + Send + Sync + 'static;
fn post_message(
&self,
uri: Arc<str>,
message: ClientJsonRpcMessage,
session_id: Option<Arc<str>>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
) -> impl Future<Output = Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>>>
+ Send
+ '_;
fn post_message_with_max_sse_event_size(
&self,
uri: Arc<str>,
message: ClientJsonRpcMessage,
session_id: Option<Arc<str>>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
_max_sse_event_size: usize,
) -> impl Future<Output = Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>>>
+ Send
+ '_ {
self.post_message(uri, message, session_id, auth_header, custom_headers)
}
fn delete_session(
&self,
uri: Arc<str>,
session_id: Arc<str>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
) -> impl Future<Output = Result<(), StreamableHttpError<Self::Error>>> + Send + '_;
fn get_stream(
&self,
uri: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
) -> impl Future<
Output = Result<
BoxStream<'static, Result<Sse, SseError>>,
StreamableHttpError<Self::Error>,
>,
> + Send
+ '_;
fn get_stream_with_max_sse_event_size(
&self,
uri: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
_max_sse_event_size: usize,
) -> impl Future<
Output = Result<
BoxStream<'static, Result<Sse, SseError>>,
StreamableHttpError<Self::Error>,
>,
> + Send
+ '_ {
self.get_stream(uri, session_id, last_event_id, auth_header, custom_headers)
}
}
#[non_exhaustive]
pub struct RetryConfig {
pub max_times: Option<usize>,
pub min_duration: Duration,
}
struct StreamableHttpClientReconnect<C> {
pub client: C,
pub session_id: Option<Arc<str>>,
pub uri: Arc<str>,
pub auth_header: Option<String>,
pub custom_headers: HashMap<HeaderName, HeaderValue>,
pub max_sse_event_size: usize,
}
impl<C: StreamableHttpClient> SseStreamReconnect for StreamableHttpClientReconnect<C> {
type Error = StreamableHttpError<C::Error>;
type Future = BoxFuture<'static, Result<BoxedSseStream, Self::Error>>;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future {
let client = self.client.clone();
let uri = self.uri.clone();
let session_id = self.session_id.clone();
let auth_header = self.auth_header.clone();
let custom_headers = self.custom_headers.clone();
let max_sse_event_size = self.max_sse_event_size;
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move {
client
.get_stream_with_max_sse_event_size(
uri,
session_id,
last_event_id,
auth_header,
custom_headers,
max_sse_event_size,
)
.await
})
}
fn map_fatal_stream_error(&mut self, error: SseError) -> Option<Self::Error> {
Some(StreamableHttpError::Sse(error))
}
}
struct SessionCleanupInfo<C> {
client: C,
uri: Arc<str>,
session_id: Arc<str>,
auth_header: Option<String>,
protocol_headers: HashMap<HeaderName, HeaderValue>,
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct StreamableHttpClientWorker<C: StreamableHttpClient> {
pub client: C,
pub config: StreamableHttpClientTransportConfig,
}
struct PostResult<C: StreamableHttpClient> {
send_request: WorkerSendRequest<StreamableHttpClientWorker<C>>,
response: Option<Result<StreamableHttpPostResponse, StreamableHttpError<C::Error>>>,
version: ProtocolVersion,
}
struct PostSession {
id: Option<Arc<str>>,
headers: HashMap<HeaderName, HeaderValue>,
version: ProtocolVersion,
cancellation: CancellationToken,
}
impl<C: StreamableHttpClient + Default> StreamableHttpClientWorker<C> {
pub fn new_simple(url: impl Into<Arc<str>>) -> Self {
Self {
client: C::default(),
config: StreamableHttpClientTransportConfig {
uri: url.into(),
..Default::default()
},
}
}
}
impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
pub fn new(client: C, config: StreamableHttpClientTransportConfig) -> Self {
Self { client, config }
}
}
impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
fn is_ordering_barrier(
message: &ClientJsonRpcMessage,
negotiated_version: &ProtocolVersion,
) -> bool {
match message {
ClientJsonRpcMessage::Request(request) => {
matches!(
&request.request,
ClientRequest::InitializeRequest(_) | ClientRequest::DiscoverRequest(_)
) || request
.request
.get_meta()
.protocol_version()
.is_some_and(|version| &version != negotiated_version)
}
ClientJsonRpcMessage::Notification(notification) => matches!(
¬ification.notification,
ClientNotification::InitializedNotification(_)
),
_ => false,
}
}
fn post_request(
client: C,
config: &StreamableHttpClientTransportConfig,
mut send_request: WorkerSendRequest<Self>,
session: PostSession,
transport_cancellation: CancellationToken,
) -> BoxFuture<'static, PostResult<C>> {
let uri = config.uri.clone();
let auth_header = config.auth_header.clone();
let max_sse_event_size = config.max_sse_event_size;
let control_request_timeout = config.control_request_timeout;
let is_control = Self::is_control_message(&send_request.message);
let cancellation = send_request
.cancellation_token()
.unwrap_or_else(|| transport_cancellation.child_token());
Box::pin(async move {
let response = tokio::select! {
biased;
_ = cancellation.cancelled() => None,
_ = send_request.responder.closed() => None,
_ = session.cancellation.cancelled() => {
Some(Err(StreamableHttpError::SessionRecoveryTimeout))
},
_ = tokio::time::sleep(control_request_timeout), if is_control => {
Some(Err(StreamableHttpError::ControlRequestTimeout))
},
response = client.post_message_with_max_sse_event_size(
uri,
send_request.message.clone(),
session.id,
auth_header,
session.headers,
max_sse_event_size,
) => Some(response),
};
PostResult {
send_request,
response,
version: session.version,
}
})
}
fn cancellation_request_id(message: &ClientJsonRpcMessage) -> Option<&RequestId> {
match message {
ClientJsonRpcMessage::Notification(notification) => match ¬ification.notification {
ClientNotification::CancelledNotification(cancelled) => {
cancelled.params.request_id.as_ref()
}
_ => None,
},
_ => None,
}
}
fn client_request_id(message: &ClientJsonRpcMessage) -> Option<RequestId> {
match message {
ClientJsonRpcMessage::Request(request) => Some(request.id.clone()),
_ => None,
}
}
fn server_response_id(message: &ServerJsonRpcMessage) -> Option<&RequestId> {
match message {
ServerJsonRpcMessage::Response(response) => Some(&response.id),
ServerJsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
}
}
fn clear_stream_response_pending(
pending_stream_response_ids: &mut HashSet<RequestId>,
message: &ServerJsonRpcMessage,
) -> Option<RequestId> {
let response_id = Self::server_response_id(message)?;
if let Some(id) = pending_stream_response_ids.take(response_id) {
return Some(id);
}
let id = RequestId::Number(response_id.numeric_string_value()?);
pending_stream_response_ids.take(&id)
}
async fn drain_queued_stream_messages(
sse_worker_rx: &mut tokio::sync::mpsc::Receiver<ServerJsonRpcMessage>,
context: &mut super::worker::WorkerContext<Self>,
pending_stream_response_ids: &mut HashSet<RequestId>,
) -> Result<(), WorkerQuitReason<StreamableHttpError<C::Error>>> {
loop {
match sse_worker_rx.try_recv() {
Ok(message) => {
let _ =
Self::clear_stream_response_pending(pending_stream_response_ids, &message);
context.send_to_handler(message).await?;
}
Err(tokio::sync::mpsc::error::TryRecvError::Empty) => return Ok(()),
Err(tokio::sync::mpsc::error::TryRecvError::Disconnected) => return Ok(()),
}
}
}
async fn fail_pending_stream_responses(
context: &mut super::worker::WorkerContext<Self>,
pending_stream_response_ids: &mut HashSet<RequestId>,
) -> Result<(), WorkerQuitReason<StreamableHttpError<C::Error>>> {
if pending_stream_response_ids.is_empty() {
return Ok(());
}
let pending_ids = std::mem::take(pending_stream_response_ids);
for id in pending_ids {
context
.send_to_handler(ServerJsonRpcMessage::error(
ErrorData::internal_error(
"streamable HTTP session was re-initialized before the response arrived",
None,
),
Some(id),
))
.await?;
}
Ok(())
}
async fn fail_pending_responses_except_retries(
context: &mut super::worker::WorkerContext<Self>,
pending_stream_response_ids: &mut HashSet<RequestId>,
recovery_posts: &VecDeque<WorkerSendRequest<Self>>,
) -> Result<(), WorkerQuitReason<StreamableHttpError<C::Error>>> {
let retry_ids: Vec<_> = recovery_posts
.iter()
.filter_map(|request| Self::client_request_id(&request.message))
.filter(|id| pending_stream_response_ids.remove(id))
.collect();
Self::fail_pending_stream_responses(context, pending_stream_response_ids).await?;
pending_stream_response_ids.extend(retry_ids);
Ok(())
}
fn fail_recovery_posts(
recovery_posts: &mut VecDeque<WorkerSendRequest<Self>>,
pending_stream_response_ids: &mut HashSet<RequestId>,
error: StreamableHttpError<C::Error>,
) {
let mut recovery_error = Some(error);
for send_request in recovery_posts.drain(..) {
let pending = Self::client_request_id(&send_request.message)
.is_none_or(|id| pending_stream_response_ids.remove(&id));
let result = if pending {
Err(recovery_error
.take()
.unwrap_or(StreamableHttpError::SessionExpired))
} else {
Ok(())
};
let _ = send_request.responder.send(result);
}
}
fn reconnecting_sse_to_jsonrpc(
stream: BoxedSseStream,
client: C,
session_id: Option<Arc<str>>,
uri: Arc<str>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
max_sse_event_size: usize,
retry_config: Arc<dyn SseRetryPolicy>,
) -> impl Stream<Item = Result<ServerJsonRpcMessage, StreamableHttpError<C::Error>>> + Send + 'static
{
SseAutoReconnectStream::new_after_event_id(
stream,
StreamableHttpClientReconnect {
client,
session_id,
uri,
auth_header,
custom_headers,
max_sse_event_size,
},
retry_config,
)
}
fn response_sse_to_jsonrpc(
stream: BoxedSseStream,
session_id: Option<Arc<str>>,
client: C,
uri: Arc<str>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
max_sse_event_size: usize,
retry_config: Arc<dyn SseRetryPolicy>,
) -> BoxStream<'static, Result<ServerJsonRpcMessage, StreamableHttpError<C::Error>>> {
Self::reconnecting_sse_to_jsonrpc(
stream,
client,
session_id,
uri,
auth_header,
custom_headers,
max_sse_event_size,
retry_config,
)
.boxed()
}
async fn run_response_stream(
mut sse_stream: BoxStream<
'static,
Result<ServerJsonRpcMessage, StreamableHttpError<C::Error>>,
>,
sse_worker_tx: tokio::sync::mpsc::Sender<ServerJsonRpcMessage>,
origin: InboundStreamOrigin,
request_ct: CancellationToken,
stream_ct: CancellationToken,
uses_modern_http: bool,
) -> Result<(), StreamableHttpError<C::Error>> {
tokio::select! {
biased;
_ = request_ct.cancelled(), if !uses_modern_http => {
stream_ct.cancelled().await;
Ok(())
}
result = Self::execute_sse_stream(
sse_stream.as_mut(), sse_worker_tx, origin, true, stream_ct.clone(),
) => result,
}
}
async fn execute_sse_stream(
sse_stream: impl Stream<Item = Result<ServerJsonRpcMessage, StreamableHttpError<C::Error>>>
+ Send,
sse_worker_tx: tokio::sync::mpsc::Sender<ServerJsonRpcMessage>,
origin: InboundStreamOrigin,
close_on_response: bool,
ct: CancellationToken,
) -> Result<(), StreamableHttpError<C::Error>> {
let mut sse_stream = std::pin::pin!(sse_stream);
loop {
let message = tokio::select! {
event = sse_stream.next() => {
event
}
_ = ct.cancelled() => {
tracing::debug!("cancelled");
break;
}
};
let Some(mut message) = message.transpose()? else {
break;
};
if let ServerJsonRpcMessage::Request(request) = &mut message {
request.request.extensions_mut().insert(origin.clone());
}
let is_response = matches!(
message,
ServerJsonRpcMessage::Response(_) | ServerJsonRpcMessage::Error(_)
);
let yield_result = sse_worker_tx.send(message).await;
if yield_result.is_err() {
tracing::trace!("streamable http transport worker dropped, exiting");
break;
}
if close_on_response && is_response {
tracing::debug!("got response, draining sse stream for connection reuse");
let _ = tokio::time::timeout(std::time::Duration::from_millis(50), async {
while sse_stream.next().await.is_some() {}
})
.await;
break;
}
}
Ok(())
}
fn spawn_common_stream(
streams: &mut tokio::task::JoinSet<SseTaskResult<C::Error>>,
client: C,
session_id: Arc<str>,
config: &StreamableHttpClientTransportConfig,
protocol_headers: HashMap<HeaderName, HeaderValue>,
sse_worker_tx: tokio::sync::mpsc::Sender<ServerJsonRpcMessage>,
transport_task_ct: CancellationToken,
) {
let uri = config.uri.clone();
let auth_header = config.auth_header.clone();
let retry_config = config.retry_config.clone();
let reconnect_uri = config.uri.clone();
let reconnect_auth_header = config.auth_header.clone();
let max_sse_event_size = config.max_sse_event_size;
streams.spawn(async move {
let result = match client
.get_stream_with_max_sse_event_size(
uri,
Some(session_id.clone()),
None,
auth_header,
protocol_headers.clone(),
max_sse_event_size,
)
.await
{
Ok(stream) => {
let sse_stream = SseAutoReconnectStream::new(
stream,
StreamableHttpClientReconnect {
client,
session_id: Some(session_id),
uri: reconnect_uri,
auth_header: reconnect_auth_header,
custom_headers: protocol_headers,
max_sse_event_size,
},
retry_config,
);
Self::execute_sse_stream(
sse_stream,
sse_worker_tx,
InboundStreamOrigin::Unassociated,
false,
transport_task_ct.child_token(),
)
.await
}
Err(StreamableHttpError::ServerDoesNotSupportSse) => {
tracing::debug!("server doesn't support sse, skip common stream");
Ok(())
}
Err(error) => {
tracing::error!("fail to get common stream: {error}");
Err(error)
}
};
(None, result)
});
}
async fn perform_reinitialization(
client: C,
saved_init_request: ClientJsonRpcMessage,
uri: Arc<str>,
auth_header: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
max_sse_event_size: usize,
) -> Result<
(
Option<Arc<str>>,
ProtocolVersion,
HashMap<HeaderName, HeaderValue>,
),
StreamableHttpError<C::Error>,
> {
let (init_msg, new_session_id_str) = client
.post_message_with_max_sse_event_size(
uri.clone(),
saved_init_request,
None,
auth_header.clone(),
custom_headers.clone(),
max_sse_event_size,
)
.await?
.expect_initialized::<C::Error>()
.await?;
let new_session_id: Option<Arc<str>> = new_session_id_str.map(|s| Arc::from(s.as_str()));
let (negotiated_version, new_protocol_headers) =
negotiate_version_headers(&init_msg, custom_headers);
let initialized_notification = ClientJsonRpcMessage::notification(
ClientNotification::InitializedNotification(InitializedNotification {
method: Default::default(),
extensions: Default::default(),
}),
);
let initialized_headers = build_request_headers(
&new_protocol_headers,
&initialized_notification,
&HashMap::new(),
&negotiated_version,
);
client
.post_message_with_max_sse_event_size(
uri,
initialized_notification,
new_session_id.clone(),
auth_header,
initialized_headers,
max_sse_event_size,
)
.await?
.expect_accepted_or_json::<C::Error>()?;
Ok((new_session_id, negotiated_version, new_protocol_headers))
}
}
impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
type Role = RoleClient;
type Error = StreamableHttpError<C::Error>;
fn is_control_message(message: &ClientJsonRpcMessage) -> bool {
match message {
ClientJsonRpcMessage::Response(_) | ClientJsonRpcMessage::Error(_) => true,
ClientJsonRpcMessage::Notification(notification) => matches!(
notification.notification,
ClientNotification::CancelledNotification(_)
),
ClientJsonRpcMessage::Request(_) => false,
}
}
fn supports_request_cancellation() -> bool {
true
}
fn err_closed() -> Self::Error {
StreamableHttpError::TransportChannelClosed
}
fn err_join(e: tokio::task::JoinError) -> Self::Error {
StreamableHttpError::TokioJoinError(e)
}
fn config(&self) -> super::worker::WorkerConfig {
super::worker::WorkerConfig {
name: Some("StreamableHttpClientWorker".into()),
channel_buffer_capacity: self.config.channel_buffer_capacity,
}
}
async fn run(
self,
mut context: super::worker::WorkerContext<Self>,
) -> Result<(), WorkerQuitReason<Self::Error>> {
let channel_buffer_capacity = self.config.channel_buffer_capacity;
let (sse_worker_tx, mut sse_worker_rx) =
tokio::sync::mpsc::channel::<ServerJsonRpcMessage>(channel_buffer_capacity);
let config = self.config.clone();
let transport_task_ct = context.cancellation_token.clone();
let _drop_guard = transport_task_ct.clone().drop_guard();
let WorkerSendRequest {
responder,
message: startup_request,
..
} = context.recv_from_handler().await?;
let is_legacy_startup = matches!(
&startup_request,
ClientJsonRpcMessage::Request(request)
if matches!(&request.request, ClientRequest::InitializeRequest(_))
);
let mut saved_init_request = is_legacy_startup.then(|| startup_request.clone());
let empty_tool_cache = HashMap::new();
let (bootstrap_version, bootstrap_headers) = if is_legacy_startup {
(ProtocolVersion::default(), config.custom_headers.clone())
} else {
request_version_headers(
&config.custom_headers,
&startup_request,
&ProtocolVersion::default(),
&empty_tool_cache,
)
};
let (message, session_id) = match self
.client
.post_message_with_max_sse_event_size(
config.uri.clone(),
startup_request,
None,
config.auth_header.clone(),
bootstrap_headers.clone(),
config.max_sse_event_size,
)
.await
{
Ok(res) => {
let _ = responder.send(Ok(()));
res.expect_initialized::<C::Error>().await.map_err(
WorkerQuitReason::fatal_context("process initialize response"),
)?
}
Err(err) => {
let msg = format!("{:?}", err);
let _ = responder.send(Err(err));
return Err(WorkerQuitReason::fatal(
StreamableHttpError::TransportChannelClosed,
msg,
));
}
};
let mut uses_modern_http = !is_legacy_startup;
let mut session_id: Option<Arc<str>> = if uses_modern_http {
None
} else if let Some(session_id) = session_id {
Some(session_id.into())
} else {
if !self.config.allow_stateless {
return Err(WorkerQuitReason::fatal(
StreamableHttpError::<C::Error>::MissingSessionIdInResponse,
"process initialize response",
));
}
None
};
let (mut negotiated_version, mut protocol_headers) = if is_legacy_startup {
negotiate_version_headers(&message, config.custom_headers.clone())
} else {
(bootstrap_version, bootstrap_headers)
};
let mut tool_header_cache: HashMap<String, Arc<JsonObject>> = HashMap::new();
let mut session_cleanup_info = session_id.as_ref().map(|sid| SessionCleanupInfo {
client: self.client.clone(),
uri: config.uri.clone(),
session_id: sid.clone(),
auth_header: config.auth_header.clone(),
protocol_headers: protocol_headers.clone(),
});
context.send_to_handler(message).await?;
if is_legacy_startup {
let initialized_notification = context.recv_from_handler().await?;
let initialized_headers = build_request_headers(
&protocol_headers,
&initialized_notification.message,
&tool_header_cache,
&negotiated_version,
);
self.client
.post_message_with_max_sse_event_size(
config.uri.clone(),
initialized_notification.message,
session_id.clone(),
config.auth_header.clone(),
initialized_headers,
config.max_sse_event_size,
)
.await
.map_err(WorkerQuitReason::fatal_context(
"send initialized notification",
))?
.expect_accepted_or_json::<C::Error>()
.map_err(WorkerQuitReason::fatal_context(
"process initialized notification response",
))?;
let _ = initialized_notification.responder.send(Ok(()));
}
#[expect(
clippy::large_enum_variant,
reason = "the event is short-lived and boxing would add allocation in the event loop"
)]
enum Event<C: StreamableHttpClient> {
ClientMessage(WorkerSendRequest<StreamableHttpClientWorker<C>>),
ControlMessage(WorkerSendRequest<StreamableHttpClientWorker<C>>),
StartPost(WorkerSendRequest<StreamableHttpClientWorker<C>>),
PostResult(PostResult<C>),
RecoveryTimeout,
ServerMessage(ServerJsonRpcMessage),
StreamResult {
request_id: Option<RequestId>,
result: Result<(), StreamableHttpError<C::Error>>,
},
}
let mut streams = tokio::task::JoinSet::new();
let mut pending_stream_response_ids = HashSet::new();
let mut request_stream_cancellations =
HashMap::<RequestId, Arc<RequestCancellationRegistration>>::new();
let mut posts = FuturesUnordered::<BoxFuture<'static, PostResult<C>>>::new();
let mut control_posts = FuturesUnordered::<BoxFuture<'static, PostResult<C>>>::new();
let mut session_cancellation = CancellationToken::new();
let mut pending_message: Option<WorkerSendRequest<Self>> = None;
let mut recovery_posts = VecDeque::<WorkerSendRequest<Self>>::new();
let mut recovery_deadline: Option<tokio::time::Instant> = None;
let mut retrying_recovery = false;
let mut barrier_in_flight = false;
let max_concurrent_requests = config.max_concurrent_requests.max(1);
let mut awaiting_fallback_initialized = false;
if let Some(session_id) = &session_id {
Self::spawn_common_stream(
&mut streams,
self.client.clone(),
session_id.clone(),
&config,
protocol_headers.clone(),
sse_worker_tx.clone(),
transport_task_ct.clone(),
);
}
let loop_result: Result<(), WorkerQuitReason<Self::Error>> = 'main_loop: loop {
if retrying_recovery && recovery_posts.is_empty() && posts.is_empty() {
retrying_recovery = false;
}
if !retrying_recovery
&& !recovery_posts.is_empty()
&& posts.is_empty()
&& control_posts.is_empty()
{
session_cancellation.cancel();
recovery_deadline = None;
let recovery = tokio::select! {
_ = transport_task_ct.cancelled() => {
break 'main_loop Err(WorkerQuitReason::Cancelled);
}
result = tokio::time::timeout(
config.session_recovery_timeout,
Self::perform_reinitialization(
self.client.clone(),
saved_init_request.clone().expect("session recovery requires an initialize request"),
config.uri.clone(),
config.auth_header.clone(),
config.custom_headers.clone(),
config.max_sse_event_size,
),
) => result.unwrap_or(Err(StreamableHttpError::SessionRecoveryTimeout)),
};
match recovery {
Ok((new_session_id, new_version, new_headers)) => {
streams.abort_all();
while streams.join_next().await.is_some() {}
request_stream_cancellations.clear();
Self::drain_queued_stream_messages(
&mut sse_worker_rx,
&mut context,
&mut pending_stream_response_ids,
)
.await?;
Self::fail_pending_responses_except_retries(
&mut context,
&mut pending_stream_response_ids,
&recovery_posts,
)
.await?;
session_id = new_session_id;
negotiated_version = new_version;
protocol_headers = new_headers;
session_cleanup_info = session_id.as_ref().map(|sid| SessionCleanupInfo {
client: self.client.clone(),
uri: config.uri.clone(),
session_id: sid.clone(),
auth_header: config.auth_header.clone(),
protocol_headers: protocol_headers.clone(),
});
context.advance_control_generation();
session_cancellation = CancellationToken::new();
if let Some(session_id) = &session_id {
Self::spawn_common_stream(
&mut streams,
self.client.clone(),
session_id.clone(),
&config,
protocol_headers.clone(),
sse_worker_tx.clone(),
transport_task_ct.clone(),
);
}
retrying_recovery = true;
}
Err(error) => {
session_cancellation = CancellationToken::new();
Self::fail_recovery_posts(
&mut recovery_posts,
&mut pending_stream_response_ids,
error,
);
}
}
continue;
}
let has_post_capacity = posts.len() < max_concurrent_requests;
let may_start = (retrying_recovery || recovery_posts.is_empty())
&& !barrier_in_flight
&& has_post_capacity;
let may_receive = may_start && pending_message.is_none() && !retrying_recovery;
let queued = if retrying_recovery {
recovery_posts.front_mut()
} else {
pending_message.as_mut()
};
let has_queued = queued.is_some();
let can_process_queued = queued.as_ref().is_some_and(|request| {
let retry_completed = retrying_recovery
&& Self::client_request_id(&request.message)
.is_some_and(|id| !pending_stream_response_ids.contains(&id));
let ordering_satisfied =
!Self::is_ordering_barrier(&request.message, &negotiated_version)
|| (posts.is_empty() && control_posts.is_empty());
retry_completed || (may_start && ordering_satisfied)
});
let event = tokio::select! {
_ = async {
if can_process_queued {
return;
}
let request = queued.expect("a POST is queued");
let cancellation = request.cancellation_token().unwrap_or_default();
tokio::select! {
_ = request.responder.closed() => {}
_ = cancellation.cancelled() => {}
}
}, if has_queued => {
let request = if retrying_recovery {
recovery_posts.pop_front()
} else {
pending_message.take()
};
Event::StartPost(request.expect("a POST is ready to start"))
}
_ = transport_task_ct.cancelled() => {
tracing::debug!("cancelled");
break 'main_loop Err(WorkerQuitReason::Cancelled);
}
message = context.from_handler_rx.recv(), if may_receive => {
match message {
Some(msg) => Event::ClientMessage(msg),
None => break 'main_loop Err(WorkerQuitReason::HandlerTerminated),
}
},
message = context.control_from_handler_rx.recv(),
if control_posts.is_empty() && !session_cancellation.is_cancelled() => {
match message {
Some(msg) => Event::ControlMessage(msg),
None => break 'main_loop Err(WorkerQuitReason::HandlerTerminated),
}
},
Some(result) = posts.next(), if !posts.is_empty() => {
Event::PostResult(result)
},
Some(result) = control_posts.next(), if !control_posts.is_empty() => {
Event::PostResult(result)
},
_ = async {
if let Some(deadline) = recovery_deadline {
tokio::time::sleep_until(deadline).await;
}
}, if recovery_deadline.is_some() => Event::RecoveryTimeout,
message = sse_worker_rx.recv() => {
let Some(message) = message else {
tracing::trace!("transport dropped, exiting");
break 'main_loop Err(WorkerQuitReason::HandlerTerminated);
};
Event::ServerMessage(message)
},
terminated_stream = streams.join_next(), if !streams.is_empty() => {
match terminated_stream {
Some(Ok((request_id, result))) => {
Event::StreamResult { request_id, result }
}
Some(Err(error)) => Event::StreamResult {
request_id: None,
result: Err(StreamableHttpError::TokioJoinError(error)),
},
None => continue,
}
}
};
match event {
Event::ClientMessage(send_request) => {
pending_message = Some(send_request);
}
Event::ControlMessage(send_request) => {
if send_request.responder.is_closed() {
continue;
}
let cancellation_request_id =
Self::cancellation_request_id(&send_request.message);
let stale = send_request.control_generation() != context.control_generation();
if stale {
let result = match cancellation_request_id {
Some(_) => Ok(()),
None => Err(StreamableHttpError::SessionExpired),
};
let _ = send_request.responder.send(result);
continue;
}
if let Some(request_id) = cancellation_request_id {
drop(request_stream_cancellations.remove(request_id));
pending_stream_response_ids.remove(request_id);
if uses_modern_http {
let _ = send_request.responder.send(Ok(()));
continue;
}
}
let (version, headers) = request_version_headers(
&protocol_headers,
&send_request.message,
&negotiated_version,
&tool_header_cache,
);
control_posts.push(Self::post_request(
self.client.clone(),
&config,
send_request,
PostSession {
id: session_id.clone(),
headers,
version,
cancellation: session_cancellation.clone(),
},
transport_task_ct.clone(),
));
}
Event::RecoveryTimeout => {
recovery_deadline = None;
session_cancellation.cancel();
tracing::warn!("old-session POSTs did not finish before the recovery deadline");
}
Event::StartPost(send_request) => {
let request_id = Self::client_request_id(&send_request.message);
let send_cancelled = send_request.responder.is_closed()
|| send_request
.cancellation_token()
.is_some_and(|token| token.is_cancelled());
let retry_completed = retrying_recovery
&& request_id
.as_ref()
.is_some_and(|id| !pending_stream_response_ids.contains(id));
if send_cancelled || retry_completed {
if retrying_recovery && let Some(id) = &request_id {
pending_stream_response_ids.remove(id);
}
let _ = send_request.responder.send(Ok(()));
continue;
}
let message = &send_request.message;
let is_fallback_initialize = saved_init_request.is_none()
&& matches!(
message,
ClientJsonRpcMessage::Request(request)
if matches!(
&request.request,
ClientRequest::InitializeRequest(_)
)
);
if is_fallback_initialize {
saved_init_request = Some(message.clone());
let WorkerSendRequest {
message, responder, ..
} = send_request;
debug_assert!(
session_id.is_none()
&& session_cleanup_info.is_none()
&& streams.is_empty(),
"discover bootstrap must not create session state"
);
uses_modern_http = false;
let response = self
.client
.post_message_with_max_sse_event_size(
config.uri.clone(),
message,
None,
config.auth_header.clone(),
config.custom_headers.clone(),
config.max_sse_event_size,
)
.await;
let response = match response {
Ok(response) => {
let _ = responder.send(Ok(()));
response
}
Err(error) => {
let _ = responder.send(Err(error));
continue;
}
};
let (initialize_response, new_session_id) = response
.expect_initialized::<C::Error>()
.await
.map_err(WorkerQuitReason::fatal_context(
"process fallback initialize response",
))?;
session_id = new_session_id.map(Arc::from);
if session_id.is_none() && !config.allow_stateless {
return Err(WorkerQuitReason::fatal(
StreamableHttpError::<C::Error>::MissingSessionIdInResponse,
"process fallback initialize response",
));
}
(negotiated_version, protocol_headers) = negotiate_version_headers(
&initialize_response,
config.custom_headers.clone(),
);
session_cleanup_info =
session_id.as_ref().map(|session_id| SessionCleanupInfo {
client: self.client.clone(),
uri: config.uri.clone(),
session_id: session_id.clone(),
auth_header: config.auth_header.clone(),
protocol_headers: protocol_headers.clone(),
});
context.send_to_handler(initialize_response).await?;
awaiting_fallback_initialized = true;
continue;
}
let barrier = Self::is_ordering_barrier(message, &negotiated_version);
debug_assert!(!barrier || (posts.is_empty() && control_posts.is_empty()));
let inline_version = match message {
ClientJsonRpcMessage::Request(request) => {
request.request.get_meta().protocol_version()
}
_ => None,
};
let (request_version, request_headers) = request_version_headers(
&protocol_headers,
message,
&negotiated_version,
&tool_header_cache,
);
if inline_version.is_some() {
negotiated_version = request_version.clone();
if let Ok(value) = HeaderValue::from_str(request_version.as_str()) {
protocol_headers
.insert(HeaderName::from_static("mcp-protocol-version"), value);
}
if let Some(cleanup) = &mut session_cleanup_info {
cleanup.protocol_headers = protocol_headers.clone();
}
}
barrier_in_flight = barrier;
if let Some(request_id) = request_id {
pending_stream_response_ids.insert(request_id);
}
posts.push(Self::post_request(
self.client.clone(),
&config,
send_request,
PostSession {
id: session_id.clone(),
headers: request_headers,
version: request_version,
cancellation: session_cancellation.clone(),
},
transport_task_ct.clone(),
));
}
Event::PostResult(PostResult {
send_request,
response,
version,
}) => {
let is_control = Self::is_control_message(&send_request.message);
if !is_control {
barrier_in_flight = false;
}
let request_id = Self::client_request_id(&send_request.message);
if request_id
.as_ref()
.is_some_and(|id| !pending_stream_response_ids.contains(id))
{
let _ = send_request.responder.send(Ok(()));
continue;
}
let will_retry =
matches!(&response, Some(Err(StreamableHttpError::SessionExpired)))
&& !is_control
&& !retrying_recovery
&& config.reinit_on_expired_session
&& saved_init_request.is_some();
let awaits_stream_response = matches!(
&response,
Some(Ok(StreamableHttpPostResponse::Accepted
| StreamableHttpPostResponse::Sse(..)))
);
if !will_retry
&& !awaits_stream_response
&& let Some(id) = &request_id
{
pending_stream_response_ids.remove(id);
}
let Some(response) = response else {
let _ = send_request.responder.send(Ok(()));
continue;
};
if will_retry {
if recovery_posts.is_empty() {
recovery_deadline =
Some(tokio::time::Instant::now() + config.session_recovery_timeout);
}
recovery_posts.push_back(send_request);
continue;
}
let request_cancellation = send_request.cancellation_registration();
let WorkerSendRequest {
message, responder, ..
} = send_request;
let is_initialized_notification = matches!(
&message,
ClientJsonRpcMessage::Notification(notification)
if matches!(
¬ification.notification,
ClientNotification::InitializedNotification(_)
)
);
let send_result = match response {
Err(e) => Err(e),
Ok(StreamableHttpPostResponse::Accepted) => {
tracing::trace!("client message accepted");
Ok(())
}
Ok(StreamableHttpPostResponse::Json(mut message, ..)) => {
cache_tools_from_response(
&mut tool_header_cache,
&mut message,
&version,
);
context.send_to_handler(message).await?;
Ok(())
}
Ok(StreamableHttpPostResponse::Sse(stream, ..)) => {
let stream_request_id = request_id;
let sse_stream = Self::response_sse_to_jsonrpc(
stream,
session_id.clone(),
self.client.clone(),
config.uri.clone(),
config.auth_header.clone(),
protocol_headers.clone(),
config.max_sse_event_size,
self.config.retry_config.clone(),
);
let request_ct = request_cancellation
.as_ref()
.map(|registration| registration.token())
.unwrap_or_else(|| transport_task_ct.child_token());
let stream_ct = if uses_modern_http {
request_ct.clone()
} else {
request_cancellation
.as_ref()
.map(|registration| registration.lifetime_token())
.unwrap_or_else(|| request_ct.clone())
};
if let (Some(request_id), Some(registration)) =
(stream_request_id.as_ref(), request_cancellation)
{
request_stream_cancellations
.insert(request_id.clone(), registration);
}
let stream_tx = sse_worker_tx.clone();
let origin = match &stream_request_id {
Some(id) => InboundStreamOrigin::OutboundRequest(id.clone()),
None => InboundStreamOrigin::Unassociated,
};
streams.spawn(async move {
let result = Self::run_response_stream(
sse_stream,
stream_tx,
origin,
request_ct,
stream_ct,
uses_modern_http,
)
.await;
(stream_request_id, result)
});
tracing::trace!("got new sse stream");
Ok(())
}
};
if send_result.is_ok()
&& awaiting_fallback_initialized
&& is_initialized_notification
{
if let Some(session_id) = &session_id {
Self::spawn_common_stream(
&mut streams,
self.client.clone(),
session_id.clone(),
&config,
protocol_headers.clone(),
sse_worker_tx.clone(),
transport_task_ct.clone(),
);
}
awaiting_fallback_initialized = false;
}
let _ = responder.send(send_result);
}
Event::ServerMessage(mut json_rpc_message) => {
if let Some(request_id) = Self::clear_stream_response_pending(
&mut pending_stream_response_ids,
&json_rpc_message,
) {
drop(request_stream_cancellations.remove(&request_id));
}
cache_tools_from_response(
&mut tool_header_cache,
&mut json_rpc_message,
&negotiated_version,
);
if let Err(e) = context.send_to_handler(json_rpc_message).await {
break 'main_loop Err(e);
}
}
Event::StreamResult { request_id, result } => {
if let Some(request_id) = request_id {
Self::drain_queued_stream_messages(
&mut sse_worker_rx,
&mut context,
&mut pending_stream_response_ids,
)
.await?;
let cancelled = request_stream_cancellations
.remove(&request_id)
.is_some_and(|registration| registration.token().is_cancelled());
if pending_stream_response_ids.remove(&request_id) && !cancelled {
context
.send_to_handler(ServerJsonRpcMessage::error(
ErrorData::transport_closed(
"streamable HTTP response stream closed before its final response",
),
Some(request_id),
))
.await?;
}
}
if result.is_err() {
tracing::warn!(
"sse client event stream terminated with error: {:?}",
result
);
}
}
}
};
transport_task_ct.cancel();
drop(posts);
drop(control_posts);
drop(pending_message);
drop(recovery_posts);
streams.abort_all();
if let Some(cleanup) = session_cleanup_info {
let cleanup_session_id = cleanup.session_id.clone();
match tokio::time::timeout(
SESSION_CLEANUP_TIMEOUT,
cleanup.client.delete_session(
cleanup.uri,
cleanup.session_id,
cleanup.auth_header,
cleanup.protocol_headers,
),
)
.await
{
Ok(Ok(_)) => {
tracing::info!(
session_id = cleanup_session_id.as_ref(),
"delete session success"
)
}
Ok(Err(StreamableHttpError::ServerDoesNotSupportDeleteSession)) => {
tracing::info!(
session_id = cleanup_session_id.as_ref(),
"server doesn't support delete session"
)
}
Ok(Err(e)) => {
tracing::error!(
session_id = cleanup_session_id.as_ref(),
"fail to delete session: {e}"
);
}
Err(_elapsed) => {
tracing::warn!(
session_id = cleanup_session_id.as_ref(),
"session cleanup timed out after {:?}",
SESSION_CLEANUP_TIMEOUT
);
}
}
}
loop_result
}
}
pub type StreamableHttpClientTransport<C> = WorkerTransport<StreamableHttpClientWorker<C>>;
impl<C: StreamableHttpClient> StreamableHttpClientTransport<C> {
pub fn with_client(client: C, config: StreamableHttpClientTransportConfig) -> Self {
let worker = StreamableHttpClientWorker::new(client, config);
WorkerTransport::spawn(worker)
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct StreamableHttpClientTransportConfig {
pub uri: Arc<str>,
pub retry_config: Arc<dyn SseRetryPolicy>,
pub channel_buffer_capacity: usize,
pub max_concurrent_requests: usize,
pub control_request_timeout: Duration,
pub session_recovery_timeout: Duration,
pub allow_stateless: bool,
pub auth_header: Option<String>,
pub custom_headers: HashMap<HeaderName, HeaderValue>,
pub max_sse_event_size: usize,
pub reinit_on_expired_session: bool,
}
impl StreamableHttpClientTransportConfig {
pub fn with_uri(uri: impl Into<Arc<str>>) -> Self {
Self {
uri: uri.into(),
..Default::default()
}
}
pub fn max_concurrent_requests(mut self, limit: usize) -> Self {
self.max_concurrent_requests = limit.max(1);
self
}
pub fn control_request_timeout(mut self, timeout: Duration) -> Self {
self.control_request_timeout = timeout;
self
}
pub fn session_recovery_timeout(mut self, timeout: Duration) -> Self {
self.session_recovery_timeout = timeout;
self
}
pub fn auth_header<T: Into<String>>(mut self, value: T) -> Self {
self.auth_header = Some(value.into());
self
}
pub fn custom_headers(mut self, custom_headers: HashMap<HeaderName, HeaderValue>) -> Self {
self.custom_headers = custom_headers;
self
}
pub fn max_sse_event_size(mut self, bytes: usize) -> Self {
self.max_sse_event_size = bytes;
self
}
pub fn reinit_on_expired_session(mut self, enable: bool) -> Self {
self.reinit_on_expired_session = enable;
self
}
}
impl Default for StreamableHttpClientTransportConfig {
fn default() -> Self {
Self {
uri: "localhost".into(),
retry_config: Arc::new(ExponentialBackoff::default()),
channel_buffer_capacity: 16,
max_concurrent_requests: 16,
control_request_timeout: Duration::from_secs(5),
session_recovery_timeout: Duration::from_secs(5),
allow_stateless: true,
auth_header: None,
custom_headers: HashMap::new(),
max_sse_event_size: DEFAULT_MAX_SSE_EVENT_SIZE,
reinit_on_expired_session: true,
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use serde_json::json;
use super::*;
use crate::{
model::{
GetExtensions, ListToolsResult, NumberOrString, ServerRequest, ServerResult, Tool,
},
service::InboundStreamOrigin,
};
#[expect(
deprecated,
reason = "Sampling is deprecated by SEP-2577 but remains the canonical restricted request"
)]
fn sampling_request_message(id: i64) -> ServerJsonRpcMessage {
use crate::model::{CreateMessageRequest, CreateMessageRequestParams, SamplingMessage};
ServerJsonRpcMessage::request(
ServerRequest::CreateMessageRequest(CreateMessageRequest::new(
CreateMessageRequestParams::new(vec![SamplingMessage::user_text("hi")], 16),
)),
NumberOrString::Number(id),
)
}
#[tokio::test]
async fn execute_sse_stream_marks_inbound_requests_with_origin() {
for origin in [
InboundStreamOrigin::Unassociated,
InboundStreamOrigin::OutboundRequest(RequestId::Number(3)),
] {
let response = ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(ListToolsResult::default()),
NumberOrString::Number(1),
);
let stream = futures::stream::iter([Ok(sampling_request_message(9)), Ok(response)]);
let (tx, mut rx) = tokio::sync::mpsc::channel(4);
StreamableHttpClientWorker::<StatelessReconnectClient>::execute_sse_stream(
stream,
tx,
origin.clone(),
false,
CancellationToken::new(),
)
.await
.expect("stream completes");
let ServerJsonRpcMessage::Request(request) =
rx.recv().await.expect("request forwarded")
else {
panic!("expected request first");
};
assert_eq!(
request.request.extensions().get::<InboundStreamOrigin>(),
Some(&origin),
"inbound requests must carry their stream origin"
);
assert!(matches!(
rx.recv().await.expect("response forwarded"),
ServerJsonRpcMessage::Response(_)
));
}
}
type ReconnectAttempt = (Option<String>, Option<String>);
#[derive(Clone, Default)]
struct StatelessReconnectClient {
reconnects: Arc<Mutex<Vec<ReconnectAttempt>>>,
}
impl StreamableHttpClient for StatelessReconnectClient {
type Error = std::io::Error;
async fn post_message(
&self,
_uri: Arc<str>,
_message: ClientJsonRpcMessage,
_session_id: Option<Arc<str>>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
Err(StreamableHttpError::UnexpectedServerResponse(
"unexpected POST".into(),
))
}
async fn delete_session(
&self,
_uri: Arc<str>,
_session_id: Arc<str>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<(), StreamableHttpError<Self::Error>> {
Ok(())
}
async fn get_stream(
&self,
_uri: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<BoxedSseStream, StreamableHttpError<Self::Error>> {
self.reconnects
.lock()
.expect("lock reconnects")
.push((session_id.map(|id| id.to_string()), last_event_id));
let response = ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(ListToolsResult::default()),
NumberOrString::Number(1),
);
Ok(futures::stream::once(async move {
Ok(Sse {
event: None,
data: Some(serde_json::to_string(&response).expect("serialize response")),
id: Some("event-1".into()),
retry: None,
})
})
.boxed())
}
}
#[tokio::test]
async fn stateless_response_reconnects_with_last_event_id() {
let initial = futures::stream::iter([Ok(Sse {
event: None,
data: None,
id: Some("event-0".into()),
retry: Some(0),
})])
.boxed();
let client = StatelessReconnectClient::default();
let reconnects = client.reconnects.clone();
let stream =
StreamableHttpClientWorker::<StatelessReconnectClient>::response_sse_to_jsonrpc(
initial,
None,
client,
Arc::from("http://localhost/mcp"),
None,
HashMap::new(),
DEFAULT_MAX_SSE_EVENT_SIZE,
Arc::new(ExponentialBackoff {
max_times: Some(1),
base_duration: Duration::ZERO,
}),
);
let mut stream = std::pin::pin!(stream);
let message = stream.next().await.expect("replayed response").unwrap();
assert!(matches!(message, ServerJsonRpcMessage::Response(_)));
assert_eq!(
reconnects.lock().expect("lock reconnects").as_slice(),
&[(None, Some("event-0".into()))]
);
}
#[derive(Clone, Default)]
struct ResumedRequestClient {
reconnects: Arc<Mutex<Vec<ReconnectAttempt>>>,
}
impl StreamableHttpClient for ResumedRequestClient {
type Error = std::io::Error;
async fn post_message(
&self,
_uri: Arc<str>,
_message: ClientJsonRpcMessage,
_session_id: Option<Arc<str>>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
Err(StreamableHttpError::UnexpectedServerResponse(
"unexpected POST".into(),
))
}
async fn delete_session(
&self,
_uri: Arc<str>,
_session_id: Arc<str>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<(), StreamableHttpError<Self::Error>> {
Ok(())
}
async fn get_stream(
&self,
_uri: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<BoxedSseStream, StreamableHttpError<Self::Error>> {
self.reconnects
.lock()
.expect("lock reconnects")
.push((session_id.map(|id| id.to_string()), last_event_id));
let request = sampling_request_message(9);
let response = ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(ListToolsResult::default()),
NumberOrString::Number(1),
);
Ok(futures::stream::iter([request, response].map(|message| {
Ok(Sse {
event: None,
data: Some(serde_json::to_string(&message).expect("serialize message")),
id: None,
retry: None,
})
}))
.chain(futures::stream::pending())
.boxed())
}
}
#[tokio::test]
async fn resumed_post_stream_requests_keep_outbound_origin() {
let initial = futures::stream::iter([Ok(Sse {
event: None,
data: None,
id: Some("e1".into()),
retry: Some(0),
})])
.boxed();
let client = ResumedRequestClient::default();
let reconnects = client.reconnects.clone();
let sse_stream =
StreamableHttpClientWorker::<ResumedRequestClient>::response_sse_to_jsonrpc(
initial,
None,
client,
Arc::from("http://localhost/mcp"),
None,
HashMap::new(),
DEFAULT_MAX_SSE_EVENT_SIZE,
Arc::new(ExponentialBackoff {
max_times: Some(1),
base_duration: Duration::ZERO,
}),
);
let origin = InboundStreamOrigin::OutboundRequest(RequestId::Number(3));
let (tx, mut rx) = tokio::sync::mpsc::channel(4);
StreamableHttpClientWorker::<ResumedRequestClient>::execute_sse_stream(
sse_stream,
tx,
origin.clone(),
true,
CancellationToken::new(),
)
.await
.expect("stream completes");
assert_eq!(
reconnects.lock().expect("lock reconnects").as_slice(),
&[(None, Some("e1".into()))],
"the request must arrive on the resumed connection"
);
let ServerJsonRpcMessage::Request(request) = rx.recv().await.expect("request forwarded")
else {
panic!("expected request first");
};
assert_eq!(
request.request.extensions().get::<InboundStreamOrigin>(),
Some(&origin),
"origin marker must survive SSE resumption"
);
assert!(matches!(
rx.recv().await.expect("response forwarded"),
ServerJsonRpcMessage::Response(_)
));
}
fn tool(name: &'static str, annotation: serde_json::Value) -> Tool {
let schema = json!({
"type": "object",
"properties": {
"value": annotation,
},
});
Tool::new(
name,
name,
Arc::new(schema.as_object().expect("object schema").clone()),
)
}
#[test]
fn cache_tools_removes_invalid_header_annotations() {
let valid = tool(
"valid",
json!({ "type": "string", "x-mcp-header": "Value" }),
);
let invalid = tool("invalid", json!({ "type": "string", "x-mcp-header": "" }));
let mut message = ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(ListToolsResult::with_all_items(vec![valid, invalid])),
NumberOrString::Number(1),
);
let mut cache = HashMap::new();
cache_tools_from_response(&mut cache, &mut message, &ProtocolVersion::V_2026_07_28);
let ServerJsonRpcMessage::Response(response) = &mut message else {
panic!("expected tools/list response");
};
let ServerResult::ListToolsResult(result) = &mut response.result else {
panic!("expected tools/list result");
};
assert_eq!(
(
result
.tools
.iter()
.map(|tool| tool.name.as_ref())
.collect::<Vec<_>>(),
cache.keys().map(String::as_str).collect::<Vec<_>>(),
),
(vec!["valid"], vec!["valid"])
);
}
#[test]
fn cache_tools_preserves_pre_standard_header_results() {
let invalid = tool("legacy", json!({ "type": "string", "x-mcp-header": "" }));
let mut message = ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(ListToolsResult::with_all_items(vec![invalid])),
NumberOrString::Number(1),
);
let mut cache = HashMap::new();
cache_tools_from_response(&mut cache, &mut message, &ProtocolVersion::V_2025_11_25);
let ServerJsonRpcMessage::Response(response) = message else {
panic!("expected tools/list response");
};
let ServerResult::ListToolsResult(result) = response.result else {
panic!("expected tools/list result");
};
assert_eq!(
result
.tools
.iter()
.map(|tool| tool.name.as_ref())
.collect::<Vec<_>>(),
vec!["legacy"]
);
}
#[cfg(feature = "transport-streamable-http-client-reqwest")]
#[test]
fn clear_stream_response_pending_accepts_stringified_numeric_id() {
let mut pending = HashSet::from([NumberOrString::Number(1)]);
let response = ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(ListToolsResult::default()),
NumberOrString::String("1".into()),
);
let matched_id =
StreamableHttpClientWorker::<reqwest::Client>::clear_stream_response_pending(
&mut pending,
&response,
);
assert_eq!(matched_id, Some(NumberOrString::Number(1)));
assert!(pending.is_empty());
}
#[cfg(feature = "transport-streamable-http-client-reqwest")]
#[test]
fn clear_stream_response_pending_prefers_exact_string_id() {
let string_id = NumberOrString::String("1".into());
let mut pending = HashSet::from([NumberOrString::Number(1), string_id.clone()]);
let response = ServerJsonRpcMessage::response(
ServerResult::ListToolsResult(ListToolsResult::default()),
string_id.clone(),
);
let matched_id =
StreamableHttpClientWorker::<reqwest::Client>::clear_stream_response_pending(
&mut pending,
&response,
);
assert_eq!(matched_id, Some(string_id));
assert_eq!(pending, HashSet::from([NumberOrString::Number(1)]));
}
}