use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use base64::Engine;
use tokio::sync::{Notify, RwLock, mpsc};
use tokio::task::JoinHandle;
use super::transport::ClientTransport;
use crate::error::{Error, Result};
use crate::protocol::{RequestId, notifications};
const MCP_METHOD_HEADER: &str = "mcp-method";
const MCP_NAME_HEADER: &str = "mcp-name";
const MCP_PARAM_HEADER_PREFIX: &str = "mcp-param-";
const BASE64_SENTINEL_PREFIX: &str = "=?base64?";
const BASE64_SENTINEL_SUFFIX: &str = "?=";
#[derive(Debug, Clone)]
struct CustomHeaderMapping {
suffix: String,
property_path: Vec<String>,
}
#[cfg(feature = "oauth-client")]
#[derive(Clone)]
struct ScopeEscalationRuntime {
handler: Arc<dyn OAuthScopeEscalationHandler>,
state: Arc<tokio::sync::Mutex<ScopeEscalationState>>,
max_attempts: usize,
}
#[cfg(feature = "oauth-client")]
struct ScopeEscalationState {
scopes: Vec<String>,
revision: usize,
}
#[cfg(feature = "oauth-client")]
impl ScopeEscalationRuntime {
fn new<P>(handler: Arc<P>, config: OAuthScopeEscalationConfig) -> Self
where
P: OAuthScopeEscalationHandler,
{
Self {
handler,
state: Arc::new(tokio::sync::Mutex::new(ScopeEscalationState {
scopes: config.initial_scopes().to_vec(),
revision: 0,
})),
max_attempts: config.maximum_attempts(),
}
}
async fn respond_to_challenge(
&self,
challenge: OAuthScopeChallenge,
resource: &str,
operation: &str,
attempt: usize,
observed_revision: usize,
) -> std::result::Result<ScopeEscalationDecision, OAuthClientError> {
let mut state = self.state.lock().await;
let previous_scopes = state.scopes.clone();
let mut requested_scopes = previous_scopes.clone();
for scope in &challenge.required_scopes {
if !requested_scopes.contains(scope) {
requested_scopes.push(scope.clone());
}
}
if requested_scopes == previous_scopes && state.revision > observed_revision {
return Ok(ScopeEscalationDecision {
revision: state.revision,
});
}
self.handler
.reauthorize(OAuthScopeEscalationRequest {
resource: resource.to_string(),
operation: operation.to_string(),
challenge,
previous_scopes,
requested_scopes: requested_scopes.clone(),
attempt,
})
.await?;
state.scopes = requested_scopes;
state.revision += 1;
Ok(ScopeEscalationDecision {
revision: state.revision,
})
}
}
#[cfg(feature = "oauth-client")]
struct ScopeEscalationDecision {
revision: usize,
}
#[cfg(feature = "oauth-client")]
use super::oauth::{
OAuthClientError, OAuthScopeChallenge, OAuthScopeEscalationConfig, OAuthScopeEscalationHandler,
OAuthScopeEscalationRequest, TokenProvider,
};
#[derive(Debug, Clone)]
pub struct HttpClientConfig {
pub headers: HashMap<String, String>,
pub auto_sse: bool,
pub channel_capacity: usize,
pub request_timeout: Duration,
pub notification_timeout: Duration,
pub sse_reconnect: bool,
pub sse_reconnect_delay: Duration,
pub max_sse_reconnect_attempts: u32,
pub session_recovery: bool,
pub max_sse_event_size: usize,
}
pub const DEFAULT_MAX_SSE_EVENT_SIZE: usize = 16 * 1024 * 1024;
impl Default for HttpClientConfig {
fn default() -> Self {
Self {
headers: HashMap::new(),
auto_sse: true,
channel_capacity: 256,
request_timeout: Duration::from_secs(30),
notification_timeout: Duration::from_secs(5),
sse_reconnect: true,
sse_reconnect_delay: Duration::from_secs(1),
max_sse_reconnect_attempts: 5,
session_recovery: true,
max_sse_event_size: DEFAULT_MAX_SSE_EVENT_SIZE,
}
}
}
impl HttpClientConfig {
pub fn bearer_token(mut self, token: impl Into<String>) -> Self {
self.headers.insert(
"Authorization".to_string(),
format!("Bearer {}", token.into()),
);
self
}
pub fn api_key_header(mut self, name: impl Into<String>, key: impl Into<String>) -> Self {
self.headers.insert(name.into(), key.into());
self
}
pub fn basic_auth(mut self, username: impl AsRef<str>, password: impl AsRef<str>) -> Self {
use base64::Engine;
let encoded = base64::engine::general_purpose::STANDARD.encode(format!(
"{}:{}",
username.as_ref(),
password.as_ref()
));
self.headers
.insert("Authorization".to_string(), format!("Basic {}", encoded));
self
}
pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.headers.insert(name.into(), value.into());
self
}
}
pub struct HttpClientTransport {
url: String,
client: reqwest::Client,
session_id: Option<String>,
protocol_version: Option<String>,
tool_header_mappings: HashMap<String, Vec<CustomHeaderMapping>>,
incoming_rx: mpsc::Receiver<String>,
incoming_tx: mpsc::Sender<String>,
sse_task: Option<JoinHandle<()>>,
request_tasks: HashMap<RequestId, JoinHandle<()>>,
last_event_id: Arc<RwLock<Option<String>>>,
sse_retry_delay: Arc<RwLock<Option<Duration>>>,
sse_reconnect_signal: Arc<Notify>,
connected: Arc<AtomicBool>,
config: HttpClientConfig,
#[cfg(feature = "oauth-client")]
token_provider: Option<Arc<dyn TokenProvider>>,
#[cfg(feature = "oauth-client")]
scope_escalation: Option<ScopeEscalationRuntime>,
}
impl HttpClientTransport {
pub fn new(url: impl Into<String>) -> Self {
Self::with_config(url, HttpClientConfig::default())
}
pub fn with_config(url: impl Into<String>, config: HttpClientConfig) -> Self {
let (tx, rx) = mpsc::channel(config.channel_capacity);
Self {
url: url.into(),
client: reqwest::Client::new(),
session_id: None,
protocol_version: None,
tool_header_mappings: HashMap::new(),
incoming_rx: rx,
incoming_tx: tx,
sse_task: None,
request_tasks: HashMap::new(),
last_event_id: Arc::new(RwLock::new(None)),
sse_retry_delay: Arc::new(RwLock::new(None)),
sse_reconnect_signal: Arc::new(Notify::new()),
connected: Arc::new(AtomicBool::new(true)),
config,
#[cfg(feature = "oauth-client")]
token_provider: None,
#[cfg(feature = "oauth-client")]
scope_escalation: None,
}
}
pub fn with_client(url: impl Into<String>, client: reqwest::Client) -> Self {
let config = HttpClientConfig::default();
let (tx, rx) = mpsc::channel(config.channel_capacity);
Self {
url: url.into(),
client,
session_id: None,
protocol_version: None,
tool_header_mappings: HashMap::new(),
incoming_rx: rx,
incoming_tx: tx,
sse_task: None,
request_tasks: HashMap::new(),
last_event_id: Arc::new(RwLock::new(None)),
sse_retry_delay: Arc::new(RwLock::new(None)),
sse_reconnect_signal: Arc::new(Notify::new()),
connected: Arc::new(AtomicBool::new(true)),
config,
#[cfg(feature = "oauth-client")]
token_provider: None,
#[cfg(feature = "oauth-client")]
scope_escalation: None,
}
}
pub fn bearer_token(mut self, token: impl Into<String>) -> Self {
self.config.headers.insert(
"Authorization".to_string(),
format!("Bearer {}", token.into()),
);
self
}
pub fn api_key(self, key: impl Into<String>) -> Self {
self.bearer_token(key)
}
pub fn api_key_header(mut self, name: impl Into<String>, key: impl Into<String>) -> Self {
self.config.headers.insert(name.into(), key.into());
self
}
pub fn basic_auth(mut self, username: impl AsRef<str>, password: impl AsRef<str>) -> Self {
use base64::Engine;
let encoded = base64::engine::general_purpose::STANDARD.encode(format!(
"{}:{}",
username.as_ref(),
password.as_ref()
));
self.config
.headers
.insert("Authorization".to_string(), format!("Basic {}", encoded));
self
}
pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.config.headers.insert(name.into(), value.into());
self
}
pub fn disable_session_recovery(mut self) -> Self {
self.config.session_recovery = false;
self
}
#[cfg(feature = "oauth-client")]
pub fn with_token_provider(mut self, provider: impl TokenProvider) -> Self {
self.token_provider = Some(Arc::new(provider));
self.scope_escalation = None;
self
}
#[cfg(feature = "oauth-client")]
pub fn with_scope_aware_token_provider<P>(
mut self,
provider: P,
config: OAuthScopeEscalationConfig,
) -> Self
where
P: TokenProvider + OAuthScopeEscalationHandler,
{
let provider = Arc::new(provider);
self.token_provider = Some(provider.clone());
self.scope_escalation = Some(ScopeEscalationRuntime::new(provider, config));
self
}
fn outgoing_custom_headers(&self, parsed: &serde_json::Value) -> Vec<(String, String)> {
if parsed.get("method").and_then(serde_json::Value::as_str) != Some("tools/call") {
return Vec::new();
}
let Some(params) = parsed.get("params") else {
return Vec::new();
};
let Some(name) = params.get("name").and_then(serde_json::Value::as_str) else {
return Vec::new();
};
let Some(mappings) = self.tool_header_mappings.get(name) else {
return Vec::new();
};
let arguments = params.get("arguments").unwrap_or(&serde_json::Value::Null);
mappings
.iter()
.filter_map(|mapping| {
let value = value_at_property_path(arguments, &mapping.property_path)?;
if value.is_null() {
return None;
}
let rendered = json_value_to_header_string(value)?;
Some((
format!("{MCP_PARAM_HEADER_PREFIX}{}", mapping.suffix),
encode_header_value(&rendered),
))
})
.collect()
}
fn normalize_incoming_message(&mut self, message: String) -> String {
if self.protocol_version.as_deref() != Some(crate::protocol::PROTOCOL_VERSION_2026_07_28) {
return message;
}
let Ok(mut parsed) = serde_json::from_str::<serde_json::Value>(&message) else {
return message;
};
let Some(tools) = parsed
.get_mut("result")
.and_then(|result| result.get_mut("tools"))
.and_then(serde_json::Value::as_array_mut)
else {
return message;
};
self.tool_header_mappings.clear();
tools.retain(|tool| {
let Some(name) = tool.get("name").and_then(serde_json::Value::as_str) else {
return false;
};
let Some(schema) = tool.get("inputSchema") else {
return false;
};
match custom_header_mappings(schema) {
Ok(mappings) => {
self.tool_header_mappings.insert(name.to_string(), mappings);
true
}
Err(error) => {
tracing::warn!(tool = %name, %error, "Excluding tool with invalid x-mcp-header annotations");
false
}
}
});
parsed.to_string()
}
fn start_sse_stream(&mut self) {
let url = self.url.clone();
let client = self.client.clone();
let session_id = self.session_id.clone().unwrap();
let protocol_version = self.protocol_version.clone();
let tx = self.incoming_tx.clone();
let last_event_id = self.last_event_id.clone();
let sse_retry_delay = self.sse_retry_delay.clone();
let reconnect_signal = self.sse_reconnect_signal.clone();
let connected = self.connected.clone();
let config = self.config.clone();
#[cfg(feature = "oauth-client")]
let token_provider = self.token_provider.clone();
self.sse_task = Some(tokio::spawn(async move {
sse_stream_loop(SseLoopParams {
url,
client,
session_id,
protocol_version,
tx,
last_event_id,
sse_retry_delay,
reconnect_signal,
connected,
config,
#[cfg(feature = "oauth-client")]
token_provider,
})
.await;
}));
}
}
fn custom_header_mappings(
schema: &serde_json::Value,
) -> std::result::Result<Vec<CustomHeaderMapping>, String> {
fn annotation_count(value: &serde_json::Value) -> usize {
match value {
serde_json::Value::Object(object) => {
usize::from(object.contains_key("x-mcp-header"))
+ object.values().map(annotation_count).sum::<usize>()
}
serde_json::Value::Array(values) => values.iter().map(annotation_count).sum::<usize>(),
_ => 0,
}
}
fn is_tchar(byte: u8) -> bool {
byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'!' | b'#'
| b'$'
| b'%'
| b'&'
| b'\''
| b'*'
| b'+'
| b'-'
| b'.'
| b'^'
| b'_'
| b'`'
| b'|'
| b'~'
)
}
fn primitive_header_type(schema: &serde_json::Value) -> bool {
match schema.get("type") {
Some(serde_json::Value::String(kind)) => {
matches!(kind.as_str(), "string" | "number" | "integer" | "boolean")
}
Some(serde_json::Value::Array(kinds)) => {
let mut primitive = false;
for kind in kinds {
match kind.as_str() {
Some("string" | "number" | "integer" | "boolean") if !primitive => {
primitive = true;
}
Some("null") => {}
_ => return false,
}
}
primitive
}
_ => false,
}
}
fn walk(
schema: &serde_json::Value,
path: &mut Vec<String>,
seen: &mut std::collections::HashSet<String>,
mappings: &mut Vec<CustomHeaderMapping>,
) -> std::result::Result<(), String> {
let Some(properties) = schema
.get("properties")
.and_then(serde_json::Value::as_object)
else {
return Ok(());
};
for (property_name, property_schema) in properties {
path.push(property_name.clone());
if let Some(annotation) = property_schema.get("x-mcp-header") {
let suffix = annotation
.as_str()
.ok_or_else(|| format!("annotation at {} is not a string", path.join(".")))?;
if suffix.is_empty() || !suffix.bytes().all(is_tchar) {
return Err(format!(
"invalid header suffix {suffix:?} at {}",
path.join(".")
));
}
if !primitive_header_type(property_schema) {
return Err(format!(
"annotation at {} is not on a primitive property",
path.join(".")
));
}
if !seen.insert(suffix.to_ascii_lowercase()) {
return Err(format!("duplicate header suffix {suffix:?}"));
}
mappings.push(CustomHeaderMapping {
suffix: suffix.to_string(),
property_path: path.clone(),
});
}
walk(property_schema, path, seen, mappings)?;
path.pop();
}
Ok(())
}
let mut mappings = Vec::new();
walk(
schema,
&mut Vec::new(),
&mut std::collections::HashSet::new(),
&mut mappings,
)?;
if annotation_count(schema) != mappings.len() {
return Err(
"x-mcp-header annotation is not statically reachable through properties".to_string(),
);
}
Ok(mappings)
}
fn value_at_property_path<'a>(
root: &'a serde_json::Value,
path: &[String],
) -> Option<&'a serde_json::Value> {
path.iter().try_fold(root, |value, key| value.get(key))
}
fn json_value_to_header_string(value: &serde_json::Value) -> Option<String> {
match value {
serde_json::Value::String(value) => Some(value.clone()),
serde_json::Value::Number(value) => Some(value.to_string()),
serde_json::Value::Bool(value) => Some(value.to_string()),
serde_json::Value::Null | serde_json::Value::Array(_) | serde_json::Value::Object(_) => {
None
}
}
}
fn encode_header_value(value: &str) -> String {
let unsafe_for_header =
value.trim() != value || value.bytes().any(|byte| !(0x20..=0x7e).contains(&byte));
if unsafe_for_header {
format!(
"{BASE64_SENTINEL_PREFIX}{}{BASE64_SENTINEL_SUFFIX}",
base64::engine::general_purpose::STANDARD.encode(value)
)
} else {
value.to_string()
}
}
#[cfg(feature = "oauth-client")]
fn bearer_headers(token: &str) -> std::result::Result<reqwest::header::HeaderMap, String> {
let value = reqwest::header::HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|_| "token provider returned an invalid bearer token".to_string())?;
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(reqwest::header::AUTHORIZATION, value);
Ok(headers)
}
fn is_jsonrpc_error_response(value: &serde_json::Value) -> bool {
value.get("error").is_some_and(serde_json::Value::is_object)
&& value.pointer("/error/code").is_some()
&& value.pointer("/error/message").is_some()
}
struct HttpRequestSendError {
message: String,
connection_failed: bool,
}
impl HttpRequestSendError {
fn request(error: reqwest::Error) -> Self {
Self {
message: format!("HTTP request failed: {error}"),
connection_failed: true,
}
}
#[cfg(feature = "oauth-client")]
fn oauth(error: OAuthClientError) -> Self {
Self {
message: error.to_string(),
connection_failed: false,
}
}
}
async fn send_http_request(
request: reqwest::RequestBuilder,
resource: &str,
operation: &str,
#[cfg(feature = "oauth-client")] token_provider: Option<Arc<dyn TokenProvider>>,
#[cfg(feature = "oauth-client")] scope_escalation: Option<ScopeEscalationRuntime>,
#[cfg(feature = "oauth-client")] initial_scope_revision: usize,
) -> std::result::Result<reqwest::Response, HttpRequestSendError> {
#[cfg(feature = "oauth-client")]
let mut request = request;
#[cfg(not(feature = "oauth-client"))]
let _ = (resource, operation);
#[cfg(feature = "oauth-client")]
let mut observed_revision = initial_scope_revision;
#[cfg(feature = "oauth-client")]
let mut attempts = 0;
loop {
#[cfg(feature = "oauth-client")]
let retry_request = request.try_clone();
let response = request
.send()
.await
.map_err(HttpRequestSendError::request)?;
#[cfg(feature = "oauth-client")]
{
let challenge = if response.status() == reqwest::StatusCode::FORBIDDEN {
scope_challenge(response.headers())
} else {
None
};
let Some(challenge) = challenge else {
return Ok(response);
};
let (Some(runtime), Some(provider), Some(mut retry_request)) = (
scope_escalation.as_ref(),
token_provider.as_ref(),
retry_request,
) else {
return Ok(response);
};
if attempts >= runtime.max_attempts {
return Ok(response);
}
attempts += 1;
let decision = runtime
.respond_to_challenge(challenge, resource, operation, attempts, observed_revision)
.await
.map_err(HttpRequestSendError::oauth)?;
observed_revision = decision.revision;
let token = provider
.get_token()
.await
.map_err(HttpRequestSendError::oauth)?;
let headers = bearer_headers(&token).map_err(|message| {
HttpRequestSendError::oauth(OAuthClientError::ScopeEscalation(message))
})?;
retry_request = retry_request.headers(headers);
request = retry_request;
}
#[cfg(not(feature = "oauth-client"))]
return Ok(response);
}
}
#[cfg(feature = "oauth-client")]
fn scope_challenge(headers: &reqwest::header::HeaderMap) -> Option<OAuthScopeChallenge> {
headers
.get_all(reqwest::header::WWW_AUTHENTICATE)
.iter()
.filter_map(|value| value.to_str().ok())
.find_map(OAuthScopeChallenge::from_www_authenticate)
}
fn http_status_error(status: reqwest::StatusCode, headers: &reqwest::header::HeaderMap) -> String {
#[cfg(feature = "oauth-client")]
if let Some(challenge) = scope_challenge(headers) {
let mut message = format!(
"server returned HTTP {status}: insufficient_scope requires {}",
challenge.required_scopes.join(" ")
);
if let Some(resource_metadata) = challenge.resource_metadata {
message.push_str(&format!(" (resource metadata: {resource_metadata})"));
}
return message;
}
#[cfg(not(feature = "oauth-client"))]
let _ = headers;
format!("server returned HTTP {status}")
}
fn operation_label(parsed: Option<&serde_json::Value>) -> String {
let Some(method) = parsed
.and_then(|value| value.get("method"))
.and_then(serde_json::Value::as_str)
else {
return "unknown".to_string();
};
let target = match method {
"tools/call" | "prompts/get" => parsed
.and_then(|value| value.pointer("/params/name"))
.and_then(serde_json::Value::as_str),
"resources/read" => parsed
.and_then(|value| value.pointer("/params/uri"))
.and_then(serde_json::Value::as_str),
"tasks/get" | "tasks/update" | "tasks/cancel" => parsed
.and_then(|value| value.pointer("/params/taskId"))
.and_then(serde_json::Value::as_str),
_ => None,
};
match target {
Some(target) => format!("{method}:{target}"),
None => method.to_string(),
}
}
#[async_trait]
impl ClientTransport for HttpClientTransport {
async fn send(&mut self, message: &str) -> Result<()> {
if !self.connected.load(Ordering::Acquire) {
return Err(Error::Transport("Transport closed".to_string()));
}
let parsed_message = serde_json::from_str::<serde_json::Value>(message).ok();
let is_notification = parsed_message
.as_ref()
.map(|v| v.get("id").is_none())
.unwrap_or(false);
let method = parsed_message
.as_ref()
.and_then(|value| value.get("method"))
.and_then(serde_json::Value::as_str);
let operation = operation_label(parsed_message.as_ref());
let outbound_version = parsed_message
.as_ref()
.and_then(|value| {
value.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion")
})
.and_then(serde_json::Value::as_str)
.map(str::to_string);
let is_modern_request =
outbound_version.as_deref() == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28);
if is_modern_request {
self.protocol_version = outbound_version.clone();
self.session_id = None;
}
let timeout = if is_notification {
self.config
.notification_timeout
.min(self.config.request_timeout)
} else {
self.config.request_timeout
};
let mut request = self
.client
.post(&self.url)
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream");
if method != Some("subscriptions/listen") {
request = request.timeout(timeout);
}
if !is_modern_request && let Some(ref session_id) = self.session_id {
request = request.header("mcp-session-id", session_id);
}
if let Some(version) = outbound_version.as_ref().or(self.protocol_version.as_ref()) {
request = request.header("mcp-protocol-version", version);
}
if let Some(method) = method {
request = request.header(MCP_METHOD_HEADER, method);
let name = match method {
"tools/call" | "prompts/get" => parsed_message
.as_ref()
.and_then(|value| value.pointer("/params/name"))
.and_then(serde_json::Value::as_str),
"resources/read" => parsed_message
.as_ref()
.and_then(|value| value.pointer("/params/uri"))
.and_then(serde_json::Value::as_str),
"tasks/get" | "tasks/update" | "tasks/cancel" => parsed_message
.as_ref()
.and_then(|value| value.pointer("/params/taskId"))
.and_then(serde_json::Value::as_str),
_ => None,
};
if let Some(name) = name {
request = request.header(MCP_NAME_HEADER, name);
}
}
if let Some(parsed) = parsed_message.as_ref() {
for (name, value) in self.outgoing_custom_headers(parsed) {
request = request.header(name, value);
}
}
for (key, value) in &self.config.headers {
request = request.header(key.as_str(), value.as_str());
}
#[cfg(feature = "oauth-client")]
let initial_scope_revision = match &self.scope_escalation {
Some(runtime) => runtime.state.lock().await.revision,
None => 0,
};
#[cfg(feature = "oauth-client")]
if let Some(ref provider) = self.token_provider {
let token = provider
.get_token()
.await
.map_err(|e| Error::Transport(format!("Token provider error: {}", e)))?;
request = request.headers(bearer_headers(&token).map_err(Error::Transport)?);
}
let request = request.body(message.to_string());
if !is_notification && (self.session_id.is_some() || is_modern_request) {
let tx = self.incoming_tx.clone();
let req_id = parsed_message
.as_ref()
.and_then(|value| value.get("id"))
.cloned();
let request_id = req_id
.clone()
.and_then(|value| serde_json::from_value(value).ok());
let is_subscription = method == Some("subscriptions/listen");
let connected = self.connected.clone();
let last_event_id = self.last_event_id.clone();
let sse_retry_delay = self.sse_retry_delay.clone();
let sse_reconnect_signal = self.sse_reconnect_signal.clone();
let max_sse_event_size = self.config.max_sse_event_size;
let request_resource = self.url.clone();
#[cfg(feature = "oauth-client")]
let token_provider = self.token_provider.clone();
#[cfg(feature = "oauth-client")]
let scope_escalation = self.scope_escalation.clone();
self.request_tasks.retain(|_, task| !task.is_finished());
let task = tokio::spawn(async move {
let response_result = send_http_request(
request,
&request_resource,
&operation,
#[cfg(feature = "oauth-client")]
token_provider,
#[cfg(feature = "oauth-client")]
scope_escalation,
#[cfg(feature = "oauth-client")]
initial_scope_revision,
)
.await;
let response = match response_result {
Ok(r) => r,
Err(e) => {
let connection_failed = e.connection_failed;
tracing::error!(error = %e.message, "Background HTTP request failed");
if let Some(id) = &req_id {
let _ = tx.send(transport_error_frame(id, &e.message)).await;
}
if connection_failed {
connected.store(false, Ordering::Release);
}
return;
}
};
let status = response.status();
if status == reqwest::StatusCode::ACCEPTED {
return;
}
if !status.is_success() {
let status_error = http_status_error(status, response.headers());
let body = response.text().await.unwrap_or_default();
if !body.is_empty()
&& let Ok(mut v) = serde_json::from_str::<serde_json::Value>(&body)
&& is_jsonrpc_error_response(&v)
{
let is_session_signal =
v.pointer("/error/code").and_then(|c| c.as_i64()) == Some(-32005);
if !is_session_signal
&& v.get("id").is_none_or(|id| id.is_null())
&& let Some(id) = &req_id
{
v["id"] = id.clone();
}
let _ = tx.send(v.to_string()).await;
return;
}
tracing::error!(status = %status, body = %body, "HTTP error from server");
if let Some(id) = &req_id {
let _ = tx.send(transport_error_frame(id, &status_error)).await;
}
connected.store(false, Ordering::Release);
return;
}
let is_sse = response
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.is_some_and(|ct| ct.contains("text/event-stream"));
if is_sse {
let mut stream = response.bytes_stream();
let mut parser = SseParser::with_limit(max_sse_event_size);
let mut had_retry = false;
let mut had_data = false;
let mut subscription_acknowledged = false;
use futures::StreamExt;
while let Some(result) = stream.next().await {
match result {
Ok(bytes) => {
let text = String::from_utf8_lossy(&bytes);
let events = match parser.feed(&text) {
Ok(events) => events,
Err(e) => {
tracing::error!(error = %e, "POST SSE stream terminated");
connected.store(false, Ordering::Release);
return;
}
};
for event in events {
if let Some(ref id) = event.id {
*last_event_id.write().await = Some(id.clone());
}
if let Some(retry_ms) = event.retry {
*sse_retry_delay.write().await =
Some(Duration::from_millis(retry_ms));
had_retry = true;
}
if !event.data.is_empty() {
had_data = true;
let value =
serde_json::from_str::<serde_json::Value>(&event.data);
let value = match value {
Ok(value) => value,
Err(error) if is_subscription => {
if let Some(id) = &req_id {
let _ = tx
.send(transport_error_frame(
id,
&format!(
"subscription stream returned invalid JSON: {error}"
),
))
.await;
}
return;
}
Err(_) => {
let _ = tx.send(event.data).await;
continue;
}
};
let is_terminal =
value.get("id").zip(req_id.as_ref()).is_some_and(
|(actual, expected)| {
json_request_ids_match(actual, expected)
},
) && (value.get("result").is_some()
|| value.get("error").is_some());
if is_subscription {
let violation = if is_terminal {
if value.get("error").is_some() {
None
} else if !subscription_acknowledged {
Some(
"subscriptions/listen completed before acknowledgment",
)
} else if !value
.pointer(
"/result/_meta/io.modelcontextprotocol~1subscriptionId",
)
.zip(req_id.as_ref())
.is_some_and(|(actual, expected)| {
json_request_ids_match(actual, expected)
})
{
Some(
"subscriptions/listen result carried a missing or mismatched subscription ID",
)
} else {
None
}
} else if value.get("method").is_some()
&& value.get("id").is_none()
{
let correlated = value
.pointer(
"/params/_meta/io.modelcontextprotocol~1subscriptionId",
)
.zip(req_id.as_ref())
.is_some_and(|(actual, expected)| {
json_request_ids_match(actual, expected)
});
let is_acknowledgment = value
.get("method")
.and_then(serde_json::Value::as_str)
== Some(
notifications::SUBSCRIPTIONS_ACKNOWLEDGED,
);
if !correlated {
Some(
"subscription notification carried a missing or mismatched subscription ID",
)
} else if !subscription_acknowledged
&& !is_acknowledgment
{
Some(
"subscription notification arrived before acknowledgment",
)
} else if subscription_acknowledged
&& is_acknowledgment
{
Some(
"subscription stream sent a duplicate acknowledgment",
)
} else {
if is_acknowledgment {
subscription_acknowledged = true;
}
None
}
} else {
Some(
"subscription stream returned an unrelated JSON-RPC message",
)
};
if let Some(message) = violation {
if let Some(id) = &req_id {
let _ = tx
.send(transport_error_frame(id, message))
.await;
}
return;
}
}
let _ = tx.send(event.data).await;
if is_terminal {
return;
}
}
}
}
Err(e) => {
tracing::warn!(error = %e, "POST SSE stream error");
break;
}
}
}
if !is_modern_request && had_retry && !had_data {
sse_reconnect_signal.notify_one();
} else {
if let Some(id) = &req_id {
let reason = if had_data {
"server closed the response stream before the final reply"
} else {
"server closed the response stream without a reply"
};
let _ = tx.send(transport_error_frame(id, reason)).await;
}
}
} else {
match response.text().await {
Ok(body) if !body.is_empty() => {
let msgs = extract_json_messages(&body);
if msgs.is_empty() {
if let Some(id) = &req_id {
let _ = tx
.send(transport_error_frame(
id,
"server returned an unparseable response body",
))
.await;
}
} else {
for msg in msgs {
let _ = tx.send(msg).await;
}
}
}
Ok(_) => {
if let Some(id) = &req_id {
let _ = tx
.send(transport_error_frame(
id,
"server returned an empty response body",
))
.await;
}
}
Err(e) => {
tracing::error!(error = %e, "Failed to read response body");
if let Some(id) = &req_id {
let _ = tx
.send(transport_error_frame(
id,
&format!("failed to read response body: {e}"),
))
.await;
}
connected.store(false, Ordering::Release);
}
}
}
});
if let Some(request_id) = request_id {
self.request_tasks.insert(request_id, task);
}
return Ok(());
}
let response = send_http_request(
request,
&self.url,
&operation,
#[cfg(feature = "oauth-client")]
self.token_provider.clone(),
#[cfg(feature = "oauth-client")]
self.scope_escalation.clone(),
#[cfg(feature = "oauth-client")]
initial_scope_revision,
)
.await
.map_err(|e| Error::Transport(e.message))?;
let status = response.status();
let new_session_id = response
.headers()
.get("mcp-session-id")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let new_protocol_version = response
.headers()
.get("mcp-protocol-version")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
if status == reqwest::StatusCode::ACCEPTED {
if !is_modern_request && let Some(sid) = new_session_id {
self.session_id = Some(sid);
}
if let Some(pv) = new_protocol_version {
self.protocol_version = Some(pv);
}
return Ok(());
}
if !status.is_success() {
#[cfg(feature = "oauth-client")]
let status_error = http_status_error(status, response.headers());
let body = response.text().await.unwrap_or_default();
if is_modern_request
&& let Ok(mut error) = serde_json::from_str::<serde_json::Value>(&body)
&& is_jsonrpc_error_response(&error)
{
if error.get("id").is_none_or(serde_json::Value::is_null)
&& let Some(id) = parsed_message.as_ref().and_then(|value| value.get("id"))
{
error["id"] = id.clone();
}
self.incoming_tx
.send(error.to_string())
.await
.map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
return Ok(());
}
if status == reqwest::StatusCode::NOT_FOUND
&& self.config.session_recovery
&& self.session_id.is_some()
{
return Err(Error::SessionExpired);
}
if status == reqwest::StatusCode::NOT_FOUND && self.session_id.is_none() {
return Err(Error::Transport(format!(
"HTTP 404 from {}: MCP endpoint not found (check the endpoint path; \
some servers serve MCP at the root, others at /mcp)",
self.url
)));
}
#[cfg(feature = "oauth-client")]
if status == reqwest::StatusCode::FORBIDDEN
&& status_error.contains("insufficient_scope")
{
return Err(Error::Transport(if body.is_empty() {
status_error
} else {
format!("{status_error}: {body}")
}));
}
return Err(Error::Transport(format!(
"HTTP {status} from server: {body}"
)));
}
if !is_modern_request && let Some(sid) = new_session_id {
let is_new_session = self.session_id.is_none();
self.session_id = Some(sid);
if is_new_session && self.config.auto_sse {
self.start_sse_stream();
}
}
if let Some(pv) = new_protocol_version {
self.protocol_version = Some(pv);
}
let body = response
.text()
.await
.map_err(|e| Error::Transport(format!("Failed to read response: {}", e)))?;
for msg in extract_json_messages(&body) {
self.incoming_tx
.send(msg)
.await
.map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
}
Ok(())
}
async fn recv(&mut self) -> Result<Option<String>> {
match self.incoming_rx.recv().await {
Some(msg) => Ok(Some(self.normalize_incoming_message(msg))),
None => {
self.connected.store(false, Ordering::Release);
Ok(None)
}
}
}
fn is_connected(&self) -> bool {
self.connected.load(Ordering::Acquire)
}
async fn close(&mut self) -> Result<()> {
self.connected.store(false, Ordering::Release);
for (_, task) in self.request_tasks.drain() {
task.abort();
}
if let Some(task) = self.sse_task.take() {
task.abort();
}
if let Some(ref session_id) = self.session_id {
let mut request = self
.client
.delete(&self.url)
.header("mcp-session-id", session_id)
.timeout(Duration::from_secs(5));
for (key, value) in &self.config.headers {
request = request.header(key.as_str(), value.as_str());
}
#[cfg(feature = "oauth-client")]
if let Some(ref provider) = self.token_provider
&& let Ok(token) = provider.get_token().await
&& let Ok(headers) = bearer_headers(&token)
{
request = request.headers(headers);
}
let _ = request.send().await;
}
self.session_id = None;
Ok(())
}
async fn reset_session(&mut self) {
tracing::info!("Resetting session for re-initialization");
for (_, task) in self.request_tasks.drain() {
task.abort();
}
if let Some(task) = self.sse_task.take() {
task.abort();
}
self.session_id = None;
self.protocol_version = None;
*self.last_event_id.write().await = None;
*self.sse_retry_delay.write().await = None;
while self.incoming_rx.try_recv().is_ok() {}
}
fn supports_session_recovery(&self) -> bool {
self.config.session_recovery
}
async fn cancel_request(&mut self, request_id: &RequestId) -> Result<()> {
if let Some(task) = self.request_tasks.remove(request_id) {
task.abort();
let _ = task.await;
}
Ok(())
}
}
struct SseLoopParams {
url: String,
client: reqwest::Client,
session_id: String,
protocol_version: Option<String>,
tx: mpsc::Sender<String>,
last_event_id: Arc<RwLock<Option<String>>>,
sse_retry_delay: Arc<RwLock<Option<Duration>>>,
reconnect_signal: Arc<Notify>,
connected: Arc<AtomicBool>,
config: HttpClientConfig,
#[cfg(feature = "oauth-client")]
token_provider: Option<Arc<dyn TokenProvider>>,
}
async fn sse_stream_loop(params: SseLoopParams) {
let SseLoopParams {
url,
client,
session_id,
protocol_version,
tx,
last_event_id,
sse_retry_delay,
reconnect_signal,
connected,
config,
#[cfg(feature = "oauth-client")]
token_provider,
} = params;
let mut reconnect_attempts = 0u32;
loop {
if !connected.load(Ordering::Acquire) {
break;
}
let mut request = client
.get(&url)
.header("Accept", "text/event-stream")
.header("mcp-session-id", &session_id);
if let Some(ref version) = protocol_version {
request = request.header("mcp-protocol-version", version);
}
for (key, value) in &config.headers {
request = request.header(key.as_str(), value.as_str());
}
#[cfg(feature = "oauth-client")]
if let Some(ref provider) = token_provider {
match provider.get_token().await {
Ok(token) => match bearer_headers(&token) {
Ok(headers) => request = request.headers(headers),
Err(error) => {
tracing::warn!(%error, "Token provider failed for SSE connection");
break;
}
},
Err(e) => {
tracing::warn!(error = %e, "Token provider failed for SSE connection");
break;
}
}
}
if let Some(ref lei) = *last_event_id.read().await {
request = request.header("Last-Event-ID", lei.clone());
}
let response = match request.send().await {
Ok(r) if r.status().is_success() => {
reconnect_attempts = 0;
r
}
Ok(r) => {
tracing::warn!(status = %r.status(), "SSE connection rejected");
break;
}
Err(e) => {
tracing::warn!(error = %e, "SSE connection failed");
if !config.sse_reconnect || reconnect_attempts >= config.max_sse_reconnect_attempts
{
break;
}
reconnect_attempts += 1;
let delay = sse_retry_delay
.read()
.await
.unwrap_or(config.sse_reconnect_delay);
tokio::time::sleep(delay).await;
continue;
}
};
let mut stream = response.bytes_stream();
let mut parser = SseParser::with_limit(config.max_sse_event_size);
use futures::StreamExt;
loop {
tokio::select! {
chunk = stream.next() => {
match chunk {
Some(Ok(bytes)) => {
let text = String::from_utf8_lossy(&bytes);
let events = match parser.feed(&text) {
Ok(events) => events,
Err(e) => {
tracing::error!(error = %e, "SSE stream terminated");
connected.store(false, Ordering::Release);
return;
}
};
for event in events {
if let Some(ref id) = event.id {
*last_event_id.write().await = Some(id.clone());
}
if let Some(retry_ms) = event.retry {
*sse_retry_delay.write().await = Some(Duration::from_millis(retry_ms));
}
if !event.data.is_empty() && tx.send(event.data).await.is_err() {
return; }
}
}
Some(Err(e)) => {
tracing::warn!(error = %e, "SSE stream error");
break;
}
None => {
tracing::debug!("SSE stream ended");
break;
}
}
}
_ = reconnect_signal.notified() => {
tracing::debug!("SSE reconnect signal received, closing current stream");
break;
}
}
}
if !config.sse_reconnect
|| !connected.load(Ordering::Acquire)
|| reconnect_attempts >= config.max_sse_reconnect_attempts
{
break;
}
reconnect_attempts += 1;
let delay = sse_retry_delay
.read()
.await
.unwrap_or(config.sse_reconnect_delay);
tracing::info!(
attempt = reconnect_attempts,
max = config.max_sse_reconnect_attempts,
delay_ms = delay.as_millis() as u64,
"Reconnecting SSE stream"
);
tokio::time::sleep(delay).await;
}
}
fn transport_error_frame(id: &serde_json::Value, message: &str) -> String {
serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"error": { "code": -32000, "message": message },
})
.to_string()
}
fn json_request_ids_match(left: &serde_json::Value, right: &serde_json::Value) -> bool {
left == right
|| match (left, right) {
(serde_json::Value::Number(number), serde_json::Value::String(value))
| (serde_json::Value::String(value), serde_json::Value::Number(number)) => number
.as_i64()
.is_some_and(|number| value.parse::<i64>() == Ok(number)),
_ => false,
}
}
fn extract_json_messages(body: &str) -> Vec<String> {
let trimmed = body.trim();
if trimmed.is_empty() {
return Vec::new();
}
let looks_like_sse = trimmed.starts_with("event:")
|| trimmed.starts_with("data:")
|| trimmed.starts_with("id:")
|| trimmed.starts_with(':');
if looks_like_sse {
let mut parser = SseParser::new();
let events = parser.feed(body).unwrap_or_default();
events.into_iter().map(|e| e.data).collect()
} else {
vec![trimmed.to_string()]
}
}
#[derive(Debug)]
struct SseEvent {
id: Option<String>,
data: String,
retry: Option<u64>,
}
struct SseParser {
buffer: String,
current_id: Option<String>,
current_data: Vec<String>,
current_retry: Option<u64>,
data_len: usize,
max_event_size: usize,
}
impl SseParser {
fn new() -> Self {
Self::with_limit(usize::MAX)
}
fn with_limit(max_event_size: usize) -> Self {
Self {
buffer: String::new(),
current_id: None,
current_data: Vec::new(),
current_retry: None,
data_len: 0,
max_event_size,
}
}
fn feed(&mut self, text: &str) -> Result<Vec<SseEvent>> {
self.buffer.push_str(text);
let mut events = Vec::new();
while let Some(newline_pos) = self.buffer.find('\n') {
let line = self.buffer[..newline_pos]
.trim_end_matches('\r')
.to_string();
self.buffer = self.buffer[newline_pos + 1..].to_string();
if line.is_empty() {
if !self.current_data.is_empty() || self.current_retry.is_some() {
events.push(SseEvent {
id: self.current_id.take(),
data: self.current_data.join("\n"),
retry: self.current_retry.take(),
});
self.current_data.clear();
self.data_len = 0;
}
self.current_id = None;
self.current_retry = None;
} else if let Some(value) = line.strip_prefix("id:") {
let trimmed = value.trim();
if !trimmed.is_empty() {
self.current_id = Some(trimmed.to_string());
}
} else if let Some(value) = line.strip_prefix("data:") {
let data = value.trim().to_string();
self.data_len += data.len();
self.current_data.push(data);
} else if let Some(value) = line.strip_prefix("retry:") {
self.current_retry = value.trim().parse().ok();
}
}
let buffered = self.buffer.len() + self.data_len;
if buffered > self.max_event_size {
return Err(Error::SseEventTooLarge {
size: buffered,
limit: self.max_event_size,
});
}
Ok(events)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_complete_event() {
let mut parser = SseParser::new();
let events = parser
.feed("id: 1\nevent: message\ndata: {\"hello\":\"world\"}\n\n")
.unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, Some("1".to_string()));
assert_eq!(events[0].data, "{\"hello\":\"world\"}");
}
#[test]
fn test_parse_multiple_events() {
let mut parser = SseParser::new();
let events = parser
.feed("id: 1\ndata: first\n\nid: 2\ndata: second\n\nid: 3\ndata: third\n\n")
.unwrap();
assert_eq!(events.len(), 3);
assert_eq!(events[0].data, "first");
assert_eq!(events[1].data, "second");
assert_eq!(events[2].data, "third");
assert_eq!(events[0].id, Some("1".to_string()));
assert_eq!(events[1].id, Some("2".to_string()));
assert_eq!(events[2].id, Some("3".to_string()));
}
#[test]
fn test_parse_partial_chunks() {
let mut parser = SseParser::new();
let events = parser.feed("id: 1\nda").unwrap();
assert!(events.is_empty());
let events = parser.feed("ta: hello\n\n").unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, Some("1".to_string()));
assert_eq!(events[0].data, "hello");
}
#[test]
fn test_parse_multiline_data() {
let mut parser = SseParser::new();
let events = parser
.feed("id: 1\ndata: line1\ndata: line2\ndata: line3\n\n")
.unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "line1\nline2\nline3");
}
#[test]
fn test_parse_comment_lines() {
let mut parser = SseParser::new();
let events = parser.feed(": keep-alive\nid: 1\ndata: hello\n\n").unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "hello");
}
#[test]
fn test_parse_event_without_id() {
let mut parser = SseParser::new();
let events = parser.feed("data: no-id-event\n\n").unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, None);
assert_eq!(events[0].data, "no-id-event");
}
#[test]
fn test_empty_data_no_event() {
let mut parser = SseParser::new();
let events = parser.feed("id: 1\n\n").unwrap();
assert!(events.is_empty());
}
#[test]
fn test_parse_crlf_line_endings() {
let mut parser = SseParser::new();
let events = parser.feed("id: 1\r\ndata: crlf\r\n\r\n").unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "crlf");
}
#[test]
fn test_parse_json_data() {
let mut parser = SseParser::new();
let json = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{"token":"t1","progress":50}}"#;
let input = format!("id: 42\nevent: message\ndata: {}\n\n", json);
let events = parser.feed(&input).unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, Some("42".to_string()));
let parsed: serde_json::Value = serde_json::from_str(&events[0].data).unwrap();
assert_eq!(parsed["method"], "notifications/progress");
}
#[test]
fn test_event_exceeding_limit_is_rejected() {
let mut parser = SseParser::with_limit(64);
let big = "data: ".to_string() + &"x".repeat(128);
let err = parser.feed(&big).unwrap_err();
match err {
Error::SseEventTooLarge { size, limit } => {
assert!(size > 64, "size {} should exceed limit", size);
assert_eq!(limit, 64);
}
other => panic!("expected SseEventTooLarge, got {:?}", other),
}
}
#[test]
fn test_accumulated_data_lines_count_toward_limit() {
let mut parser = SseParser::with_limit(64);
let mut result = Ok(Vec::new());
for _ in 0..10 {
result = parser.feed("data: 0123456789\n");
if result.is_err() {
break;
}
}
assert!(matches!(result, Err(Error::SseEventTooLarge { .. })));
}
#[test]
fn test_events_within_limit_pass() {
let mut parser = SseParser::with_limit(64);
let events = parser.feed("data: hello\n\ndata: world\n\n").unwrap();
assert_eq!(events.len(), 2);
}
#[test]
fn test_default_config() {
let config = HttpClientConfig::default();
assert!(config.auto_sse);
assert_eq!(config.channel_capacity, 256);
assert_eq!(config.request_timeout, Duration::from_secs(30));
assert!(config.sse_reconnect);
assert_eq!(config.sse_reconnect_delay, Duration::from_secs(1));
assert_eq!(config.max_sse_reconnect_attempts, 5);
assert!(config.headers.is_empty());
}
#[test]
fn test_new_transport() {
let transport = HttpClientTransport::new("http://localhost:3000");
assert_eq!(transport.url, "http://localhost:3000");
assert!(transport.session_id.is_none());
assert!(transport.protocol_version.is_none());
assert!(transport.is_connected());
}
#[test]
fn test_with_config() {
let config = HttpClientConfig {
request_timeout: Duration::from_secs(60),
sse_reconnect: false,
..Default::default()
};
let transport = HttpClientTransport::with_config("http://example.com", config);
assert_eq!(transport.url, "http://example.com");
assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
assert!(!transport.config.sse_reconnect);
}
#[test]
fn test_with_client() {
let client = reqwest::Client::new();
let transport = HttpClientTransport::with_client("http://example.com", client);
assert_eq!(transport.url, "http://example.com");
assert!(transport.is_connected());
}
#[test]
fn test_bearer_token() {
let transport =
HttpClientTransport::new("http://localhost:3000").bearer_token("sk-test-token");
assert_eq!(
transport.config.headers.get("Authorization").unwrap(),
"Bearer sk-test-token"
);
}
#[test]
fn test_api_key() {
let transport = HttpClientTransport::new("http://localhost:3000").api_key("sk-api-key-123");
assert_eq!(
transport.config.headers.get("Authorization").unwrap(),
"Bearer sk-api-key-123"
);
}
#[test]
fn test_api_key_header() {
let transport =
HttpClientTransport::new("http://localhost:3000").api_key_header("X-API-Key", "my-key");
assert_eq!(transport.config.headers.get("X-API-Key").unwrap(), "my-key");
assert!(!transport.config.headers.contains_key("Authorization"));
}
#[test]
fn test_basic_auth() {
let transport =
HttpClientTransport::new("http://localhost:3000").basic_auth("admin", "secret");
let header = transport.config.headers.get("Authorization").unwrap();
assert!(header.starts_with("Basic "));
use base64::Engine;
let decoded = base64::engine::general_purpose::STANDARD
.decode(header.strip_prefix("Basic ").unwrap())
.unwrap();
assert_eq!(String::from_utf8(decoded).unwrap(), "admin:secret");
}
#[test]
fn test_custom_header() {
let transport = HttpClientTransport::new("http://localhost:3000")
.header("X-Custom", "value1")
.header("X-Another", "value2");
assert_eq!(transport.config.headers.get("X-Custom").unwrap(), "value1");
assert_eq!(transport.config.headers.get("X-Another").unwrap(), "value2");
}
#[test]
fn test_chaining_with_config() {
let config = HttpClientConfig {
request_timeout: Duration::from_secs(60),
..Default::default()
};
let transport =
HttpClientTransport::with_config("http://localhost:3000", config).bearer_token("tk");
assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
assert_eq!(
transport.config.headers.get("Authorization").unwrap(),
"Bearer tk"
);
}
#[test]
fn test_last_auth_wins() {
let transport = HttpClientTransport::new("http://localhost:3000")
.bearer_token("token1")
.basic_auth("user", "pass");
let header = transport.config.headers.get("Authorization").unwrap();
assert!(header.starts_with("Basic "));
}
#[test]
fn test_config_bearer_token() {
let config = HttpClientConfig::default().bearer_token("tk-123");
assert_eq!(
config.headers.get("Authorization").unwrap(),
"Bearer tk-123"
);
}
#[test]
fn test_config_header() {
let config = HttpClientConfig::default().header("X-Foo", "bar");
assert_eq!(config.headers.get("X-Foo").unwrap(), "bar");
}
#[test]
fn test_config_api_key_header() {
let config = HttpClientConfig::default().api_key_header("X-Key", "secret");
assert_eq!(config.headers.get("X-Key").unwrap(), "secret");
}
#[test]
fn test_config_basic_auth() {
let config = HttpClientConfig::default().basic_auth("user", "pw");
let header = config.headers.get("Authorization").unwrap();
assert!(header.starts_with("Basic "));
}
#[test]
fn sep_2243_encodes_only_unsafe_values() {
assert_eq!(encode_header_value("us west 1"), "us west 1");
assert_eq!(encode_header_value(""), "");
assert_eq!(encode_header_value(" padded "), "=?base64?IHBhZGRlZCA=?=");
assert_eq!(
encode_header_value("Hello, 世界"),
"=?base64?SGVsbG8sIOS4lueVjA==?="
);
}
#[test]
fn oauth_error_body_is_not_misclassified_as_jsonrpc() {
assert!(!is_jsonrpc_error_response(&serde_json::json!({
"error": "insufficient_scope",
"error_description": "Token has insufficient scope"
})));
assert!(is_jsonrpc_error_response(&serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"error": {
"code": -32022,
"message": "Unsupported protocol version"
}
})));
}
#[test]
fn sep_2243_validates_custom_header_annotations() {
let mappings = custom_header_mappings(&serde_json::json!({
"type": "object",
"properties": {
"region": {"type": "string", "x-mcp-header": "Region"},
"priority": {"type": "integer", "x-mcp-header": "Priority"},
"ratio": {"type": "number", "x-mcp-header": "Ratio"}
}
}))
.unwrap();
assert_eq!(mappings.len(), 3);
for invalid in [
serde_json::json!({
"type": "object",
"properties": {"value": {"type": "object", "x-mcp-header": "Value"}}
}),
serde_json::json!({
"type": "object",
"properties": {
"a": {"type": "string", "x-mcp-header": "Region"},
"b": {"type": "string", "x-mcp-header": "region"}
}
}),
serde_json::json!({
"type": "object",
"properties": {"value": {"type": "string", "x-mcp-header": "Bad Header"}}
}),
] {
assert!(custom_header_mappings(&invalid).is_err());
}
}
#[test]
fn sep_2243_filters_invalid_tools_and_caches_valid_mappings() {
let mut transport = HttpClientTransport::new("http://localhost:3000");
transport.protocol_version = Some(crate::protocol::PROTOCOL_VERSION_2026_07_28.to_string());
let normalized = transport.normalize_incoming_message(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"result": {
"tools": [
{
"name": "valid",
"inputSchema": {
"type": "object",
"properties": {
"region": {"type": "string", "x-mcp-header": "Region"}
}
}
},
{
"name": "invalid",
"inputSchema": {
"type": "object",
"properties": {
"value": {"type": "array", "x-mcp-header": "Value"}
}
}
}
]
}
})
.to_string(),
);
let parsed: serde_json::Value = serde_json::from_str(&normalized).unwrap();
assert_eq!(parsed["result"]["tools"].as_array().unwrap().len(), 1);
assert!(transport.tool_header_mappings.contains_key("valid"));
assert!(!transport.tool_header_mappings.contains_key("invalid"));
}
}