use self::mcp_session::McpSession;
use crate::{
error::{Error, ErrorCode},
transport::http::{ClientRuntimeContext, MCP_SESSION_ID, get_mcp_session_id},
types::Message,
};
use futures_util::{StreamExt, TryStreamExt};
use reqwest::header::{CACHE_CONTROL, HeaderName};
use reqwest::{
RequestBuilder,
header::{ACCEPT, CONTENT_TYPE},
};
use std::sync::Arc;
use std::time::Duration;
#[cfg(feature = "client-tls")]
use tls_config::ClientTlsConfig;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
pub(super) mod mcp_session;
#[cfg(feature = "client-oauth")]
pub(crate) mod oauth;
#[cfg(feature = "client-tls")]
pub(crate) mod tls_config;
#[derive(Clone)]
enum ClientAuth {
None,
Static(Arc<str>),
#[cfg(feature = "client-oauth")]
OAuth(Arc<oauth::OAuthSession>),
}
impl ClientAuth {
async fn fresh_bearer(&self) -> Option<Arc<str>> {
match self {
ClientAuth::None => None,
ClientAuth::Static(token) => Some(token.clone()),
#[cfg(feature = "client-oauth")]
ClientAuth::OAuth(session) => session.refreshed_bearer().await,
}
}
fn from_static(access_token: Option<Box<[u8]>>) -> Self {
match access_token {
Some(token) => ClientAuth::Static(String::from_utf8_lossy(&token).into()),
None => ClientAuth::None,
}
}
}
const LAST_EVENT_ID: HeaderName = HeaderName::from_static("last-event-id");
const SSE_RECONNECT_DELAY: Duration = Duration::from_secs(3);
const STREAM_ENDED_BEFORE_RESPONSE: &str = "POST SSE stream ended before the response arrived";
#[cfg(not(feature = "legacy-spec"))]
fn routing_hints(msg: &Message) -> Option<(&str, Option<String>)> {
match msg {
Message::Request(r) => Some((r.method.as_str(), name_param(r))),
Message::Notification(n) => Some((n.method.as_str(), None)),
Message::Batch(_) | Message::Response(_) => None,
}
}
#[cfg(not(feature = "legacy-spec"))]
fn name_param(req: &crate::types::Request) -> Option<String> {
#[cfg(feature = "tasks")]
{
use crate::types::task::commands as tasks;
if matches!(
req.method.as_str(),
tasks::GET | tasks::UPDATE | tasks::CANCEL
) {
let raw = req.params.as_ref()?.as_object()?.get("taskId")?.as_str()?;
return Some(crate::transport::http::encode_header_value(raw));
}
}
let field = match req.method.as_str() {
crate::types::tool::commands::CALL | crate::types::prompt::commands::GET => "name",
crate::types::resource::commands::READ => "uri",
_ => return None,
};
let raw = req.params.as_ref()?.as_object()?.get(field)?.as_str()?;
Some(crate::transport::http::encode_header_value(raw))
}
#[cfg(not(feature = "legacy-spec"))]
fn param_headers(
msg: &Message,
registry: &crate::shared::param_headers::Registry,
) -> Vec<(String, String)> {
let Message::Request(req) = msg else {
return Vec::new();
};
if req.method != crate::types::tool::commands::CALL {
return Vec::new();
}
let Some(params) = req.params.as_ref().and_then(|p| p.as_object()) else {
return Vec::new();
};
let Some(name) = params.get("name").and_then(|n| n.as_str()) else {
return Vec::new();
};
let Some(entry) = registry.get(name) else {
return Vec::new();
};
let Some(headers) = entry.usable() else {
return Vec::new();
};
let args = params.get("arguments").cloned().unwrap_or_default();
crate::shared::param_headers::extract(headers, &args)
}
pub(super) async fn connect(rt: ClientRuntimeContext, token: CancellationToken) {
let session = Arc::new(McpSession::new(
rt.url,
token,
#[cfg(not(feature = "legacy-spec"))]
rt.peer_mode.clone(),
));
#[cfg(feature = "client-oauth")]
let auth = match rt.oauth {
Some(oauth) => ClientAuth::OAuth(oauth),
None => ClientAuth::from_static(rt.access_token),
};
#[cfg(not(feature = "client-oauth"))]
let auth = ClientAuth::from_static(rt.access_token);
tokio::join!(
handle_connection(
session.clone(),
rt.rx,
rt.tx.clone(),
auth.clone(),
#[cfg(not(feature = "legacy-spec"))]
rt.param_headers.clone(),
#[cfg(feature = "client-tls")]
rt.tls_config.clone()
),
start_sse_connection(
session.clone(),
rt.tx.clone(),
auth.clone(),
#[cfg(feature = "client-tls")]
rt.tls_config.clone()
)
);
}
async fn handle_connection(
session: Arc<McpSession>,
mut sender_rx: mpsc::Receiver<Message>,
recv_tx: mpsc::Sender<Result<Message, Error>>,
auth: ClientAuth,
#[cfg(not(feature = "legacy-spec"))] param_registry: crate::shared::param_headers::Registry,
#[cfg(feature = "client-tls")] tls_config: Option<ClientTlsConfig>,
) {
#[cfg(not(feature = "client-tls"))]
let client = match create_client() {
Ok(client) => client,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "HTTP client error: {_err:#}");
return;
}
};
#[cfg(feature = "client-tls")]
let client = match create_client(tls_config) {
Ok(client) => client,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "HTTP client error: {_err:#}");
return;
}
};
let token = session.cancellation_token();
loop {
tokio::select! {
biased;
_ = token.cancelled() => return,
req = sender_rx.recv() => {
let Some(req) = req else {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Unexpected messaging error");
break;
};
#[cfg(not(feature = "legacy-spec"))]
abort_cancelled_stream(&req, &session);
#[cfg(not(feature = "legacy-spec"))]
let abort = track_listen(&req, &session);
crate::spawn_fair!(send_request(
client.clone(),
session.clone(),
req,
recv_tx.clone(),
auth.clone(),
#[cfg(not(feature = "legacy-spec"))]
param_registry.clone(),
#[cfg(not(feature = "legacy-spec"))]
abort,
));
}
}
}
}
#[cfg(not(feature = "legacy-spec"))]
fn mirrored_param_headers(
session: &McpSession,
req: &Message,
registry: &crate::shared::param_headers::Registry,
) -> Vec<(String, String)> {
if session.is_legacy() {
return Vec::new();
}
param_headers(req, registry)
}
fn build_post(
client: &reqwest::Client,
session: &McpSession,
req: &Message,
bearer: Option<&str>,
#[cfg(not(feature = "legacy-spec"))] mirrored: &[(String, String)],
) -> RequestBuilder {
let mut resp = client
.post(session.url())
.json(req)
.header(ACCEPT, "application/json, text/event-stream");
if let Some(session_id) = session.session_id() {
resp = resp.header(MCP_SESSION_ID, session_id.to_string())
}
#[cfg(not(feature = "legacy-spec"))]
if !session.is_legacy() {
if let Some((method, name)) = routing_hints(req) {
resp = resp.header(crate::transport::http::MCP_METHOD, method);
if let Some(n) = name {
resp = resp.header(crate::transport::http::MCP_NAME, n);
}
}
for (name, value) in mirrored {
resp = resp.header(
name.as_str(),
crate::transport::http::encode_header_value(value),
);
}
resp = resp.header(
crate::transport::http::MCP_PROTOCOL_VERSION,
crate::LATEST_PROTOCOL_VERSION,
);
}
if let Some(bearer) = bearer {
resp = resp.bearer_auth(bearer)
}
resp
}
async fn send_request(
client: reqwest::Client,
session: Arc<McpSession>,
req: Message,
resp_tx: mpsc::Sender<Result<Message, Error>>,
auth: ClientAuth,
#[cfg(not(feature = "legacy-spec"))] param_registry: crate::shared::param_headers::Registry,
#[cfg(not(feature = "legacy-spec"))] abort: ListenAbort,
) {
#[cfg(not(feature = "legacy-spec"))]
if abort.is_tracked() {
let session_token = session.cancellation_token();
tokio::select! {
_ = exchange(client, session, req, resp_tx, auth, param_registry) => {}
_ = abort.cancelled() => {}
_ = session_token.cancelled() => {}
}
return;
}
exchange(
client,
session,
req,
resp_tx,
auth,
#[cfg(not(feature = "legacy-spec"))]
param_registry,
)
.await
}
async fn exchange(
client: reqwest::Client,
session: Arc<McpSession>,
req: Message,
resp_tx: mpsc::Sender<Result<Message, Error>>,
auth: ClientAuth,
#[cfg(not(feature = "legacy-spec"))] param_registry: crate::shared::param_headers::Registry,
) {
let bearer = auth.fresh_bearer().await;
#[cfg(not(feature = "legacy-spec"))]
let mirrored = mirrored_param_headers(&session, &req, ¶m_registry);
let sent = build_post(
&client,
&session,
&req,
bearer.as_deref(),
#[cfg(not(feature = "legacy-spec"))]
&mirrored,
)
.send()
.await;
let resp = match sent {
Ok(resp) => resp,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to send HTTP request: {}", _err);
return;
}
};
#[cfg(feature = "client-oauth")]
let resp = match (&auth, resp.status()) {
(ClientAuth::OAuth(oauth), status)
if status == reqwest::StatusCode::UNAUTHORIZED
|| (status == reqwest::StatusCode::FORBIDDEN
&& insufficient_scope(resp.headers())) =>
{
let challenge = bearer_challenge(resp.headers());
match oauth
.authorize(challenge.as_deref(), bearer.as_deref())
.await
{
Ok(fresh) => {
let retried = build_post(
&client,
&session,
&req,
Some(&fresh),
#[cfg(not(feature = "legacy-spec"))]
&mirrored,
)
.send()
.await;
match retried {
Ok(retried) => retried,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(
logger = "neva",
"Failed to resend HTTP request: {}",
_err
);
return;
}
}
}
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "OAuth authorization failed: {}", _err);
resp
}
}
}
_ => resp,
};
if let Message::Notification(_) = &req {
return;
}
if let Message::Batch(ref batch) = req
&& !batch.has_requests()
{
return;
}
if !session.has_session_id()
&& let Some(session_id) = get_mcp_session_id(resp.headers())
{
session.set_session_id(session_id);
}
if let Message::Request(r) = &req
&& r.method == crate::commands::INIT
{
let token = session.cancellation_token();
session.notify_session_initialized();
tokio::select! {
biased;
_ = token.cancelled() => return,
_ = session.sse_ready() => {},
}
}
let status = resp.status();
if is_event_stream(resp.headers()) {
let stream = sse_stream::SseStream::from_bytes_stream(resp.bytes_stream());
let ids = request_ids(&req);
let Drained {
mut owed,
last_event_id,
retry,
} = drain_post_sse(stream, &resp_tx, &ids).await;
if !owed.is_empty()
&& resumable(&session)
&& let Some(last_id) = last_event_id
{
owed = resume_stream(&client, &session, &auth, &last_id, retry, &resp_tx, &owed).await;
}
if !owed.is_empty() {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", STREAM_ENDED_BEFORE_RESPONSE);
for id in owed {
let resp = crate::types::Response::error(
id,
Error::new(ErrorCode::InternalError, STREAM_ENDED_BEFORE_RESPONSE),
);
if resp_tx.send(Ok(Message::Response(resp))).await.is_err() {
break;
}
}
}
return;
}
match resp.json::<Message>().await {
Ok(msg) => {
if let Err(_err) = resp_tx.send(Ok(msg)).await {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to send response: {}", _err);
}
}
Err(err) => {
#[cfg(feature = "tracing")]
tracing::error!(
logger = "neva",
"Failed to parse HTTP response ({}): {}",
status,
err
);
let (code, reason) = parse_failure(status, &err);
for id in request_ids(&req) {
let resp = crate::types::Response::error(id, Error::new(code, reason.clone()));
if resp_tx.send(Ok(Message::Response(resp))).await.is_err() {
break;
}
}
}
}
}
#[cfg(not(feature = "legacy-spec"))]
struct ListenAbort {
tokens: Vec<CancellationToken>,
session: Arc<McpSession>,
ids: Vec<crate::types::RequestId>,
}
#[cfg(not(feature = "legacy-spec"))]
impl ListenAbort {
fn is_tracked(&self) -> bool {
!self.tokens.is_empty()
}
async fn cancelled(&self) {
if self.tokens.is_empty() {
std::future::pending::<()>().await;
}
futures_util::future::select_all(self.tokens.iter().map(|t| Box::pin(t.cancelled()))).await;
}
}
#[cfg(not(feature = "legacy-spec"))]
impl Drop for ListenAbort {
fn drop(&mut self) {
for id in &self.ids {
self.session.untrack_stream(id);
}
}
}
#[cfg(not(feature = "legacy-spec"))]
fn track_listen(req: &Message, session: &Arc<McpSession>) -> ListenAbort {
let ids = match req {
Message::Request(r) if r.method == crate::types::subscription::commands::LISTEN => {
request_ids(req)
}
_ => Vec::new(),
};
let tokens = ids
.iter()
.map(|id| session.track_stream(id.clone()))
.collect();
ListenAbort {
tokens,
session: session.clone(),
ids,
}
}
#[cfg(not(feature = "legacy-spec"))]
fn abort_cancelled_stream(msg: &Message, session: &McpSession) {
let Message::Notification(notification) = msg else {
return;
};
if notification.method != crate::types::notification::commands::CANCELLED {
return;
}
if let Some(id) = notification
.params
.as_ref()
.and_then(|p| p.get("requestId"))
.and_then(|v| serde_json::from_value::<crate::types::RequestId>(v.clone()).ok())
{
session.abort_stream(&id);
}
}
#[inline]
fn parse_failure(status: reqwest::StatusCode, err: &impl std::fmt::Display) -> (ErrorCode, String) {
let unsupported_route = matches!(status.as_u16(), 400 | 404 | 405 | 406);
let code = if status.is_success() || unsupported_route {
ErrorCode::ParseError
} else {
ErrorCode::InternalError
};
(code, format!("HTTP {status}: {err}"))
}
fn is_event_stream(headers: &reqwest::header::HeaderMap) -> bool {
headers
.get(CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.and_then(|ct| ct.split(';').next())
.is_some_and(|essence| essence.trim().eq_ignore_ascii_case("text/event-stream"))
}
fn request_ids(msg: &Message) -> Vec<crate::types::RequestId> {
match msg {
Message::Request(r) => vec![r.id()],
Message::Batch(batch) => batch
.iter()
.filter_map(|envelope| match envelope {
crate::types::MessageEnvelope::Request(r) => Some(r.id()),
_ => None,
})
.collect(),
_ => Vec::new(),
}
}
async fn start_sse_connection(
session: Arc<McpSession>,
resp_tx: mpsc::Sender<Result<Message, Error>>,
auth: ClientAuth,
#[cfg(feature = "client-tls")] tls_config: Option<ClientTlsConfig>,
) {
let token = session.cancellation_token();
tokio::select! {
biased;
_ = token.cancelled() => (),
_ = session.initialized() => {
tokio::spawn(handle_sse_connection(
session.clone(),
resp_tx,
auth,
#[cfg(feature = "client-tls")]
tls_config
));
}
}
}
async fn handle_sse_connection(
session: Arc<McpSession>,
resp_tx: mpsc::Sender<Result<Message, Error>>,
auth: ClientAuth,
#[cfg(feature = "client-tls")] tls_config: Option<ClientTlsConfig>,
) {
#[cfg(not(feature = "client-tls"))]
let client = match create_client() {
Ok(client) => client,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "SSE client error: {_err:#}");
return;
}
};
#[cfg(feature = "client-tls")]
let client = match create_client(tls_config) {
Ok(client) => client,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "SSE client error: {_err:#}");
return;
}
};
let token = session.cancellation_token();
#[cfg(feature = "client-oauth")]
let mut reauthorized = false;
let mut streamed = false;
loop {
let bearer = auth.fresh_bearer().await;
let mut req = client
.get(session.url())
.header(ACCEPT, "application/json, text/event-stream")
.header(CACHE_CONTROL, "no-cache");
if let Some(ref bearer) = bearer {
req = req.bearer_auth(bearer);
}
if let Some(session_id) = session.session_id() {
req = req.header(MCP_SESSION_ID, session_id.to_string());
}
if let Some(last_id) = session.last_event_id() {
req = req.header(LAST_EVENT_ID, last_id);
}
let resp = match req.send().await {
Ok(resp) => resp,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to send SSE request: {}", _err);
session.cancellation_token().cancel();
return;
}
};
#[cfg(feature = "client-oauth")]
if (resp.status() == reqwest::StatusCode::UNAUTHORIZED
|| (resp.status() == reqwest::StatusCode::FORBIDDEN
&& insufficient_scope(resp.headers())))
&& !reauthorized
&& let ClientAuth::OAuth(oauth) = &auth
{
let challenge = bearer_challenge(resp.headers());
match oauth
.authorize(challenge.as_deref(), bearer.as_deref())
.await
{
Ok(_) => {
reauthorized = true;
continue;
}
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "OAuth authorization failed: {}", _err);
}
}
}
if resp.status() == reqwest::StatusCode::METHOD_NOT_ALLOWED
|| (resp.status() == reqwest::StatusCode::NOT_FOUND && !streamed)
{
#[cfg(feature = "tracing")]
tracing::debug!(
logger = "neva",
"server offers no standalone SSE stream ({}); continuing over POST only",
resp.status()
);
session.notify_sse_initialized();
return;
}
if !resp.status().is_success() {
#[cfg(feature = "tracing")]
tracing::error!(
logger = "neva",
"SSE request failed with status: {}",
resp.status()
);
session.cancellation_token().cancel();
return;
}
#[cfg(feature = "client-oauth")]
{
reauthorized = false;
}
let mut stream = sse_stream::SseStream::from_bytes_stream(resp.bytes_stream())
.fuse()
.map_ok(|event| handle_event(event, &session, &resp_tx))
.map_err(handle_error);
streamed = true;
session.notify_sse_initialized();
loop {
tokio::select! {
biased;
_ = token.cancelled() => return,
fut = stream.next() => {
let Some(Ok(fut)) = fut else {
#[cfg(feature = "tracing")]
tracing::info!(logger = "neva", "SSE stream ended, reconnecting");
break;
};
fut.await;
}
}
}
tokio::select! {
biased;
_ = token.cancelled() => return,
_ = tokio::time::sleep(session.retry_delay(SSE_RECONNECT_DELAY)) => {}
}
}
}
async fn drain_post_sse<S>(
mut stream: S,
resp_tx: &mpsc::Sender<Result<Message, Error>>,
ids: &[crate::types::RequestId],
) -> Drained
where
S: futures_util::Stream<Item = Result<sse_stream::Sse, sse_stream::Error>> + Unpin,
{
let mut owed = ids.to_vec();
let mut last_event_id = None;
let mut retry = None;
while !owed.is_empty()
&& let Some(event) = stream.next().await
{
match event {
Ok(sse) => {
if let Some(ms) = sse.retry {
retry = Some(ms);
}
if let Some(id) = sse.id.clone() {
last_event_id = Some(id);
}
if is_message_event(&sse) {
forward_sse_message(sse, resp_tx, &mut owed).await;
}
}
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "SSE POST stream error: {}", _err);
break;
}
}
}
Drained {
owed,
last_event_id,
retry,
}
}
#[derive(Debug)]
struct Drained {
owed: Vec<crate::types::RequestId>,
last_event_id: Option<String>,
retry: Option<u64>,
}
#[cfg(feature = "client-oauth")]
fn insufficient_scope(headers: &reqwest::header::HeaderMap) -> bool {
use volga_oauth_client::{BearerChallenge, OAuthErrorCode};
bearer_challenge(headers)
.and_then(|challenge| BearerChallenge::parse(&challenge).ok())
.is_some_and(|challenge| {
matches!(challenge.error(), Some(OAuthErrorCode::InsufficientScope))
})
}
#[cfg(feature = "client-oauth")]
fn bearer_challenge(headers: &reqwest::header::HeaderMap) -> Option<String> {
use volga_oauth_client::{BearerChallenge, OAuthErrorCode};
let mut first = None;
for value in headers
.get_all(reqwest::header::WWW_AUTHENTICATE)
.iter()
.filter_map(|value| value.to_str().ok())
{
for challenge in bearer_challenges(value) {
let Ok(parsed) = BearerChallenge::parse(&challenge) else {
continue;
};
if matches!(parsed.error(), Some(OAuthErrorCode::InsufficientScope)) {
return Some(challenge);
}
first.get_or_insert(challenge);
}
}
first
}
#[cfg(feature = "client-oauth")]
fn bearer_challenges(value: &str) -> Vec<String> {
let mut groups: Vec<Vec<&str>> = Vec::new();
for element in list_elements(value) {
let rest = element.trim_start();
let token_end = rest
.find(|c: char| c.is_whitespace() || c == '=')
.unwrap_or(rest.len());
let starts_challenge = !rest[token_end..].trim_start().starts_with('=');
if starts_challenge {
groups.push(vec![element]);
} else if let Some(group) = groups.last_mut() {
group.push(element);
}
}
groups
.into_iter()
.filter(|group| {
group[0]
.split_ascii_whitespace()
.next()
.is_some_and(|scheme| scheme.eq_ignore_ascii_case("Bearer"))
})
.map(|group| group.join(", "))
.collect()
}
#[cfg(feature = "client-oauth")]
fn list_elements(value: &str) -> Vec<&str> {
let mut elements = Vec::new();
let mut start = 0;
let mut quoted = false;
let mut escaped = false;
for (i, byte) in value.bytes().enumerate() {
if escaped {
escaped = false;
continue;
}
match byte {
b'\\' if quoted => escaped = true,
b'"' => quoted = !quoted,
b',' if !quoted => {
elements.push(value[start..i].trim());
start = i + 1;
}
_ => {}
}
}
elements.push(value[start..].trim());
elements.retain(|element| !element.is_empty());
elements
}
fn resumable(
#[cfg_attr(feature = "legacy-spec", allow(unused_variables))] session: &McpSession,
) -> bool {
#[cfg(not(feature = "legacy-spec"))]
{
session.is_legacy()
}
#[cfg(feature = "legacy-spec")]
{
true
}
}
async fn resume_stream(
client: &reqwest::Client,
session: &McpSession,
auth: &ClientAuth,
last_event_id: &str,
retry: Option<u64>,
resp_tx: &mpsc::Sender<Result<Message, Error>>,
ids: &[crate::types::RequestId],
) -> Vec<crate::types::RequestId> {
let delay = retry.map_or(SSE_RECONNECT_DELAY, std::time::Duration::from_millis);
let token = session.cancellation_token();
tokio::select! {
biased;
_ = token.cancelled() => return ids.to_vec(),
_ = tokio::time::sleep(delay) => {}
}
#[cfg_attr(not(feature = "client-oauth"), allow(unused_mut))]
let mut bearer = auth.fresh_bearer().await;
#[cfg(feature = "client-oauth")]
let mut reauthorized = false;
#[cfg_attr(not(feature = "client-oauth"), allow(clippy::never_loop))]
let resp = loop {
let mut req = client
.get(session.url())
.header(ACCEPT, "application/json, text/event-stream")
.header(CACHE_CONTROL, "no-cache")
.header(LAST_EVENT_ID, last_event_id);
if let Some(session_id) = session.session_id() {
req = req.header(MCP_SESSION_ID, session_id.to_string());
}
if let Some(bearer) = bearer.as_deref() {
req = req.bearer_auth(bearer);
}
let resp = match req.send().await {
Ok(resp) => resp,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to resume SSE stream: {}", _err);
return ids.to_vec();
}
};
if resp.status().is_success() {
break resp;
}
#[cfg(feature = "client-oauth")]
if !reauthorized
&& (resp.status() == reqwest::StatusCode::UNAUTHORIZED
|| (resp.status() == reqwest::StatusCode::FORBIDDEN
&& insufficient_scope(resp.headers())))
&& let ClientAuth::OAuth(oauth) = auth
{
let challenge = bearer_challenge(resp.headers());
if let Ok(fresh) = oauth
.authorize(challenge.as_deref(), bearer.as_deref())
.await
{
bearer = Some(fresh);
reauthorized = true;
continue;
}
}
#[cfg(feature = "tracing")]
tracing::debug!(
logger = "neva",
"SSE resumption refused with status: {}",
resp.status()
);
return ids.to_vec();
};
let stream = sse_stream::SseStream::from_bytes_stream(resp.bytes_stream());
tokio::select! {
biased;
_ = token.cancelled() => ids.to_vec(),
drained = drain_post_sse(stream, resp_tx, ids) => drained.owed,
}
}
fn is_message_event(event: &sse_stream::Sse) -> bool {
match &event.event {
None => true,
Some(kind) => kind.trim() == "message",
}
}
async fn handle_event(
event: sse_stream::Sse,
session: &Arc<McpSession>,
resp_tx: &mpsc::Sender<Result<Message, Error>>,
) {
if let Some(retry) = event.retry {
session.set_retry(retry);
}
let id = event.id.clone();
let delivered = if is_message_event(&event) {
handle_msg(event, resp_tx).await
} else {
#[cfg(feature = "tracing")]
tracing::debug!(logger = "neva", event = ?event);
true
};
if delivered && let Some(id) = id {
session.set_last_event_id(id);
}
}
#[inline]
fn handle_error(_err: sse_stream::Error) {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "SSE Error: {}", _err);
}
async fn handle_msg(
event: sse_stream::Sse,
resp_tx: &mpsc::Sender<Result<Message, Error>>,
) -> bool {
let Some(data) = event.data else {
return false;
};
let msg = match serde_json::from_str::<Message>(&data) {
Ok(msg) => msg,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to parse SSE event: {}", _err);
return false;
}
};
if let Err(_err) = resp_tx.send(Ok(msg)).await {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to send server request: {}", _err);
return false;
}
true
}
async fn forward_sse_message(
event: sse_stream::Sse,
resp_tx: &mpsc::Sender<Result<Message, Error>>,
owed: &mut Vec<crate::types::RequestId>,
) {
let Some(data) = event.data else {
return;
};
let msg = match serde_json::from_str::<Message>(&data) {
Ok(msg) => msg,
Err(_err) => {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to parse SSE POST event: {}", _err);
return;
}
};
let answered: Vec<_> = match &msg {
Message::Response(resp) => vec![resp.full_id()],
Message::Batch(batch) => batch
.iter()
.filter_map(|env| match env {
crate::types::MessageEnvelope::Response(resp) => Some(resp.full_id()),
_ => None,
})
.collect(),
_ => Vec::new(),
};
if let Err(_err) = resp_tx.send(Ok(msg)).await {
#[cfg(feature = "tracing")]
tracing::error!(logger = "neva", "Failed to send response: {}", _err);
return;
}
owed.retain(|id| !answered.contains(id));
}
#[inline]
#[cfg(not(feature = "client-tls"))]
fn create_client() -> Result<reqwest::Client, Error> {
reqwest::Client::builder().build().map_err(Error::from)
}
#[inline]
#[cfg(feature = "client-tls")]
fn create_client(mut tls_config: Option<ClientTlsConfig>) -> Result<reqwest::Client, Error> {
let mut builder = reqwest::ClientBuilder::new();
if let Some(ca_cert) = tls_config.as_mut().and_then(|tls| tls.ca.take()) {
builder = builder.add_root_certificate(ca_cert);
}
if let Some(identity) = tls_config.as_mut().and_then(|tls| tls.identity.take()) {
builder = builder.identity(identity);
}
if tls_config.is_some_and(|tls| !tls.certs_verification) {
builder = builder.danger_accept_invalid_certs(true);
}
builder.build().map_err(Error::from)
}
impl From<reqwest::Error> for Error {
#[inline]
fn from(err: reqwest::Error) -> Self {
Error::new(ErrorCode::ParseError, err.to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transport::http::ServiceUrl;
#[test]
#[cfg(feature = "client-oauth")]
fn only_the_challenge_error_parameter_says_the_scope_is_short() {
let headers_of = |value: &str| {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::WWW_AUTHENTICATE,
value.parse().expect("a header value"),
);
headers
};
let challenged = |value: &str| insufficient_scope(&headers_of(value));
assert!(challenged(r#"Bearer error="insufficient_scope""#));
assert!(challenged(
r#"Bearer realm="mcp", error="insufficient_scope", scope="admin""#
));
assert!(!challenged(
r#"Bearer error="invalid_token", error_description="missing insufficient_scope claim""#
));
assert!(!challenged(
r#"Bearer error="invalid_token", scope="insufficient_scope_admin""#
));
assert!(!challenged(r#"Bearer realm="mcp""#));
assert!(!challenged("Basic realm=\"mcp\""));
assert!(!insufficient_scope(&reqwest::header::HeaderMap::new()));
let mut headers = reqwest::header::HeaderMap::new();
headers.append(
reqwest::header::WWW_AUTHENTICATE,
r#"Basic realm="legacy""#.parse().expect("a header value"),
);
headers.append(
reqwest::header::WWW_AUTHENTICATE,
r#"Bearer error="insufficient_scope", scope="admin""#
.parse()
.expect("a header value"),
);
assert!(
insufficient_scope(&headers),
"the Bearer challenge counts wherever in the list it sits"
);
assert!(challenged(
r#"Basic realm="legacy", Bearer error="insufficient_scope""#
));
let mut headers = reqwest::header::HeaderMap::new();
headers.append(
reqwest::header::WWW_AUTHENTICATE,
r#"Bearer realm="legacy", error="invalid_token""#
.parse()
.expect("a header value"),
);
headers.append(
reqwest::header::WWW_AUTHENTICATE,
r#"Bearer realm="mcp", error="insufficient_scope", scope="admin""#
.parse()
.expect("a header value"),
);
assert!(
insufficient_scope(&headers),
"the challenge that names the code is the one that answers"
);
assert_eq!(
bearer_challenge(&headers).as_deref(),
Some(r#"Bearer realm="mcp", error="insufficient_scope", scope="admin""#),
"the flow must be given the challenge that says what is missing"
);
let mut plain = reqwest::header::HeaderMap::new();
plain.append(
reqwest::header::WWW_AUTHENTICATE,
r#"Basic realm="legacy""#.parse().expect("a header value"),
);
plain.append(
reqwest::header::WWW_AUTHENTICATE,
r#"Bearer resource_metadata="https://rs.example/.well-known/oauth-protected-resource""#
.parse()
.expect("a header value"),
);
assert_eq!(
bearer_challenge(&plain).as_deref(),
Some(
r#"Bearer resource_metadata="https://rs.example/.well-known/oauth-protected-resource""#
)
);
assert!(!insufficient_scope(&plain));
assert!(challenged(
r#"Bearer realm="legacy", error="invalid_token", Bearer realm="mcp", error="insufficient_scope", scope="admin""#
));
let mut combined = reqwest::header::HeaderMap::new();
combined.insert(
reqwest::header::WWW_AUTHENTICATE,
r#"Bearer realm="legacy", error="invalid_token", Bearer realm="mcp", error="insufficient_scope", scope="admin""#
.parse()
.expect("a header value"),
);
assert_eq!(
bearer_challenge(&combined).as_deref(),
Some(r#"Bearer realm="mcp", error="insufficient_scope", scope="admin""#),
"the applicable challenge is handed over on its own"
);
let spaced = r#"Bearer error = "insufficient_scope", scope = "admin""#;
assert!(challenged(spaced));
let parsed = volga_oauth_client::BearerChallenge::parse(
bearer_challenge(&headers_of(spaced))
.as_deref()
.expect("a challenge"),
)
.expect("it parses");
assert_eq!(parsed.scope(), Some("admin"));
}
#[test]
#[cfg(feature = "client-oauth")]
fn a_header_value_is_split_on_challenge_boundaries() {
assert_eq!(
bearer_challenges(r#"Bearer scope="a,b", error="insufficient_scope""#),
vec![r#"Bearer scope="a,b", error="insufficient_scope""#],
"a quoted comma is part of the value, not a list separator"
);
assert_eq!(
bearer_challenges(r#"Basic realm="legacy", Bearer realm="mcp""#),
vec![r#"Bearer realm="mcp""#],
"the other scheme's parameters stay with it"
);
assert_eq!(
bearer_challenges(r#"Bearer, Bearer error="insufficient_scope""#),
vec![r#"Bearer"#, r#"Bearer error="insufficient_scope""#],
"a bare challenge is still a challenge"
);
assert!(bearer_challenges(r#"Basic realm="legacy""#).is_empty());
assert_eq!(
bearer_challenges(r#"Bearer error="insufficient_scope", scope = "admin""#),
vec![r#"Bearer error="insufficient_scope", scope = "admin""#]
);
assert!(bearer_challenges("Basic dXNlcjpwYXNz==").is_empty());
assert_eq!(
bearer_challenges(r#"Bearer error_description="say \" then, stop""#),
vec![r#"Bearer error_description="say \" then, stop""#]
);
}
fn make_session() -> Arc<McpSession> {
Arc::new(McpSession::new(
ServiceUrl::default(),
CancellationToken::new(),
#[cfg(not(feature = "legacy-spec"))]
Default::default(),
))
}
#[cfg(not(feature = "legacy-spec"))]
#[test]
fn a_retried_post_mirrors_what_the_first_one_did() {
use crate::shared::param_headers::{ParamHeader, Registration};
let session = make_session();
let registry: crate::shared::param_headers::Registry = Default::default();
registry.insert(
"route".to_string(),
Registration::new(
vec![ParamHeader {
path: vec!["region".into()],
header: "Region".into(),
}],
0,
true,
),
);
let req = Message::Request(crate::types::Request::new(
Some(crate::types::RequestId::Number(1)),
crate::types::tool::commands::CALL,
Some(serde_json::json!({
"name": "route",
"arguments": { "region": "us-west1" }
})),
));
let mirrored = mirrored_param_headers(&session, &req, ®istry);
assert_eq!(
mirrored,
vec![("Mcp-Param-Region".to_string(), "us-west1".to_string())],
"the grace covers this call"
);
assert!(
mirrored_param_headers(&session, &req, ®istry).is_empty(),
"and reading is what spends it -- hence reading once"
);
let client = create_client(
#[cfg(feature = "client-tls")]
None,
)
.expect("a client");
for attempt in ["first", "retry"] {
let built = build_post(&client, &session, &req, None, &mirrored)
.build()
.expect("a request");
assert_eq!(
built
.headers()
.get("Mcp-Param-Region")
.and_then(|v| v.to_str().ok()),
Some("us-west1"),
"the {attempt} attempt must carry the mirrored header"
);
}
}
#[cfg(feature = "client-oauth")]
#[tokio::test]
async fn a_resumption_refused_for_its_token_authorizes_and_tries_again() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((mut stream, _)) = listener.accept().await {
let mut buf = [0u8; 8192];
let read = stream.read(&mut buf).await.unwrap_or(0);
let request = String::from_utf8_lossy(&buf[..read]).to_string();
let root = format!("http://{addr}");
let resp = if request.starts_with("GET /mcp") {
if request.contains("Bearer granted-token") {
let body = "id: 2\nevent: message\ndata: \
{\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\n\n";
format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
)
} else {
format!(
"HTTP/1.1 401 Unauthorized\r\nWWW-Authenticate: Bearer resource_metadata=\"{root}/.well-known/oauth-protected-resource/mcp\"\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
)
}
} else {
let body = if request.contains("/.well-known/oauth-protected-resource") {
format!(r#"{{"resource":"{root}/mcp","authorization_servers":["{root}"]}}"#)
} else if request.contains("/.well-known/") {
format!(
r#"{{"issuer":"{root}","token_endpoint":"{root}/token",
"authorization_endpoint":"{root}/authorize",
"registration_endpoint":"{root}/register",
"response_types_supported":["code"]}}"#
)
} else if request.contains("/register") {
r#"{"client_id":"registered-client"}"#.to_string()
} else {
r#"{"access_token":"granted-token","token_type":"Bearer","expires_in":3600}"#
.to_string()
};
format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
)
};
let _ = stream.write_all(resp.as_bytes()).await;
}
});
let url = format!("http://{addr}/mcp");
let session = Arc::new(McpSession::new(
ServiceUrl::from(addr.to_string().as_str()),
CancellationToken::new(),
#[cfg(not(feature = "legacy-spec"))]
Default::default(),
));
let config = oauth::OAuthClientConfig::default()
.require_https(false)
.with_handler(EchoesState);
let auth = ClientAuth::OAuth(Arc::new(
oauth::OAuthSession::new(config, &url).expect("a session"),
));
let (tx, mut rx) = mpsc::channel(2);
let owed = resume_stream(
&create_client(
#[cfg(feature = "client-tls")]
None,
)
.expect("a client"),
&session,
&auth,
"1",
Some(0),
&tx,
&[crate::types::RequestId::Number(1)],
)
.await;
assert!(
owed.is_empty(),
"the resumption must recover the answer rather than give up on a 401"
);
assert!(matches!(rx.try_recv(), Ok(Ok(Message::Response(_)))));
}
#[cfg(feature = "client-oauth")]
struct EchoesState;
#[cfg(feature = "client-oauth")]
impl crate::auth::oauth::AuthorizationHandler for EchoesState {
fn redirect_uri(
&self,
) -> futures_util::future::BoxFuture<'_, Result<String, crate::error::Error>> {
Box::pin(async { Ok("http://127.0.0.1:8919/callback".to_string()) })
}
fn authorize(
&self,
url: String,
) -> futures_util::future::BoxFuture<
'_,
Result<crate::auth::oauth::CallbackParams, crate::error::Error>,
> {
Box::pin(async move {
let state = url
.split(['?', '&'])
.find_map(|param| param.strip_prefix("state="))
.ok_or_else(|| {
Error::new(
ErrorCode::InvalidRequest,
"the authorization URL carried no `state`",
)
})?
.to_owned();
Ok(crate::auth::oauth::CallbackParams {
code: "the-code".into(),
state,
iss: None,
})
})
}
}
const VALID_MSG: &str = r#"{"jsonrpc":"2.0","method":"ping"}"#;
#[test]
fn parse_failure_classifies_statuses() {
let cases = [
(200, ErrorCode::ParseError),
(202, ErrorCode::ParseError),
(400, ErrorCode::ParseError),
(404, ErrorCode::ParseError),
(405, ErrorCode::ParseError),
(406, ErrorCode::ParseError),
(401, ErrorCode::InternalError),
(403, ErrorCode::InternalError),
(407, ErrorCode::InternalError),
(429, ErrorCode::InternalError),
(500, ErrorCode::InternalError),
(502, ErrorCode::InternalError),
(503, ErrorCode::InternalError),
(504, ErrorCode::InternalError),
];
for (status, expected) in cases {
let status = reqwest::StatusCode::from_u16(status).unwrap();
let (code, reason) = parse_failure(status, &"boom");
assert_eq!(code, expected, "wrong code for HTTP {status}");
assert!(
reason.contains(status.as_str()),
"the status must be carried in the message, got: {reason}"
);
}
}
#[tokio::test]
async fn it_advances_last_event_id_on_successful_delivery() {
let session = make_session();
let (tx, mut rx) = mpsc::channel(1);
let event = sse_stream::Sse::default().id("evt-1").data(VALID_MSG);
handle_event(event, &session, &tx).await;
assert_eq!(session.last_event_id(), Some("evt-1".to_string()));
assert!(rx.try_recv().is_ok(), "message should have been delivered");
}
#[tokio::test]
async fn it_does_not_advance_last_event_id_on_parse_failure() {
let session = make_session();
let (tx, _rx) = mpsc::channel(1);
let event = sse_stream::Sse::default()
.id("evt-bad")
.data("not { valid json");
handle_event(event, &session, &tx).await;
assert!(session.last_event_id().is_none());
}
#[tokio::test]
async fn it_does_not_advance_last_event_id_when_channel_closed() {
let session = make_session();
let (tx, rx) = mpsc::channel(1);
drop(rx);
let event = sse_stream::Sse::default().id("evt-dropped").data(VALID_MSG);
handle_event(event, &session, &tx).await;
assert!(session.last_event_id().is_none());
}
#[tokio::test]
async fn it_advances_last_event_id_for_non_message_event() {
let session = make_session();
let (tx, _rx) = mpsc::channel(1);
let event = sse_stream::Sse::default()
.id("evt-keepalive")
.event("keepalive");
handle_event(event, &session, &tx).await;
assert_eq!(session.last_event_id(), Some("evt-keepalive".to_string()));
}
#[tokio::test]
async fn it_does_not_advance_last_event_id_when_data_is_absent() {
let session = make_session();
let (tx, _rx) = mpsc::channel(1);
let event = sse_stream::Sse::default().id("evt-empty");
handle_event(event, &session, &tx).await;
assert!(session.last_event_id().is_none());
}
#[test]
fn explicitly_named_message_events_count_as_messages() {
let cases = [
(None, true),
(Some("message"), true),
(Some(" message "), true),
(Some("Message"), false),
(Some("keepalive"), false),
(Some("endpoint"), false),
];
for (kind, expected) in cases {
let event = match kind {
Some(kind) => sse_stream::Sse::default().event(kind),
None => sse_stream::Sse::default(),
};
assert_eq!(
is_message_event(&event),
expected,
"wrong verdict for event type {kind:?}"
);
}
}
#[tokio::test]
async fn named_message_events_are_delivered_on_both_sse_paths() {
let response = r#"{"jsonrpc":"2.0","id":1,"result":{}}"#;
let session = make_session();
let (tx, mut rx) = mpsc::channel(1);
let event = sse_stream::Sse::default()
.id("evt-named")
.event("message")
.data(response);
handle_event(event, &session, &tx).await;
assert!(rx.try_recv().is_ok(), "GET frame must be delivered");
assert_eq!(session.last_event_id(), Some("evt-named".to_string()));
let (tx, mut rx) = mpsc::channel(4);
let frames = vec![
Ok(sse_stream::Sse::default()
.event("message")
.data(r#"{"jsonrpc":"2.0","method":"notifications/message"}"#)),
Ok(sse_stream::Sse::default().event("message").data(response)),
];
assert!(
drain_post_sse(
futures_util::stream::iter(frames),
&tx,
&[crate::types::RequestId::Number(1)],
)
.await
.owed
.is_empty(),
"the POST stream must leave nothing owed"
);
assert!(matches!(rx.try_recv(), Ok(Ok(Message::Notification(_)))));
assert!(matches!(rx.try_recv(), Ok(Ok(Message::Response(_)))));
}
#[tokio::test]
async fn a_priming_frame_still_states_where_to_resume_from() {
let session = make_session();
session.set_last_event_id("get-stream-7".to_string());
session.set_retry(9_000);
let (tx, mut rx) = mpsc::channel(2);
let mut priming = sse_stream::Sse::default().id("event-1");
priming.retry = Some(500);
let frames = vec![Ok(priming)];
let drained = drain_post_sse(
futures_util::stream::iter(frames),
&tx,
&[crate::types::RequestId::Number(1)],
)
.await;
assert_eq!(
drained.owed,
vec![crate::types::RequestId::Number(1)],
"a priming frame answers nothing"
);
assert!(rx.try_recv().is_err(), "and delivers nothing");
assert_eq!(
drained.last_event_id,
Some("event-1".to_string()),
"this stream resumes from where this stream got to"
);
assert_eq!(
drained.retry,
Some(500),
"and after the delay this stream was given"
);
assert_eq!(
session.last_event_id(),
Some("get-stream-7".to_string()),
"leaving the standalone GET's own position alone"
);
assert_eq!(
session.retry_delay(SSE_RECONNECT_DELAY),
Duration::from_millis(9_000),
"and its own reconnection delay with it"
);
}
#[tokio::test]
async fn drain_post_sse_skips_other_event_types() {
let (tx, mut rx) = mpsc::channel(2);
let frames = vec![
Ok(sse_stream::Sse::default().event("keepalive").data("{}")),
Ok(sse_stream::Sse::default()
.event("endpoint")
.data(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#)),
];
assert_eq!(
drain_post_sse(
futures_util::stream::iter(frames),
&tx,
&[crate::types::RequestId::Number(1)],
)
.await
.owed,
vec![crate::types::RequestId::Number(1)]
);
assert!(rx.try_recv().is_err(), "no frame should be delivered");
}
#[tokio::test]
async fn forward_sse_message_flags_only_terminal_replies() {
let cases = [
(
r#"{"jsonrpc":"2.0","method":"notifications/message"}"#,
false,
),
(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#, true),
(r#"[{"jsonrpc":"2.0","id":1,"result":{}}]"#, true),
(
r#"[{"jsonrpc":"2.0","method":"notifications/subscriptions/acknowledged"},
{"jsonrpc":"2.0","method":"notifications/tools/list_changed"}]"#,
false,
),
(
r#"[{"jsonrpc":"2.0","method":"notifications/message"},
{"jsonrpc":"2.0","id":1,"result":{}}]"#,
true,
),
(r#"{"jsonrpc":"2.0","id":9,"result":{}}"#, false),
(r#"[{"jsonrpc":"2.0","id":9,"result":{}}]"#, false),
];
for (frame, terminal) in cases {
let (tx, mut rx) = mpsc::channel(1);
let event = sse_stream::Sse::default().data(frame);
let mut owed = vec![crate::types::RequestId::Number(1)];
forward_sse_message(event, &tx, &mut owed).await;
assert_eq!(owed.is_empty(), terminal, "wrong terminal flag for {frame}");
assert!(rx.try_recv().is_ok(), "{frame} should still be delivered");
}
}
#[tokio::test]
async fn a_batch_is_struck_off_one_answer_at_a_time() {
let ids = [
crate::types::RequestId::Number(1),
crate::types::RequestId::Number(2),
];
let (tx, _rx) = mpsc::channel(2);
let mut owed = ids.to_vec();
forward_sse_message(
sse_stream::Sse::default().data(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#),
&tx,
&mut owed,
)
.await;
assert_eq!(
owed,
vec![crate::types::RequestId::Number(2)],
"only the answered request may be struck off"
);
forward_sse_message(
sse_stream::Sse::default().data(r#"{"jsonrpc":"2.0","id":2,"result":{}}"#),
&tx,
&mut owed,
)
.await;
assert!(owed.is_empty(), "the batch is now fully answered");
}
#[tokio::test]
async fn draining_stops_once_nothing_is_owed() {
let (tx, mut rx) = mpsc::channel(8);
let frames = futures_util::stream::iter(vec![
Ok(sse_stream::Sse::default().data(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#)),
Ok(sse_stream::Sse::default()
.data(r#"{"jsonrpc":"2.0","method":"notifications/message"}"#)),
])
.chain(futures_util::stream::pending());
let owed = tokio::time::timeout(
Duration::from_secs(1),
drain_post_sse(Box::pin(frames), &tx, &[crate::types::RequestId::Number(1)]),
)
.await
.expect("the drain must return instead of holding the session stream open");
assert!(owed.owed.is_empty());
assert!(matches!(rx.try_recv(), Ok(Ok(Message::Response(_)))));
assert!(
rx.try_recv().is_err(),
"nothing past the answer belongs to this exchange"
);
}
#[tokio::test]
async fn forward_sse_message_reports_unparseable_frame_as_unanswered() {
let (tx, mut rx) = mpsc::channel(1);
let event = sse_stream::Sse::default().data("not json");
let mut owed = vec![crate::types::RequestId::Number(1)];
forward_sse_message(event, &tx, &mut owed).await;
assert_eq!(owed, vec![crate::types::RequestId::Number(1)]);
assert!(
rx.try_recv().is_err(),
"a malformed frame must not reach the receive loop"
);
}
#[test]
fn event_stream_media_type_is_matched_case_insensitively() {
let cases = [
("text/event-stream", true),
("Text/Event-Stream", true),
("TEXT/EVENT-STREAM; charset=utf-8", true),
("text/event-stream ;charset=utf-8", true),
("application/json", false),
("text/event-streaming", false),
("application/json, text/event-stream", false),
];
for (value, expected) in cases {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(CONTENT_TYPE, value.parse().unwrap());
assert_eq!(
is_event_stream(&headers),
expected,
"wrong verdict for content-type {value:?}"
);
}
assert!(
!is_event_stream(&reqwest::header::HeaderMap::new()),
"a reply without content-type is not an SSE stream"
);
}
}
#[cfg(test)]
#[cfg(not(feature = "legacy-spec"))]
mod routing_hints_tests {
use super::{name_param, routing_hints};
use crate::transport::http::encode_header_value;
use crate::types::notification::Notification;
use crate::types::{Message, Request, RequestId};
use serde_json::json;
#[test]
fn request_yields_method_and_no_name() {
let req = Request::new::<()>(Some(RequestId::Number(1)), "tools/list", None);
let msg = Message::Request(req);
let hints = routing_hints(&msg).unwrap();
assert_eq!(hints.0, "tools/list");
assert!(hints.1.is_none());
}
#[test]
fn tools_call_yields_method_and_tool_name() {
let req = Request::new(
Some(RequestId::Number(1)),
"tools/call",
Some(json!({"name": "echo", "arguments": {}})),
);
let msg = Message::Request(req);
let hints = routing_hints(&msg).unwrap();
assert_eq!(hints.0, "tools/call");
assert_eq!(hints.1.as_deref(), Some("echo"));
}
#[test]
#[cfg(not(feature = "legacy-spec"))]
fn name_header_is_sourced_per_method() {
use crate::types::{Request, RequestId};
use serde_json::json;
let cases = [
("tools/call", json!({ "name": "echo" }), Some("echo")),
(
"prompts/get",
json!({ "name": "greeting" }),
Some("greeting"),
),
(
"resources/read",
json!({ "uri": "file:///a.txt" }),
Some("file:///a.txt"),
),
("tools/list", json!({}), None),
];
for (method, params, expected) in cases {
let req = Request::new(Some(RequestId::Number(1)), method, Some(params));
assert_eq!(name_param(&req).as_deref(), expected, "method: {method}");
}
}
#[test]
#[cfg(not(feature = "legacy-spec"))]
fn header_values_are_encoded_when_not_ascii_safe() {
assert_eq!(encode_header_value("us-west1"), "us-west1");
assert_eq!(encode_header_value("caf\u{e9}"), "=?base64?Y2Fmw6k=?=");
assert_eq!(encode_header_value(" lead"), "=?base64?IGxlYWQ=?=");
assert_eq!(encode_header_value("trail "), "=?base64?dHJhaWwg?=");
assert_eq!(encode_header_value("a\nb"), "=?base64?YQpi?=");
assert_eq!(encode_header_value("a\tb"), "a\tb");
assert_eq!(encode_header_value("\tindented"), "=?base64?CWluZGVudGVk?=");
assert_eq!(encode_header_value("trailing\t"), "=?base64?dHJhaWxpbmcJ?=");
assert_eq!(
encode_header_value("=?base64?zzz?="),
"=?base64?PT9iYXNlNjQ/enp6Pz0=?="
);
}
#[test]
fn notification_yields_method_only() {
let n = Notification::new("notifications/cancelled", None);
let msg = Message::Notification(n);
let hints = routing_hints(&msg).unwrap();
assert_eq!(hints.0, "notifications/cancelled");
assert!(hints.1.is_none());
}
}