#![allow(rustdoc::private_intra_doc_links)]
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::{SharedSender, 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, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
#[cfg(not(target_arch = "wasm32"))]
use tokio::sync::{mpsc, watch};
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;
const CLIENT_RECEIVE_QUEUE_CAPACITY: usize = 64;
const INITIAL_SSE_RECONNECT_DELAY: Duration = Duration::from_secs(1);
const MAX_SSE_RECONNECT_DELAY: Duration = Duration::from_secs(30);
const MIN_SSE_RECONNECT_DELAY: Duration = Duration::from_millis(500);
const RECONNECT_BUDGET_RESET_UPTIME: Duration = Duration::from_secs(30);
fn budget_reset_earned(uptime: Duration) -> bool {
uptime >= RECONNECT_BUDGET_RESET_UPTIME
}
const SSE_RECONNECT_GROWTH: f64 = 1.5;
const MAX_SSE_RECONNECT_ATTEMPTS: u32 = 2;
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(Debug, Clone, Copy, PartialEq, Eq)]
enum OutboundFrame {
InitializedNotification,
Other,
}
impl OutboundFrame {
fn of(message: &TransportMessage) -> Self {
match message {
TransportMessage::Notification(crate::types::Notification::Client(
crate::types::ClientNotification::Initialized,
)) => Self::InitializedNotification,
_ => Self::Other,
}
}
}
struct RequestParts<'a> {
config: &'a Arc<RwLock<StreamableHttpTransportConfig>>,
protocol_version: &'a Arc<RwLock<Option<String>>>,
v2_mode: &'a Arc<AtomicBool>,
cold_vend_gate: &'a Arc<ColdVendGate>,
}
impl RequestParts<'_> {
fn is_v2(&self) -> bool {
self.v2_mode.load(Ordering::Relaxed)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum TerminalKind {
InvalidMessage,
Request,
Closed,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum StreamKind {
Session,
PostResponse,
Transport,
}
impl StreamKind {
fn describe(self) -> &'static str {
match self {
Self::Session => "the GET session stream",
Self::PostResponse => "this call's own POST response stream",
Self::Transport => "this transport",
}
}
}
#[derive(Clone, Debug)]
struct TerminalReason {
kind: TerminalKind,
message: String,
stream: StreamKind,
}
impl TerminalReason {
fn to_error(&self) -> Error {
let message = format!("{} ended: {}", self.stream.describe(), self.message);
Error::Transport(match self.kind {
TerminalKind::InvalidMessage => TransportError::InvalidMessage(message),
TerminalKind::Request => TransportError::Request(message),
TerminalKind::Closed => TransportError::ConnectionClosed,
})
}
}
fn terminal_reason_of(error: &Error, stream: StreamKind) -> TerminalReason {
match error {
Error::Transport(TransportError::InvalidMessage(message)) => TerminalReason {
kind: TerminalKind::InvalidMessage,
message: message.clone(),
stream,
},
Error::Transport(TransportError::Request(message)) => TerminalReason {
kind: TerminalKind::Request,
message: message.clone(),
stream,
},
other => TerminalReason {
kind: TerminalKind::Request,
message: other.to_string(),
stream,
},
}
}
#[derive(Clone)]
struct ReaderDelivery {
sender: mpsc::Sender<Result<TransportMessage>>,
stream: StreamKind,
terminal: Arc<RwLock<Option<TerminalReason>>>,
terminal_signal: Arc<watch::Sender<u64>>,
shutdown: Arc<watch::Sender<bool>>,
}
impl ReaderDelivery {
fn is_closed(&self) -> bool {
self.sender.is_closed()
}
}
fn latch_terminal_reason(delivery: &ReaderDelivery, error: &Error) {
latch_reason(
&delivery.terminal,
&delivery.terminal_signal,
terminal_reason_of(error, delivery.stream),
);
}
fn latch_reason(
terminal: &Arc<RwLock<Option<TerminalReason>>>,
signal: &watch::Sender<u64>,
reason: TerminalReason,
) {
{
let mut slot = terminal.write();
if slot.is_some() {
return;
}
*slot = Some(reason);
}
signal.send_modify(|generation| *generation += 1);
}
struct PostReaderGuard {
counter: Arc<AtomicUsize>,
signal: Arc<watch::Sender<u64>>,
}
impl PostReaderGuard {
fn acquire(counter: &Arc<AtomicUsize>, signal: &Arc<watch::Sender<u64>>) -> Self {
counter.fetch_add(1, Ordering::SeqCst);
Self {
counter: Arc::clone(counter),
signal: Arc::clone(signal),
}
}
}
impl Drop for PostReaderGuard {
fn drop(&mut self) {
if self.counter.fetch_sub(1, Ordering::SeqCst) == 1 {
self.signal.send_modify(|generation| *generation += 1);
}
}
}
fn drain_or_latch(
receiver: &mut mpsc::Receiver<Result<TransportMessage>>,
overflow: &RwLock<std::collections::VecDeque<Result<TransportMessage>>>,
terminal: &Arc<RwLock<Option<TerminalReason>>>,
open_post_readers: &AtomicUsize,
) -> Option<Result<TransportMessage>> {
match receiver.try_recv() {
Ok(queued) => return Some(queued),
Err(mpsc::error::TryRecvError::Disconnected) => {
return Some(
overflow
.write()
.pop_front()
.unwrap_or(Err(Error::Transport(TransportError::ConnectionClosed))),
);
},
Err(mpsc::error::TryRecvError::Empty) => {},
}
let overflowed = overflow.write().pop_front();
if let Some(overflowed) = overflowed {
return Some(overflowed);
}
if open_post_readers.load(Ordering::SeqCst) > 0 {
return None;
}
let latched = terminal.read().as_ref().map(TerminalReason::to_error);
latched.map(Err)
}
#[derive(Debug)]
struct ColdVendGate {
primed: AtomicBool,
lock: tokio::sync::Mutex<()>,
}
impl ColdVendGate {
fn new() -> Self {
Self {
primed: AtomicBool::new(false),
lock: tokio::sync::Mutex::new(()),
}
}
async fn vend(&self, provider: &Arc<dyn AuthProvider>) -> Result<String> {
if self.primed.load(Ordering::Relaxed) {
return provider.get_access_token().await;
}
let _cold = self.lock.lock().await;
let token = provider.get_access_token().await?;
self.primed.store(true, Ordering::Relaxed);
Ok(token)
}
fn mark_cold(&self) {
self.primed.store(false, Ordering::Relaxed);
}
}
#[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::Receiver<Result<TransportMessage>>>>,
sender: mpsc::Sender<Result<TransportMessage>>,
caller_overflow: Arc<RwLock<std::collections::VecDeque<Result<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,
terminal: Arc<RwLock<Option<TerminalReason>>>,
terminal_signal: Arc<watch::Sender<u64>>,
open_post_readers: Arc<AtomicUsize>,
shutdown: Arc<watch::Sender<bool>>,
refresh_lock: Arc<tokio::sync::Mutex<()>>,
cold_vend_gate: Arc<ColdVendGate>,
token_generation: Arc<AtomicU64>,
restart_lock: Arc<tokio::sync::Mutex<()>>,
}
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::channel(CLIENT_RECEIVE_QUEUE_CAPACITY);
let (terminal_signal, _) = watch::channel(0u64);
let (shutdown, _) = watch::channel(false);
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,
terminal: Arc::new(RwLock::new(None)),
terminal_signal: Arc::new(terminal_signal),
open_post_readers: Arc::new(AtomicUsize::new(0)),
shutdown: Arc::new(shutdown),
caller_overflow: Arc::new(RwLock::new(std::collections::VecDeque::new())),
refresh_lock: Arc::new(tokio::sync::Mutex::new(())),
cold_vend_gate: Arc::new(ColdVendGate::new()),
token_generation: Arc::new(AtomicU64::new(0)),
restart_lock: Arc::new(tokio::sync::Mutex::new(())),
}
}
fn reader_delivery(&self, stream: StreamKind) -> ReaderDelivery {
ReaderDelivery {
sender: self.sender.clone(),
stream,
terminal: Arc::clone(&self.terminal),
terminal_signal: Arc::clone(&self.terminal_signal),
shutdown: Arc::clone(&self.shutdown),
}
}
#[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 queue_from_caller(&self, message: TransportMessage) -> Result<()> {
let full = match self.sender.try_send(Ok(message)) {
Ok(()) => return Ok(()),
Err(mpsc::error::TrySendError::Full(message)) => message,
Err(mpsc::error::TrySendError::Closed(_)) => {
return Err(Error::Transport(TransportError::Send(
"the client receive queue is closed".to_string(),
)))
},
};
self.caller_overflow.write().push_back(full);
self.terminal_signal
.send_modify(|generation| *generation += 1);
Ok(())
}
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(parts: &RequestParts<'_>, headers: &hyper::HeaderMap) {
if parts.is_v2() {
return;
}
if let Some(value) = headers.get(MCP_SESSION_ID) {
if let Ok(text) = value.to_str() {
if parts.config.read().session_id.as_deref() == Some(text) {
return;
}
parts.config.write().session_id = Some(text.to_string());
}
}
}
#[cfg(not(feature = "v1-compat"))]
const fn capture_session_header(_parts: &RequestParts<'_>, _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, cursor: Option<String>) -> Result<()> {
let _restart = self.restart_lock.lock().await;
let handle = self.abort_handle.write().take();
if let Some(handle) = handle {
handle.abort();
}
let context = self.sse_reader_context();
let Some(body) = context.open_sse_once(cursor).await? else {
return Ok(());
};
{
let mut slot = self.terminal.write();
if slot
.as_ref()
.is_some_and(|reason| reason.stream == StreamKind::Session)
{
*slot = None;
}
}
let handle = tokio::spawn(async move { context.run_session_stream(body).await });
*self.abort_handle.write() = Some(handle);
Ok(())
}
async fn build_sse_get_request(
parts: &RequestParts<'_>,
#[cfg(feature = "v1-compat")] resumption_token: Option<String>,
#[cfg(not(feature = "v1-compat"))] _ignored_cursor: Option<String>,
) -> Result<Request<Full<Bytes>>> {
let url = parts.config.read().url.clone();
let mut request = Self::build_request_from_parts(
parts,
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())?;
Ok(request)
}
fn sse_reader_context(&self) -> SseReaderContext {
SseReaderContext {
client: self.client.clone(),
config: Arc::clone(&self.config),
protocol_version: Arc::clone(&self.protocol_version),
v2_mode: Arc::clone(&self.v2_mode),
cold_vend_gate: Arc::clone(&self.cold_vend_gate),
delivery: self.reader_delivery(StreamKind::Session),
last_event_id: Arc::clone(&self.last_event_id),
on_resumption: self.resumption_callback(),
max_collected_body_bytes: self.max_collected_body_bytes,
}
}
fn spawn_sse_reader(&self, body: hyper::body::Incoming) -> tokio::task::JoinHandle<()> {
let delivery = self.reader_delivery(StreamKind::PostResponse);
let on_resumption = self.resumption_callback();
let last_event_id = Arc::clone(&self.last_event_id);
let max_buffer_size = self.max_collected_body_bytes;
let in_flight = PostReaderGuard::acquire(&self.open_post_readers, &self.terminal_signal);
tokio::spawn(async move {
let _in_flight = in_flight;
let mut cursor: Option<String> = None;
let end = read_sse_body(
&delivery,
&last_event_id,
on_resumption.as_ref(),
body,
max_buffer_size,
&mut cursor,
)
.await;
if let SseBodyEnd::Dropped {
cause: Some(error), ..
} = end
{
latch_terminal_reason(&delivery, &error);
}
})
}
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
}
fn request_parts(&self) -> RequestParts<'_> {
RequestParts {
config: &self.config,
protocol_version: &self.protocol_version,
v2_mode: &self.v2_mode,
cold_vend_gate: &self.cold_vend_gate,
}
}
async fn build_request_with_middleware(
&self,
method: Method,
url: &str,
body: Vec<u8>,
) -> Result<Request<Full<Bytes>>> {
Self::build_request_from_parts(&self.request_parts(), method, url, body).await
}
async fn build_request_from_parts(
parts: &RequestParts<'_>,
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 = parts.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 = parts.cold_vend_gate.vend(&auth_provider).await?;
request_builder = request_builder.header("Authorization", format!("Bearer {}", token));
true
} else {
false
};
let is_v2 = parts.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) = parts.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 crate::shared::wire_trace::enabled() {
let rendered = crate::shared::wire_trace::render_headers(
headers
.iter()
.filter_map(|(k, v)| v.to_str().ok().map(|v| (k.as_str(), v))),
);
tracing::debug!(
target: crate::shared::wire_trace::WIRE_TARGET,
direction = "request",
%url,
method = %method,
headers = %rendered,
body = %crate::shared::wire_trace::render_body(&body),
"outgoing MCP request"
);
}
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::process_headers_from(&self.request_parts(), response.headers());
}
fn process_headers_from(parts: &RequestParts<'_>, headers: &hyper::HeaderMap) {
Self::capture_session_header(parts, headers);
if let Some(protocol_version) = headers.get(MCP_PROTOCOL_VERSION) {
if let Ok(protocol_version_str) = protocol_version.to_str() {
*parts.protocol_version.write() = Some(protocol_version_str.to_string());
}
}
}
pub async fn send_with_options(
&mut self,
message: TransportMessage,
options: SendOptions,
) -> Result<()> {
self.send_with_options_shared(message, options).await
}
async fn send_with_options_shared(
&self,
message: TransportMessage,
options: SendOptions,
) -> Result<()> {
if let Some(token) = options.resumption_cursor() {
self.start_sse(Some(token)).await?;
return Ok(());
}
let outbound = OutboundFrame::of(&message);
let body_bytes = crate::shared::StdioTransport::serialize_message(&message)?;
self.post_body(body_bytes, outbound).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?;
let presented_generation = self.token_generation.load(Ordering::SeqCst);
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 {
if crate::shared::wire_trace::enabled() {
let rendered = crate::shared::wire_trace::render_headers(
response
.headers()
.iter()
.filter_map(|(k, v)| v.to_str().ok().map(|v| (k.as_str(), v))),
);
tracing::debug!(
target: crate::shared::wire_trace::WIRE_TARGET,
direction = "response",
status = response.status().as_u16(),
headers = %rendered,
"incoming MCP response"
);
}
return Ok(response);
}
let auth_provider = self.config.read().auth_provider.clone();
let Some(provider) = auth_provider else {
return Ok(response);
};
let retry_request = {
let _refresh = self.refresh_lock.lock().await;
if self.token_generation.load(Ordering::SeqCst) == presented_generation {
provider.on_unauthorized().await?;
self.token_generation.fetch_add(1, Ordering::SeqCst);
self.cold_vend_gate.mark_cold();
}
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)?;
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>, outbound: OutboundFrame) -> Result<()> {
let response = self.post_once(body_bytes).await?;
self.process_response_headers(&response);
if response.status() == StatusCode::ACCEPTED {
if outbound == OutboundFrame::InitializedNotification {
let _ = self.start_sse(None).await;
}
return Ok(());
}
if !response.status().is_success() {
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.queue_from_caller(message)?;
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());
if content_type.contains(TEXT_EVENT_STREAM) {
drop(self.spawn_sse_reader(response.into_body()));
return 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.queue_from_caller(msg)?;
}
} else {
let msg_parsed = crate::shared::StdioTransport::parse_message(&modified_body)?;
self.queue_from_caller(msg_parsed)?;
}
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.queue_from_caller(msg)?;
}
} else {
let msg_parsed = crate::shared::StdioTransport::parse_message(&modified_body)?;
self.queue_from_caller(msg_parsed)?;
}
} 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 SharedSender for StreamableHttpTransport {
async fn send_shared(&self, message: TransportMessage) -> Result<()> {
self.send_with_options_shared(message, SendOptions::default())
.await
}
async fn send_raw_shared(&self, body: Vec<u8>) -> Result<()> {
self.post_body(body, OutboundFrame::Other).await
}
}
#[async_trait]
impl Transport for StreamableHttpTransport {
async fn send(&mut self, message: TransportMessage) -> Result<()> {
self.send_with_options(message, SendOptions::default())
.await
}
fn shared_sender(&self) -> Option<Arc<dyn SharedSender>> {
Some(Arc::new(self.clone()))
}
async fn receive(&mut self) -> Result<TransportMessage> {
let mut signal = self.terminal_signal.subscribe();
let mut receiver = self.receiver.lock().await;
loop {
if let Some(outcome) = drain_or_latch(
&mut receiver,
&self.caller_overflow,
&self.terminal,
&self.open_post_readers,
) {
return outcome;
}
tokio::select! {
biased;
queued = receiver.recv() => {
return queued
.ok_or_else(|| Error::Transport(TransportError::ConnectionClosed))?;
},
signalled = signal.changed() => {
if signalled.is_err() {
return receiver
.recv()
.await
.ok_or_else(|| Error::Transport(TransportError::ConnectionClosed))?;
}
},
}
}
}
async fn close(&mut self) -> Result<()> {
self.shutdown.send_replace(true);
let handle = self.abort_handle.write().take();
if let Some(handle) = handle {
handle.abort();
}
latch_reason(
&self.terminal,
&self.terminal_signal,
TerminalReason {
kind: TerminalKind::Closed,
message: "closed by the application".to_string(),
stream: StreamKind::Transport,
},
);
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, OutboundFrame::Other).await
}
}
const MAX_ECHOED_SSE_FRAME: usize = 200;
struct SseReadState {
body: hyper::body::Incoming,
parser: SseParser,
bytes: Vec<u8>,
pending: std::collections::VecDeque<crate::shared::sse_parser::SseEvent>,
done: bool,
}
impl SseReadState {
fn new(body: hyper::body::Incoming, max_buffer_size: usize) -> Self {
Self {
body,
parser: SseParser::with_max_buffer_size(max_buffer_size),
bytes: Vec::new(),
pending: std::collections::VecDeque::new(),
done: false,
}
}
}
enum SseFrameStop {
Dropped(Error),
Corrupt(Error),
Shutdown,
}
enum SseBodyEnd {
Dropped {
cause: Option<Error>,
retry: Option<Duration>,
},
Ended,
}
async fn read_sse_body(
delivery: &ReaderDelivery,
last_event_id: &Arc<RwLock<Option<String>>>,
on_resumption: Option<&Arc<dyn Fn(String) + Send + Sync>>,
body: hyper::body::Incoming,
max_buffer_size: usize,
cursor: &mut Option<String>,
) -> SseBodyEnd {
let mut state = SseReadState::new(body, max_buffer_size);
let mut progress = SseProgress::default();
let mut shutdown = delivery.shutdown.subscribe();
if *shutdown.borrow() {
return SseBodyEnd::Ended;
}
loop {
if !drain_pending_events(
delivery,
last_event_id,
on_resumption,
&mut state,
&mut progress,
cursor,
)
.await
{
return SseBodyEnd::Ended;
}
if state.done {
return progress.dropped(None);
}
if let Some(stop) = read_next_sse_frame(&mut state, delivery, &mut shutdown).await {
return end_of_frame_stop(stop, delivery, &progress);
}
}
}
fn end_of_frame_stop(
stop: SseFrameStop,
delivery: &ReaderDelivery,
progress: &SseProgress,
) -> SseBodyEnd {
match stop {
SseFrameStop::Dropped(error) => progress.dropped(Some(error)),
SseFrameStop::Shutdown => SseBodyEnd::Ended,
SseFrameStop::Corrupt(error) => {
latch_terminal_reason(delivery, &error);
SseBodyEnd::Ended
},
}
}
#[derive(Default)]
struct SseProgress {
retry: Option<Duration>,
}
impl SseProgress {
fn dropped(&self, cause: Option<Error>) -> SseBodyEnd {
SseBodyEnd::Dropped {
cause,
retry: self.retry,
}
}
}
async fn drain_pending_events(
delivery: &ReaderDelivery,
last_event_id: &Arc<RwLock<Option<String>>>,
on_resumption: Option<&Arc<dyn Fn(String) + Send + Sync>>,
state: &mut SseReadState,
progress: &mut SseProgress,
cursor: &mut Option<String>,
) -> bool {
while let Some(event) = state.pending.pop_front() {
if let Some(millis) = event.retry {
progress.retry = Some(Duration::from_millis(millis));
}
if !deliver_sse_event(delivery, last_event_id, on_resumption, event, cursor).await {
return false;
}
}
true
}
#[cfg(feature = "v1-compat")]
fn reconnect_cursor(cursor: Option<&str>) -> Option<String> {
cursor.map(ToString::to_string)
}
#[cfg(not(feature = "v1-compat"))]
const fn reconnect_cursor(_ignored_cursor: Option<&str>) -> Option<String> {
None
}
async fn read_next_sse_frame(
state: &mut SseReadState,
delivery: &ReaderDelivery,
shutdown: &mut watch::Receiver<bool>,
) -> Option<SseFrameStop> {
let polled = tokio::select! {
biased;
frame = state.body.frame() => Some(frame),
() = delivery.sender.closed() => None,
_ = shutdown.changed() => None,
};
let Some(frame) = polled else {
state.done = true;
return Some(SseFrameStop::Shutdown);
};
match frame {
None => {
state.done = true;
None
},
Some(Err(e)) => {
state.done = true;
Some(SseFrameStop::Dropped(Error::Transport(
TransportError::Request(e.to_string()),
)))
},
Some(Ok(frame)) => {
if let Some(chunk) = frame.data_ref() {
state.bytes.extend_from_slice(chunk);
let text = crate::shared::sse_parser::take_utf8_prefix(&mut state.bytes);
state
.pending
.extend(drain_sse_events(&mut state.parser, &text));
if let Some(error) = sse_stream_overflow(&state.parser) {
state.done = true;
return Some(SseFrameStop::Corrupt(error));
}
}
None
},
}
}
fn drain_sse_events(
parser: &mut SseParser,
chunk: &str,
) -> Vec<crate::shared::sse_parser::SseEvent> {
parser
.feed(chunk)
.into_iter()
.filter(|event| event.event.as_deref().is_none_or(|name| name == "message"))
.collect()
}
fn sse_stream_overflow(parser: &SseParser) -> Option<Error> {
if !parser.overflowed() {
return None;
}
Some(Error::Transport(TransportError::InvalidMessage(format!(
"a session-stream chunk pushed the buffered stream state past the {}-byte parser bound; \
the buffered bytes were discarded and the stream was ended",
parser.max_buffer_size()
))))
}
fn unparseable_sse_frame(cause: &Error, data: &str) -> Error {
Error::Transport(TransportError::InvalidMessage(format!(
"a session-stream frame did not parse as a JSON-RPC message ({cause}); the stream was \
ended. Frame: {}",
truncate_sse_frame(data)
)))
}
fn truncate_sse_frame(text: &str) -> String {
let mut boundary = None;
for (index, (offset, _)) in text.char_indices().enumerate() {
if index == MAX_ECHOED_SSE_FRAME {
boundary = Some(offset);
break;
}
}
let Some(boundary) = boundary else {
return text.to_string();
};
let mut out = String::with_capacity(boundary + '…'.len_utf8());
out.push_str(&text[..boundary]);
out.push('…');
out
}
async fn deliver_sse_event(
delivery: &ReaderDelivery,
last_event_id: &Arc<RwLock<Option<String>>>,
on_resumption: Option<&Arc<dyn Fn(String) + Send + Sync>>,
event: crate::shared::sse_parser::SseEvent,
cursor: &mut Option<String>,
) -> bool {
if let Some(id) = &event.id {
*last_event_id.write() = Some(id.clone());
*cursor = Some(id.clone());
if let Some(callback) = on_resumption {
callback(id.clone());
}
}
match crate::shared::StdioTransport::parse_message(event.data.as_bytes()) {
Ok(message) => delivery.sender.send(Ok(message)).await.is_ok(),
Err(cause) => {
latch_terminal_reason(delivery, &unparseable_sse_frame(&cause, &event.data));
false
},
}
}
#[cfg(any(feature = "fuzzing", test))]
#[doc(hidden)]
#[must_use]
#[allow(clippy::type_complexity)]
pub fn decode_sse_chunks_for_fuzz(
chunks: &[&[u8]],
max_buffer_size: usize,
) -> (
Vec<std::result::Result<TransportMessage, String>>,
Vec<bool>,
Vec<usize>,
Vec<usize>,
) {
let mut parser = SseParser::with_max_buffer_size(max_buffer_size);
let mut bytes: Vec<u8> = Vec::new();
let mut outcomes = Vec::new();
let mut overflowed = Vec::with_capacity(chunks.len());
let mut peak_buffered_bytes = Vec::with_capacity(chunks.len());
let mut undecoded_tail_bytes = Vec::with_capacity(chunks.len());
for chunk in chunks {
bytes.extend_from_slice(chunk);
let text = crate::shared::sse_parser::take_utf8_prefix(&mut bytes);
outcomes.extend(
drain_sse_events(&mut parser, &text)
.into_iter()
.map(|event| {
crate::shared::StdioTransport::parse_message(event.data.as_bytes())
.map_err(|error| error.to_string())
}),
);
overflowed.push(sse_stream_overflow(&parser).is_some());
peak_buffered_bytes.push(parser.buffered_bytes());
undecoded_tail_bytes.push(bytes.len());
}
(
outcomes,
overflowed,
peak_buffered_bytes,
undecoded_tail_bytes,
)
}
struct SseReaderContext {
client: Client<
hyper_rustls::HttpsConnector<hyper_util::client::legacy::connect::HttpConnector>,
Full<Bytes>,
>,
config: Arc<RwLock<StreamableHttpTransportConfig>>,
protocol_version: Arc<RwLock<Option<String>>>,
v2_mode: Arc<AtomicBool>,
cold_vend_gate: Arc<ColdVendGate>,
delivery: ReaderDelivery,
last_event_id: Arc<RwLock<Option<String>>>,
on_resumption: Option<Arc<dyn Fn(String) + Send + Sync>>,
max_collected_body_bytes: usize,
}
impl SseReaderContext {
fn request_parts(&self) -> RequestParts<'_> {
RequestParts {
config: &self.config,
protocol_version: &self.protocol_version,
v2_mode: &self.v2_mode,
cold_vend_gate: &self.cold_vend_gate,
}
}
async fn open_sse_once(&self, cursor: Option<String>) -> Result<Option<hyper::body::Incoming>> {
let request =
StreamableHttpTransport::build_sse_get_request(&self.request_parts(), cursor).await?;
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(None);
}
if !response.status().is_success() {
return Err(Error::Transport(TransportError::Request(format!(
"SSE request failed with status: {}",
response.status()
))));
}
StreamableHttpTransport::process_headers_from(&self.request_parts(), response.headers());
Ok(Some(response.into_body()))
}
async fn run_session_stream(self, mut body: hyper::body::Incoming) {
let mut attempt: u32 = 0;
let mut cursor: Option<String> = None;
loop {
let opened_at = std::time::Instant::now();
let end = read_sse_body(
&self.delivery,
&self.last_event_id,
self.on_resumption.as_ref(),
body,
self.max_collected_body_bytes,
&mut cursor,
)
.await;
let SseBodyEnd::Dropped { cause, retry } = end else {
return;
};
if budget_reset_earned(opened_at.elapsed()) {
attempt = 0;
}
if attempt >= MAX_SSE_RECONNECT_ATTEMPTS {
latch_terminal_reason(
&self.delivery,
&reconnect_budget_exhausted(attempt, cause.as_ref()),
);
return;
}
if self.delivery.is_closed() {
return;
}
tokio::time::sleep(next_reconnect_delay(attempt, retry)).await;
if self.delivery.is_closed() {
return;
}
attempt += 1;
match self
.open_sse_once(reconnect_cursor(cursor.as_deref()))
.await
{
Ok(Some(reopened)) => body = reopened,
Ok(None) => {
latch_terminal_reason(&self.delivery, &reconnect_stream_gone(attempt));
return;
},
Err(error) => {
latch_terminal_reason(&self.delivery, &reconnect_open_failed(attempt, &error));
return;
},
}
}
}
}
fn next_reconnect_delay(attempt: u32, server_retry: Option<Duration>) -> Duration {
if let Some(retry) = server_retry {
return retry
.max(MIN_SSE_RECONNECT_DELAY)
.min(MAX_SSE_RECONNECT_DELAY);
}
let exponent = i32::try_from(attempt).unwrap_or(i32::MAX);
let seconds = INITIAL_SSE_RECONNECT_DELAY.as_secs_f64() * SSE_RECONNECT_GROWTH.powi(exponent);
Duration::try_from_secs_f64(seconds)
.unwrap_or(MAX_SSE_RECONNECT_DELAY)
.min(MAX_SSE_RECONNECT_DELAY)
}
fn reconnect_budget_exhausted(attempts: u32, cause: Option<&Error>) -> Error {
let because = cause.map_or_else(
|| "the peer ended the body".to_string(),
|error| format!("the last attempt ended with: {error}"),
);
Error::Transport(TransportError::Request(format!(
"the session stream was dropped and its {MAX_SSE_RECONNECT_ATTEMPTS}-attempt reconnect \
budget (MAX_SSE_RECONNECT_ATTEMPTS) is exhausted after {attempts} attempt(s); {because}"
)))
}
fn reconnect_stream_gone(attempts: u32) -> Error {
Error::Transport(TransportError::Request(format!(
"the session stream was dropped and reconnect attempt {attempts} was answered 405 Method \
Not Allowed: the server no longer offers a GET session stream, so there is nothing left \
to resume"
)))
}
fn reconnect_open_failed(attempts: u32, cause: &Error) -> Error {
Error::Transport(TransportError::Request(format!(
"the session stream was dropped and reconnect attempt {attempts} could not re-open it \
({cause}); the stream was ended. The client deliberately does NOT re-handshake: a silent \
re-`initialize` would mint a new session and orphan every in-flight correlation"
)))
}
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 json_body_of(len: usize) -> String {
let empty = r#"{"jsonrpc":"2.0","id":42,"result":{"tools":[],"pad":""}}"#;
let padding = len
.checked_sub(empty.len())
.expect("requested length must fit one frame");
let body = format!(
r#"{{"jsonrpc":"2.0","id":42,"result":{{"tools":[],"pad":"{}"}}}}"#,
"p".repeat(padding)
);
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_over_the_parser_bound_ends_the_stream_with_a_named_error() {
let mut server = MockServer::new_async().await;
let body = format!("data: {}", "p".repeat(CAP * 2));
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("send returns as soon as the reader task is spawned");
let error = tokio::time::timeout(QUIET_WINDOW, transport.receive())
.await
.expect("the terminal error must be dispatched")
.expect_err("an over-bound chunk must END the stream, not be silently dropped");
assert_over_cap_refusal(&error, CAP);
}
#[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_over_the_parser_bound_ends_the_stream_with_a_named_error() {
let mut server = MockServer::new_async().await;
let body = format!("data: {}", "p".repeat(CAP * 2));
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("start_sse returns as soon as its reader task is spawned");
let error = tokio::time::timeout(QUIET_WINDOW, transport.receive())
.await
.expect("the terminal error must be dispatched")
.expect_err("an over-bound chunk must END the stream, not be silently dropped");
assert_over_cap_refusal(&error, CAP);
}
#[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", APPLICATION_JSON)
.with_body(json_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}"
);
}
}
mod reconnect_delay_bounds {
use super::*;
const STEP: Duration = Duration::from_millis(1);
#[test]
fn a_peer_asking_for_zero_still_waits_the_floor() {
assert_eq!(
next_reconnect_delay(0, Some(Duration::ZERO)),
MIN_SSE_RECONNECT_DELAY,
"`retry: 0` is the CR-01 input: honoured verbatim it turns the reconnect loop \
into a request flood that also re-mints an access token per iteration"
);
}
#[test]
fn a_peer_value_below_the_floor_is_raised_to_it() {
assert_eq!(
next_reconnect_delay(0, Some(MIN_SSE_RECONNECT_DELAY.saturating_sub(STEP))),
MIN_SSE_RECONNECT_DELAY,
"one millisecond under the floor is still under the floor"
);
}
#[test]
fn a_peer_value_exactly_at_the_floor_is_left_alone() {
assert_eq!(
next_reconnect_delay(0, Some(MIN_SSE_RECONNECT_DELAY)),
MIN_SSE_RECONNECT_DELAY,
"the bound is inclusive at its lower end"
);
}
#[test]
fn a_peer_value_just_above_the_floor_is_honoured_verbatim() {
assert_eq!(
next_reconnect_delay(0, Some(MIN_SSE_RECONNECT_DELAY + STEP)),
MIN_SSE_RECONNECT_DELAY + STEP,
"the floor must bound a hostile value, not overwrite a reasonable one — a peer \
that asks for a legitimate wait still gets the wait it asked for"
);
}
#[test]
fn a_peer_value_just_below_the_ceiling_is_honoured_verbatim() {
assert_eq!(
next_reconnect_delay(0, Some(MAX_SSE_RECONNECT_DELAY.saturating_sub(STEP))),
MAX_SSE_RECONNECT_DELAY.saturating_sub(STEP),
"the ceiling is likewise a bound, not an overwrite"
);
}
#[test]
fn a_peer_value_exactly_at_the_ceiling_is_left_alone() {
assert_eq!(
next_reconnect_delay(0, Some(MAX_SSE_RECONNECT_DELAY)),
MAX_SSE_RECONNECT_DELAY,
"the bound is inclusive at its upper end too"
);
}
#[test]
fn a_peer_value_above_the_ceiling_is_lowered_to_it() {
assert_eq!(
next_reconnect_delay(0, Some(MAX_SSE_RECONNECT_DELAY + STEP)),
MAX_SSE_RECONNECT_DELAY,
"an uncapped peer value parks a client's reader task for a duration the peer chose"
);
}
#[test]
fn a_saturating_peer_value_neither_panics_nor_escapes_the_ceiling() {
assert_eq!(
next_reconnect_delay(0, Some(Duration::MAX)),
MAX_SSE_RECONNECT_DELAY,
"`Duration::MAX` is reachable from the wire: `retry:` is parsed as u64 \
milliseconds and nothing about the parse bounds it"
);
}
#[test]
fn a_saturating_attempt_count_falls_back_to_the_ceiling() {
assert_eq!(
next_reconnect_delay(u32::MAX, None),
MAX_SSE_RECONNECT_DELAY,
"the exponential curve must SATURATE rather than overflow: an `unwrap` on the \
Duration conversion here would panic inside a client's reader task"
);
}
#[test]
fn the_computed_curve_never_falls_under_the_floor() {
for attempt in 0..8u32 {
let delay = next_reconnect_delay(attempt, None);
assert!(
delay >= MIN_SSE_RECONNECT_DELAY && delay <= MAX_SSE_RECONNECT_DELAY,
"attempt {attempt} produced {delay:?}, outside the two-sided bound"
);
}
}
#[test]
fn a_quiet_stream_that_stayed_up_earns_a_fresh_budget() {
assert!(
budget_reset_earned(RECONNECT_BUDGET_RESET_UPTIME),
"a stream that stayed up past the threshold having emitted only keep-alive \
comments is a WORKING stream: nothing else distinguishes an idle MCP session \
from a dead one, and refusing it here kills the session stream permanently"
);
assert!(
budget_reset_earned(Duration::MAX),
"and more uptime cannot make it less true"
);
}
#[test]
fn a_short_bounce_never_earns_a_fresh_budget() {
assert!(
!budget_reset_earned(Duration::ZERO),
"this is the CR-01 shape exactly: a body that ends immediately. Refunding here \
makes the reconnect loop unbounded for any budget value"
);
assert!(
!budget_reset_earned(
RECONNECT_BUDGET_RESET_UPTIME.saturating_sub(Duration::from_millis(1))
),
"one millisecond under the threshold is still a bounce — which is what keeps \
T-118.2-04-01's keeps-closing peer bounded at MAX_SSE_RECONNECT_ATTEMPTS"
);
}
#[test]
fn a_stream_that_stayed_up_earns_a_fresh_budget() {
assert!(
budget_reset_earned(RECONNECT_BUDGET_RESET_UPTIME),
"the threshold is inclusive"
);
assert!(
budget_reset_earned(RECONNECT_BUDGET_RESET_UPTIME * 120),
"a stream that worked for an hour and then blinked must not inherit a spent \
budget — that is the case D-03 exists for"
);
}
}
mod sse_reader_properties {
use super::*;
const TINY_BOUND: usize = 64;
fn assert_retention_bounded(peaks: &[usize], tails: &[usize], bound: usize) {
for (index, held) in peaks.iter().copied().enumerate() {
assert!(
held <= bound,
"the parser retained {held} bytes after chunk {index} under a {bound}-byte \
bound (peaks: {peaks:?})"
);
}
for (index, tail) in tails.iter().copied().enumerate() {
assert!(
tail <= 3,
"the undecoded UTF-8 tail was {tail} bytes after chunk {index}; the longest \
incomplete character is 3 bytes, so anything more means take_utf8_prefix \
stopped draining (tails: {tails:?})"
);
}
}
proptest::proptest! {
#[test]
fn property_arbitrary_bytes_never_panic_or_grow_the_reader(
bytes in proptest::collection::vec(proptest::prelude::any::<u8>(), 0..512),
) {
let (_outcomes, _overflowed, peaks, tails) =
decode_sse_chunks_for_fuzz(&[&bytes], TINY_BOUND);
assert_retention_bounded(&peaks, &tails, TINY_BOUND);
}
#[test]
fn property_chunked_arbitrary_bytes_never_panic_or_grow_the_reader(
bytes in proptest::collection::vec(proptest::prelude::any::<u8>(), 0..512),
) {
let chunks: Vec<&[u8]> = if bytes.is_empty() {
vec![&bytes[..]]
} else {
bytes.chunks(7).collect()
};
let (_outcomes, _overflowed, peaks, tails) =
decode_sse_chunks_for_fuzz(&chunks, TINY_BOUND);
assert_retention_bounded(&peaks, &tails, TINY_BOUND);
}
#[test]
fn property_a_valid_frame_survives_any_chunk_split(
split in 1usize..80,
) {
let frame = "id: e1\nevent: message\ndata: \
{\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\",\
\"params\":{\"progressToken\":\"t\u{00e9}\",\"progress\":1}}\n\n";
let raw = frame.as_bytes();
let at = split.min(raw.len());
let chunks: Vec<&[u8]> = vec![&raw[..at], &raw[at..]];
let (outcomes, overflowed, _peaks, tails) =
decode_sse_chunks_for_fuzz(&chunks, DEFAULT_MAX_COLLECTED_BODY_BYTES);
assert_eq!(
outcomes.len(),
1,
"a split at byte {at} yielded {} message(s), not 1",
outcomes.len()
);
assert!(
outcomes[0].is_ok(),
"a split at byte {at} corrupted the payload: {:?}",
outcomes[0]
);
assert!(
!overflowed.iter().any(|seen| *seen),
"a frame well under the bound must not overflow it"
);
assert!(
tails.iter().all(|tail| *tail <= 3),
"the undecoded UTF-8 tail must stay under one character: {tails:?}"
);
}
#[test]
fn property_next_reconnect_delay_stays_inside_both_bounds(
attempt in proptest::prelude::any::<u32>(),
retry_millis in proptest::option::of(proptest::prelude::any::<u64>()),
) {
let server_retry = retry_millis.map(Duration::from_millis);
let delay = next_reconnect_delay(attempt, server_retry);
assert!(
delay >= MIN_SSE_RECONNECT_DELAY,
"attempt {attempt} with retry {retry_millis:?} produced {delay:?}, under the \
{MIN_SSE_RECONNECT_DELAY:?} floor — an unfloored wait is a request flood"
);
assert!(
delay <= MAX_SSE_RECONNECT_DELAY,
"attempt {attempt} with retry {retry_millis:?} produced {delay:?}, over the \
{MAX_SSE_RECONNECT_DELAY:?} ceiling — an uncapped wait parks the reader"
);
}
#[test]
fn property_next_reconnect_delay_is_pure(
attempt in proptest::prelude::any::<u32>(),
retry_millis in proptest::option::of(proptest::prelude::any::<u64>()),
) {
let server_retry = retry_millis.map(Duration::from_millis);
assert_eq!(
next_reconnect_delay(attempt, server_retry),
next_reconnect_delay(attempt, server_retry),
"two identical calls disagreed for attempt {attempt}, retry {retry_millis:?}"
);
}
}
}
mod latch_gate {
use super::*;
fn no_overflow() -> RwLock<std::collections::VecDeque<Result<TransportMessage>>> {
RwLock::new(std::collections::VecDeque::new())
}
fn offline_transport() -> StreamableHttpTransport {
StreamableHttpTransport::new(
StreamableHttpTransportConfigBuilder::new(
url::Url::parse("http://127.0.0.1:1/mcp").expect("the fixture URL parses"),
)
.build(),
)
}
fn queue_capacity_tags() -> i64 {
i64::try_from(CLIENT_RECEIVE_QUEUE_CAPACITY)
.expect("the queue capacity is a small constant")
}
fn tagged(n: i64) -> TransportMessage {
TransportMessage::Response(crate::types::JSONRPCResponse {
jsonrpc: "2.0".to_string(),
id: crate::types::RequestId::Number(n),
payload: crate::types::jsonrpc::ResponsePayload::Result(serde_json::Value::Null),
})
}
async fn receive_within(
transport: &mut StreamableHttpTransport,
) -> Result<TransportMessage> {
tokio::time::timeout(Duration::from_secs(5), transport.receive())
.await
.expect("receive() must resolve — parking here IS the defect this arm fences")
}
fn tag_of(message: &TransportMessage) -> i64 {
match message {
TransportMessage::Response(response) => match &response.id {
crate::types::RequestId::Number(n) => *n,
other @ crate::types::RequestId::String(_) => {
panic!("the fixture only mints numeric ids, got {other:?}")
},
},
other => panic!("the fixture only mints responses, got {other:?}"),
}
}
#[tokio::test]
async fn a_full_queue_diverts_a_caller_send_instead_of_failing() {
let transport = offline_transport();
for tag in 0..queue_capacity_tags() {
transport
.queue_from_caller(tagged(tag))
.expect("the bounded queue accepts up to its capacity");
}
assert!(
transport.sender.try_send(Ok(tagged(-1))).is_err(),
"the bounded queue must actually be FULL, or this arm proves nothing"
);
transport.queue_from_caller(tagged(-2)).expect(
"a caller-task send must NEVER fail on a full queue: the only consumer that could \
drain it is the caller itself, which never reaches its pump if this returns Err",
);
assert_eq!(
transport.caller_overflow.read().len(),
1,
"the diverted message must be RETAINED, not dropped — D-04's never-silently-drop \
rule applies to the overflow lane too"
);
}
#[tokio::test]
async fn the_overflow_lane_is_drained_after_the_bounded_queue() {
let mut transport = offline_transport();
for tag in 0..queue_capacity_tags() {
transport
.queue_from_caller(tagged(tag))
.expect("fills the bounded queue");
}
transport
.queue_from_caller(tagged(9999))
.expect("diverts to the overflow lane");
for tag in 0..queue_capacity_tags() {
assert_eq!(
tag_of(
&receive_within(&mut transport)
.await
.expect("a queued message")
),
tag,
"the bounded queue must drain in order, and drain FIRST"
);
}
assert_eq!(
tag_of(
&receive_within(&mut transport)
.await
.expect("the overflowed message")
),
9999,
"and the overflow lane follows it, preserving global FIFO"
);
}
#[tokio::test]
async fn receive_after_close_reports_connection_closed_rather_than_parking() {
let mut transport = offline_transport();
transport.close().await.expect("close on an idle transport");
let error = receive_within(&mut transport)
.await
.expect_err("a closed transport must not hand back a message");
assert!(
matches!(error, Error::Transport(TransportError::ConnectionClosed)),
"a close is not a stream failure and must not be reported as one; got {error:?}"
);
}
#[tokio::test]
async fn a_close_still_delivers_messages_that_arrived_before_it() {
let mut transport = offline_transport();
transport
.queue_from_caller(tagged(7))
.expect("the queue is empty");
transport.close().await.expect("close");
assert_eq!(
tag_of(
&receive_within(&mut transport)
.await
.expect("the pre-close message")
),
7,
"a queued message must be delivered ahead of any reason, close included"
);
assert!(
matches!(
receive_within(&mut transport).await,
Err(Error::Transport(TransportError::ConnectionClosed))
),
"and only then does the close surface"
);
}
#[tokio::test]
async fn a_close_does_not_overwrite_an_earlier_stream_reason() {
let mut transport = offline_transport();
let (delivery, _receiver) = delivery_for(StreamKind::Session, &transport.terminal);
latch_terminal_reason(
&delivery,
&Error::Transport(TransportError::InvalidMessage(
"a corrupt frame".to_string(),
)),
);
transport.close().await.expect("close");
let error = receive_within(&mut transport)
.await
.expect_err("the latch surfaces");
assert!(
matches!(error, Error::Transport(TransportError::InvalidMessage(_))),
"the FIRST reason is the causal one and must survive a later close; got {error:?}"
);
}
#[test]
fn the_reset_seam_leaves_another_streams_reason_alone() {
for (stream, cleared) in [
(StreamKind::Session, true),
(StreamKind::PostResponse, false),
(StreamKind::Transport, false),
] {
let terminal: Arc<RwLock<Option<TerminalReason>>> = Arc::new(RwLock::new(None));
let (delivery, _receiver) = delivery_for(stream, &terminal);
latch_terminal_reason(
&delivery,
&Error::Transport(TransportError::Request("a reason".to_string())),
);
{
let mut slot = terminal.write();
if slot
.as_ref()
.is_some_and(|reason| reason.stream == StreamKind::Session)
{
*slot = None;
}
}
assert_eq!(
terminal.read().is_none(),
cleared,
"a re-opened SESSION stream forgives its OWN reason and nothing else; \
{stream:?} was handled wrongly"
);
}
}
fn delivery_for(
stream: StreamKind,
terminal: &Arc<RwLock<Option<TerminalReason>>>,
) -> (ReaderDelivery, mpsc::Receiver<Result<TransportMessage>>) {
let (sender, receiver) = mpsc::channel(CLIENT_RECEIVE_QUEUE_CAPACITY);
let (terminal_signal, _) = watch::channel(0u64);
let (shutdown, _) = watch::channel(false);
(
ReaderDelivery {
sender,
stream,
terminal: Arc::clone(terminal),
terminal_signal: Arc::new(terminal_signal),
shutdown: Arc::new(shutdown),
},
receiver,
)
}
fn wake_signal() -> Arc<watch::Sender<u64>> {
let (signal, _) = watch::channel(0u64);
Arc::new(signal)
}
fn budget_reason() -> Error {
reconnect_budget_exhausted(MAX_SSE_RECONNECT_ATTEMPTS, None)
}
#[test]
fn latch_gate_boundary_at_zero_and_one_in_flight_readers() {
let terminal: Arc<RwLock<Option<TerminalReason>>> = Arc::new(RwLock::new(None));
let (delivery, mut receiver) = delivery_for(StreamKind::Session, &terminal);
let open_post_readers = Arc::new(AtomicUsize::new(0));
let wake = wake_signal();
latch_terminal_reason(&delivery, &budget_reason());
let at_zero =
drain_or_latch(&mut receiver, &no_overflow(), &terminal, &open_post_readers);
assert!(
matches!(at_zero, Some(Err(_))),
"with an empty queue, a set latch and NO reader in flight, the latched reason \
must be surfaced — that is the CR-02 contract and this gate does not weaken it"
);
let guard = PostReaderGuard::acquire(&open_post_readers, &wake);
assert_eq!(
open_post_readers.load(Ordering::SeqCst),
1,
"the guard must count itself the moment it is acquired, synchronously"
);
let at_one =
drain_or_latch(&mut receiver, &no_overflow(), &terminal, &open_post_readers);
assert!(
at_one.is_none(),
"with a POST-response reader still live, `drain_or_latch` must answer None and \
the caller must keep waiting. Surfacing the latch here hands a caller another \
stream's diagnosis as its own result (BLOCKER 1)"
);
drop(guard);
let after =
drain_or_latch(&mut receiver, &no_overflow(), &terminal, &open_post_readers);
assert!(
matches!(after, Some(Err(_))),
"once the last reader is gone the latch must be surfaced again, or the gate has \
traded a permanent failure for a permanent hang (T-118.2-19-03)"
);
}
#[test]
fn a_queued_message_still_wins_over_a_set_latch() {
let terminal: Arc<RwLock<Option<TerminalReason>>> = Arc::new(RwLock::new(None));
let (delivery, mut receiver) = delivery_for(StreamKind::Session, &terminal);
let open_post_readers = Arc::new(AtomicUsize::new(0));
latch_terminal_reason(&delivery, &budget_reason());
delivery
.sender
.try_send(Ok(TransportMessage::Notification(
crate::types::Notification::Client(
crate::types::ClientNotification::Initialized,
),
)))
.expect("the bounded queue has capacity for one message");
let drained =
drain_or_latch(&mut receiver, &no_overflow(), &terminal, &open_post_readers);
assert!(
matches!(drained, Some(Ok(_))),
"the queue is drained before the latch is consulted; a message that arrived \
before the failure must still be delivered ahead of it"
);
}
#[test]
fn post_reader_guard_returns_the_count_to_zero() {
let counter = Arc::new(AtomicUsize::new(0));
let wake = wake_signal();
{
let _empty = PostReaderGuard::acquire(&counter, &wake);
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
assert_eq!(
counter.load(Ordering::SeqCst),
0,
"a POST-response reader that delivered zero events must still return the count \
to zero, or an empty answer permanently gates the transport"
);
{
let _first = PostReaderGuard::acquire(&counter, &wake);
let _second = PostReaderGuard::acquire(&counter, &wake);
assert_eq!(counter.load(Ordering::SeqCst), 2);
{
let _third = PostReaderGuard::acquire(&counter, &wake);
assert_eq!(counter.load(Ordering::SeqCst), 3);
}
assert_eq!(
counter.load(Ordering::SeqCst),
2,
"one reader finishing must not clear the others"
);
}
assert_eq!(
counter.load(Ordering::SeqCst),
0,
"every nested guard must unwind to zero"
);
let unwound = Arc::clone(&counter);
let unwound_wake = Arc::clone(&wake);
let panicked = std::panic::catch_unwind(std::panic::AssertUnwindSafe(move || {
let _guard = PostReaderGuard::acquire(&unwound, &unwound_wake);
panic!("a reader failing mid-body");
}));
assert!(panicked.is_err(), "the arm must actually have panicked");
assert_eq!(
counter.load(Ordering::SeqCst),
0,
"a reader that panicked mid-body must still return the count to zero — that is \
why the count is RAII and not an explicit decrement at each exit"
);
}
#[test]
fn the_last_post_reader_out_wakes_a_parked_consumer() {
let counter = Arc::new(AtomicUsize::new(0));
let wake = wake_signal();
let mut observer = wake.subscribe();
let before = *observer.borrow_and_update();
let first = PostReaderGuard::acquire(&counter, &wake);
let second = PostReaderGuard::acquire(&counter, &wake);
drop(first);
assert_eq!(
*wake.borrow(),
before,
"a reader finishing while ANOTHER is still live changes nothing a consumer \
could act on — the gate is still closed, so waking would be a spurious \
re-poll"
);
drop(second);
assert_eq!(
*wake.borrow(),
before + 1,
"the LAST reader out must bump the generation, or a consumer parked while the \
gate was closed never learns that it re-opened. A clean reader exit latches \
nothing, so this is the only wake on that path"
);
assert_eq!(
counter.load(Ordering::SeqCst),
0,
"and the count must be zero when that wake is raised, so the woken consumer \
re-reads an OPEN gate rather than being sent back to sleep"
);
}
#[test]
fn clearing_an_already_clear_latch_is_a_no_op() {
let terminal: Arc<RwLock<Option<TerminalReason>>> = Arc::new(RwLock::new(None));
let (delivery, mut receiver) = delivery_for(StreamKind::Session, &terminal);
let open_post_readers = Arc::new(AtomicUsize::new(0));
latch_terminal_reason(&delivery, &budget_reason());
assert!(terminal.read().is_some(), "the latch is set");
*terminal.write() = None;
assert!(terminal.read().is_none(), "one reset clears it");
*terminal.write() = None;
assert!(
terminal.read().is_none(),
"clearing an already-clear latch must be a no-op"
);
assert!(
drain_or_latch(&mut receiver, &no_overflow(), &terminal, &open_post_readers)
.is_none(),
"a recovered transport with an empty queue and nothing in flight must WAIT, not \
answer the stale reason"
);
let (post_delivery, _post_receiver) = delivery_for(StreamKind::PostResponse, &terminal);
latch_terminal_reason(
&post_delivery,
&Error::Transport(TransportError::Request("a fresh reason".to_string())),
);
let resurrected = terminal.read().clone().expect("the fresh reason is stored");
assert_eq!(
resurrected.stream,
StreamKind::PostResponse,
"after a reset, the next write wins outright — the cleared reason must not come \
back"
);
assert!(
resurrected.message.contains("a fresh reason"),
"got {:?}",
resurrected.message
);
}
#[test]
fn write_once_holds_under_two_racing_latch_writers() {
let terminal: Arc<RwLock<Option<TerminalReason>>> = Arc::new(RwLock::new(None));
let (session, _session_rx) = delivery_for(StreamKind::Session, &terminal);
let (post, _post_rx) = delivery_for(StreamKind::PostResponse, &terminal);
let barrier = Arc::new(std::sync::Barrier::new(2));
let session_barrier = Arc::clone(&barrier);
let post_barrier = Arc::clone(&barrier);
let session_writer = std::thread::spawn(move || {
session_barrier.wait();
latch_terminal_reason(&session, &budget_reason());
});
let post_writer = std::thread::spawn(move || {
post_barrier.wait();
latch_terminal_reason(
&post,
&Error::Transport(TransportError::Request(
"the POST response stream dropped".to_string(),
)),
);
});
session_writer.join().expect("the session writer finishes");
post_writer.join().expect("the POST writer finishes");
let stored = terminal.read().clone().expect("one of them won");
let is_one_of_the_two = (stored.stream == StreamKind::Session
&& stored.message.contains("reconnect budget"))
|| (stored.stream == StreamKind::PostResponse
&& stored.message.contains("the POST response stream dropped"));
assert!(
is_one_of_the_two,
"exactly ONE reason must be stored, whole and unmixed: a slot holding one \
writer's stream kind beside the other's message would tell a caller a \
falsehood about which stream ended. Got {stored:?}"
);
}
#[test]
fn a_post_reader_in_flight_gates_a_reason_from_either_stream() {
for stream in [StreamKind::Session, StreamKind::PostResponse] {
let terminal: Arc<RwLock<Option<TerminalReason>>> = Arc::new(RwLock::new(None));
let (delivery, mut receiver) = delivery_for(stream, &terminal);
let open_post_readers = Arc::new(AtomicUsize::new(0));
let wake = wake_signal();
let _guard = PostReaderGuard::acquire(&open_post_readers, &wake);
latch_terminal_reason(&delivery, &budget_reason());
assert!(
drain_or_latch(&mut receiver, &no_overflow(), &terminal, &open_post_readers)
.is_none(),
"a {stream:?} reason must not pre-empt a caller whose own POST-response \
stream is still live"
);
}
}
#[test]
fn to_error_names_the_stream_and_preserves_the_message_body() {
let session = terminal_reason_of(&budget_reason(), StreamKind::Session);
let rendered = session.to_error().to_string();
assert!(
rendered.contains("reconnect budget"),
"the message BODY must survive the stream-name prefix verbatim, or fences 12 \
and 13 stop measuring what they were written to measure. Got {rendered:?}"
);
assert!(
rendered.contains("the GET session stream"),
"a caller must be able to tell an unrelated stream's diagnosis from its own \
(T-118.2-19-02). Got {rendered:?}"
);
assert!(
matches!(
session.to_error(),
Error::Transport(TransportError::Request(_))
),
"a spent budget is a LIFECYCLE end, not corruption; the variant must not move"
);
let post = terminal_reason_of(
&unparseable_sse_frame(
&Error::Transport(TransportError::InvalidMessage("bad json".to_string())),
"{",
),
StreamKind::PostResponse,
);
let post_rendered = post.to_error().to_string();
assert!(
post_rendered.contains("this call's own POST response stream"),
"the POST half must name itself distinctly from the session stream. Got \
{post_rendered:?}"
);
assert!(
matches!(
post.to_error(),
Error::Transport(TransportError::InvalidMessage(_))
),
"a parse failure is CORRUPTION; the D-02/D-05 taxonomy must not move"
);
assert_ne!(
rendered, post_rendered,
"two streams ending for different reasons must render differently"
);
}
proptest::proptest! {
#[test]
fn property_the_in_flight_count_is_exactly_the_guards_held(
held in 0usize..32,
) {
let counter = Arc::new(AtomicUsize::new(0));
let wake = wake_signal();
let mut guards = Vec::with_capacity(held);
for expected in 1..=held {
guards.push(PostReaderGuard::acquire(&counter, &wake));
assert_eq!(
counter.load(Ordering::SeqCst),
expected,
"the count must equal the guards held at every prefix"
);
}
while let Some(guard) = guards.pop() {
let before = counter.load(Ordering::SeqCst);
drop(guard);
assert_eq!(
counter.load(Ordering::SeqCst),
before - 1,
"each drop must return exactly one"
);
}
assert_eq!(
counter.load(Ordering::SeqCst),
0,
"every guard dropped must leave the count at exactly zero, for any number \
of concurrently outstanding streaming POSTs"
);
}
}
}
#[derive(Debug)]
struct CachingProbe {
cached: StdMutex<Option<String>>,
vends: AtomicUsize,
purges: AtomicUsize,
minted: AtomicUsize,
release: watch::Sender<bool>,
}
impl CachingProbe {
fn primed(token: &str) -> Self {
let (release, _) = watch::channel(false);
Self {
cached: StdMutex::new(Some(token.to_string())),
vends: AtomicUsize::new(0),
purges: AtomicUsize::new(0),
minted: AtomicUsize::new(0),
release,
}
}
}
#[async_trait]
impl AuthProvider for CachingProbe {
async fn get_access_token(&self) -> Result<String> {
let cached = self.cached.lock().unwrap().clone();
if let Some(token) = cached {
return Ok(token);
}
self.vends.fetch_add(1, Ordering::SeqCst);
let mut release = self.release.subscribe();
if !*release.borrow() {
let _ = release.wait_for(|open| *open).await;
}
let serial = self.minted.fetch_add(1, Ordering::SeqCst) + 1;
let token = format!("vended-{serial}");
*self.cached.lock().unwrap() = Some(token.clone());
Ok(token)
}
async fn on_unauthorized(&self) -> Result<()> {
self.purges.fetch_add(1, Ordering::SeqCst);
*self.cached.lock().unwrap() = None;
tokio::task::yield_now().await;
Ok(())
}
}
async fn spawn_401_then_ok_listener(
unauthorized: usize,
refresh_lock: Arc<tokio::sync::Mutex<()>>,
seen: Arc<StdMutex<Vec<(u16, bool)>>>,
) -> Url {
use hyper::service::service_fn;
use hyper_util::server::conn::auto::Builder as ServerBuilder;
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
let refresh_lock = Arc::clone(&refresh_lock);
let seen = Arc::clone(&seen);
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 refresh_lock = Arc::clone(&refresh_lock);
let seen = Arc::clone(&seen);
async move {
let _ = req.collect().await;
let lock_free = refresh_lock.try_lock().is_ok();
let status = {
let mut seen = seen.lock().unwrap();
let status = if seen.len() < unauthorized {
401u16
} else {
200u16
};
seen.push((status, lock_free));
status
};
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":42,"result":{}}"#
} else {
r#"{"error":"unauthorized"}"#
})))
.unwrap(),
)
}
}),
)
.await;
});
}
});
Url::parse(&format!("http://127.0.0.1:{}", addr.port())).unwrap()
}
type RecordedRequest = (String, String, Vec<(String, String)>, Vec<u8>);
async fn spawn_recording_listener(seen: Arc<StdMutex<Vec<RecordedRequest>>>) -> Url {
use hyper::service::service_fn;
use hyper_util::server::conn::auto::Builder as ServerBuilder;
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
let seen = Arc::clone(&seen);
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 seen = Arc::clone(&seen);
async move {
let method = req.method().to_string();
let path = req.uri().path().to_string();
let mut headers: Vec<(String, String)> = req
.headers()
.iter()
.map(|(name, value)| {
(
name.as_str().to_string(),
value.to_str().unwrap_or("<binary>").to_string(),
)
})
.collect();
headers.sort();
let body = req.collect().await.unwrap().to_bytes().to_vec();
seen.lock().unwrap().push((method, path, headers, body));
Ok::<_, hyper::Error>(
HyperResponse::builder()
.status(200u16)
.header("content-type", "application/json")
.body(Full::new(Bytes::from(
r#"{"jsonrpc":"2.0","id":42,"result":{}}"#,
)))
.unwrap(),
)
}
}),
)
.await;
});
}
});
Url::parse(&format!("http://127.0.0.1:{}", addr.port())).unwrap()
}
#[tokio::test]
async fn the_shared_handle_writes_the_same_request_as_the_exclusive_path() {
let seen: Arc<StdMutex<Vec<RecordedRequest>>> = Arc::new(StdMutex::new(Vec::new()));
let url = spawn_recording_listener(Arc::clone(&seen)).await;
let mut transport = make_transport(url, None);
transport
.send(list_tools_message())
.await
.expect("the exclusive path reaches the listener");
let handle = transport
.shared_sender()
.expect("this transport offers a shared-send path");
handle
.send_shared(list_tools_message())
.await
.expect("the shared path reaches the listener");
let recorded = seen.lock().unwrap().clone();
assert_eq!(
recorded.len(),
2,
"both sends must have reached the wire, or this fence compares nothing"
);
assert_eq!(
recorded[0], recorded[1],
"a frame sent through the shared handle must produce a BYTE-IDENTICAL request to one \
sent through the exclusive `&mut` path — same method, same path, same header block, \
same body. A difference here means the handle is a second, hand-rolled POST path \
rather than the same core reached differently (T-118.2-23-03)"
);
}
#[test]
fn stdio_offers_no_shared_send_path() {
let transport = crate::shared::StdioTransport::new();
assert!(
transport.shared_sender().is_none(),
"StdioTransport owns its own I/O and must keep the default `None`, so a client over \
it sends through the exclusive `&mut` path byte-for-byte as it does today"
);
}
#[tokio::test]
async fn a_solo_401_recovery_purges_once_and_bumps_the_generation_once() {
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 provider = Arc::new(CountingProvider::new("initial-token"));
let url = Url::parse(&server.url()).unwrap();
let mut transport = make_transport(url, Some(provider.clone() as Arc<dyn AuthProvider>));
let _ = transport.send(ping_message()).await;
assert_eq!(
provider.unauthorized_count.load(Ordering::SeqCst),
1,
"a solo caller purges exactly once; the second 401 — on the retry — is returned \
unchanged, which is STRUCTURAL and not something the new lock provides"
);
assert_eq!(
transport.token_generation.load(Ordering::SeqCst),
1,
"one completed refresh must move the generation exactly one step"
);
assert!(
transport.refresh_lock.try_lock().is_ok(),
"the refresh lock must be released on every exit path, including the error one"
);
}
#[tokio::test]
async fn a_401_with_no_provider_is_returned_unchanged_and_moves_no_generation() {
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(1)
.create_async()
.await;
let url = Url::parse(&server.url()).unwrap();
let mut transport = make_transport(url, None);
let result = transport.send(ping_message()).await;
assert!(
result.is_err(),
"a 401 with no provider is an ordinary failed request"
);
assert_eq!(
transport.token_generation.load(Ordering::SeqCst),
0,
"the no-provider path returns BEFORE the guarded region, so nothing may move"
);
assert!(
transport.refresh_lock.try_lock().is_ok(),
"the no-provider path must not take the lock at all"
);
}
#[tokio::test]
async fn a_401_on_a_current_generation_still_refreshes() {
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(4)
.create_async()
.await;
let provider = Arc::new(CountingProvider::new("initial-token"));
let url = Url::parse(&server.url()).unwrap();
let mut transport = make_transport(url, Some(provider.clone() as Arc<dyn AuthProvider>));
let _ = transport.send(ping_message()).await;
let _ = transport.send(ping_message()).await;
assert_eq!(
provider.unauthorized_count.load(Ordering::SeqCst),
2,
"each caller whose captured generation is CURRENT must refresh; skipping the second \
would leave it presenting an invalid token forever"
);
assert_eq!(
transport.token_generation.load(Ordering::SeqCst),
2,
"two genuinely new 401s are two refreshes, so two generation steps"
);
}
#[tokio::test]
async fn the_retry_post_is_sent_outside_the_refresh_lock() {
let seen = Arc::new(StdMutex::new(Vec::new()));
let provider = Arc::new(CountingProvider::new("initial-token"));
let refresh_lock = Arc::new(tokio::sync::Mutex::new(()));
let url = spawn_401_then_ok_listener(1, Arc::clone(&refresh_lock), Arc::clone(&seen)).await;
let mut transport = make_transport(url, Some(provider.clone() as Arc<dyn AuthProvider>));
transport.refresh_lock = refresh_lock;
let _ = transport
.send_with_options(list_tools_message(), SendOptions::default())
.await;
let seen = seen.lock().unwrap().clone();
assert_eq!(
seen.len(),
2,
"expected the original POST and exactly one retry; observed {seen:?}"
);
assert_eq!(seen[0].0, 401, "the first attempt is the 401");
assert_eq!(seen[1].0, 200, "the retry is served normally");
assert!(
seen[1].1,
"the RETRY POST must be on the wire with the refresh lock RELEASED. A held lock here \
means the guarded region was drawn around the send as well as the build, which \
re-creates the whole-transport bottleneck the client-side guard is being removed to \
escape"
);
}
#[tokio::test]
async fn two_concurrent_401s_take_one_purge_and_one_vend() {
let seen = Arc::new(StdMutex::new(Vec::new()));
let provider = Arc::new(CachingProbe::primed("primed-token"));
let refresh_lock = Arc::new(tokio::sync::Mutex::new(()));
let url = spawn_401_then_ok_listener(2, Arc::clone(&refresh_lock), Arc::clone(&seen)).await;
let mut transport = make_transport(url, Some(provider.clone() as Arc<dyn AuthProvider>));
transport.refresh_lock = Arc::clone(&refresh_lock);
let mut one = transport.clone();
let mut two = transport.clone();
let first = tokio::spawn(async move {
one.send_with_options(list_tools_message(), SendOptions::default())
.await
});
let second = tokio::spawn(async move {
two.send_with_options(list_tools_message(), SendOptions::default())
.await
});
let deadline = tokio::time::Instant::now() + Duration::from_millis(750);
loop {
let answered_401 = seen.lock().unwrap().iter().filter(|e| e.0 == 401).count();
if answered_401 >= 2 || tokio::time::Instant::now() >= deadline {
break;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
let deadline = tokio::time::Instant::now() + Duration::from_millis(750);
while provider.vends.load(Ordering::SeqCst) < 2 {
if tokio::time::Instant::now() >= deadline {
break;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
let _ = provider.release.send_replace(true);
let _ = first.await.unwrap();
let _ = second.await.unwrap();
assert_eq!(
provider.vends.load(Ordering::SeqCst),
1,
"exactly ONE vend may occur across the whole recovery — the loser's retry token comes \
from the cache the winner warmed, not from a second round trip to the IdP"
);
assert_eq!(
provider.purges.load(Ordering::SeqCst),
1,
"the loser's captured generation was already superseded, so it has nothing to purge"
);
assert_eq!(
transport.token_generation.load(Ordering::SeqCst),
1,
"one refresh, one generation step"
);
}
}