use crate::{
agent::cancellation::AgentCancellation,
config::{DEFAULT_MCP_TIMEOUT_SECONDS, McpHttpServerConfig},
mcp::{
McpError, McpResult,
jsonrpc::{JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, RequestId},
oauth::{TokenProvider, parse_www_authenticate_resource_metadata},
protocol::{METHOD_INITIALIZE, METHOD_INITIALIZED, PROTOCOL_VERSION},
sse::{
DEFAULT_MAX_SSE_EVENT_BYTES, SseDecoder, jsonrpc_message_from_event, read_sse_event,
},
},
};
use reqwest::{
StatusCode,
blocking::{Client, Response},
header::{ACCEPT, AUTHORIZATION, CONTENT_TYPE, HeaderMap},
};
use serde_json::Value;
use std::{
collections::HashMap,
io::{BufReader, Read},
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
mpsc::{self, SyncSender},
},
thread::{self, JoinHandle},
time::{Duration, Instant},
};
use tokio::sync::watch;
pub(crate) const MAX_MCP_HTTP_ERROR_BODY_BYTES: usize = 1024;
#[cfg(not(test))]
pub(crate) const MAX_MCP_HTTP_RESPONSE_BYTES: usize = 10 * 1024 * 1024;
#[cfg(test)]
pub(crate) const MAX_MCP_HTTP_RESPONSE_BYTES: usize = 1024;
const HEADER_MCP_SESSION_ID: &str = "Mcp-Session-Id";
const HEADER_MCP_PROTOCOL_VERSION: &str = "MCP-Protocol-Version";
const ACCEPT_POST: &str = "application/json, text/event-stream";
const ACCEPT_GET: &str = "text/event-stream, application/json";
const HTTP_SHUTDOWN_REQUEST_TIMEOUT: Duration = Duration::from_millis(500);
type PendingMap = Arc<Mutex<HashMap<RequestId, SyncSender<McpResult<Value>>>>>;
pub(crate) struct HttpConnection {
client: Client,
url: String,
display_url: String,
server_name: Option<String>,
headers: HeaderMap,
token_provider: Option<TokenProvider>,
session_id: Arc<Mutex<Option<String>>>,
protocol_version: Arc<Mutex<Option<String>>>,
timeout: Duration,
pending: PendingMap,
closed: Arc<AtomicBool>,
notification_stream_unavailable: Arc<Mutex<Option<String>>>,
last_event_id: Arc<Mutex<Option<String>>>,
shutdown_sender: watch::Sender<bool>,
get_thread: Mutex<Option<JoinHandle<()>>>,
}
impl std::fmt::Debug for HttpConnection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HttpConnection")
.field("url", &self.display_url)
.field("server_name", &self.server_name)
.field(
"headers",
&format_args!("<{} headers redacted>", self.headers.len()),
)
.field("timeout", &self.timeout)
.finish_non_exhaustive()
}
}
impl HttpConnection {
#[cfg(test)]
pub(crate) fn connect(config: &McpHttpServerConfig) -> McpResult<Self> {
Self::connect_named(None, config, None)
}
pub(crate) fn connect_named(
server_name: Option<&str>,
config: &McpHttpServerConfig,
mc_home: Option<&std::path::Path>,
) -> McpResult<Self> {
let headers = crate::mcp::headers::resolve_http_headers(&config.headers)?;
if config.oauth.is_some() && has_auth_header(&headers) {
return Err(McpError::Config(
"MCP HTTP OAuth cannot be combined with Authorization or Proxy-Authorization headers"
.to_string(),
));
}
let timeout = Duration::from_secs(config.timeout.unwrap_or(DEFAULT_MCP_TIMEOUT_SECONDS));
let client = Client::builder()
.redirect(reqwest::redirect::Policy::none())
.timeout(timeout)
.build()
.map_err(sanitize_reqwest_error)?;
let token_provider = match &config.oauth {
Some(oauth) => {
let server_name = server_name.ok_or_else(|| {
McpError::Config("MCP OAuth HTTP connection requires server name".to_string())
})?;
let mc_home = mc_home.ok_or_else(|| {
McpError::Config("MCP OAuth HTTP connection requires MC_HOME paths".to_string())
})?;
Some(TokenProvider::new(
mc_home.to_path_buf(),
server_name.to_string(),
config.url.clone(),
oauth.clone(),
client.clone(),
))
}
None => None,
};
let (shutdown_sender, _) = watch::channel(false);
Ok(Self {
client,
url: config.url.clone(),
display_url: sanitize_url(&config.url),
server_name: server_name.map(ToString::to_string),
headers,
token_provider,
session_id: Arc::new(Mutex::new(None)),
protocol_version: Arc::new(Mutex::new(None)),
timeout,
pending: Arc::new(Mutex::new(HashMap::new())),
closed: Arc::new(AtomicBool::new(false)),
notification_stream_unavailable: Arc::new(Mutex::new(None)),
last_event_id: Arc::new(Mutex::new(None)),
shutdown_sender,
get_thread: Mutex::new(None),
})
}
pub(crate) fn send_request(
&self,
id: RequestId,
method: &str,
params: Option<Value>,
cancellation: Option<&AgentCancellation>,
) -> McpResult<Value> {
if self.closed.load(Ordering::SeqCst) {
return Err(McpError::Transport("MCP HTTP server is closed".to_string()));
}
if let Some(cancellation) = cancellation
&& let Err(error) = cancellation.check()
{
return Err(McpError::Transport(error.to_string()));
}
let (sender, receiver) = mpsc::sync_channel(1);
self.pending
.lock()
.map_err(|_| McpError::Transport("pending request map poisoned".to_string()))?
.insert(id.clone(), sender);
let message = serde_json::to_value(JsonRpcRequest::new(id.clone(), method, params))
.map_err(McpError::transport)?;
let posted = self.post_json(message, true, Some(&id));
match posted {
Ok(Some(value)) => {
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
if method == METHOD_INITIALIZE {
self.capture_initialize_result(&value);
}
if let Some(cancellation) = cancellation
&& let Err(error) = cancellation.check()
{
return Err(McpError::Transport(error.to_string()));
}
Ok(value)
}
Ok(None) => self.wait_for_pending(id, receiver, cancellation),
Err(error) => {
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
Err(error)
}
}
}
pub(crate) fn send_notification(&self, method: &str, params: Option<Value>) -> McpResult<()> {
if self.closed.load(Ordering::SeqCst) {
return Err(McpError::Transport("MCP HTTP server is closed".to_string()));
}
let message = serde_json::to_value(JsonRpcNotification::new(method, params))
.map_err(McpError::transport)?;
match self.post_json(message, false, None)? {
None => {
if method == METHOD_INITIALIZED {
self.start_notification_stream();
}
Ok(())
}
Some(_) => {
if method == METHOD_INITIALIZED {
self.start_notification_stream();
}
Ok(())
}
}
}
pub(crate) fn shutdown(&mut self) {
if self.closed.swap(true, Ordering::SeqCst) {
return;
}
let _ = self.shutdown_sender.send(true);
fail_pending(
&self.pending,
McpError::Transport("MCP HTTP server closed".to_string()),
);
if let Ok(mut handle) = self.get_thread.lock()
&& let Some(handle) = handle.take()
{
let _ = handle.join();
}
self.delete_session_best_effort();
}
fn post_json(
&self,
body: Value,
expect_response: bool,
request_id: Option<&RequestId>,
) -> McpResult<Option<Value>> {
let mut response = self.send_post_json(body.clone(), false)?;
if response.status() == StatusCode::UNAUTHORIZED && self.token_provider.is_some() {
response = self.send_post_json(body, true)?;
if response.status() == StatusCode::UNAUTHORIZED {
return Err(McpError::Config(format!(
"authentication failed after token refresh; run magi-code mcp login {}",
self.server_name.as_deref().unwrap_or("<server>")
)));
}
}
self.capture_session_id(response.headers());
match response.status() {
StatusCode::OK if expect_response => self.handle_ok_response(response, request_id),
StatusCode::OK => Err(McpError::Transport(format!(
"MCP HTTP notification expected 202 or 204 from {}",
self.display_url
))),
StatusCode::ACCEPTED | StatusCode::NO_CONTENT if !expect_response => Ok(None),
StatusCode::ACCEPTED | StatusCode::NO_CONTENT => Ok(None),
status if status.is_client_error() || status.is_server_error() => {
Err(http_status_error(
status,
response,
&self.display_url,
&self.headers,
self.server_name.as_deref(),
))
}
status => Err(McpError::Transport(format!(
"MCP HTTP unexpected status {status} from {}",
self.display_url
))),
}
}
fn handle_ok_response(
&self,
response: Response,
request_id: Option<&RequestId>,
) -> McpResult<Option<Value>> {
let content_type = response
.headers()
.get(CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or("")
.to_ascii_lowercase();
if content_type.contains("text/event-stream") {
return self.read_sse_response(response, request_id);
}
if content_type.contains("application/json") || content_type.is_empty() {
let mut bytes = Vec::new();
let mut limited = response.take(MAX_MCP_HTTP_RESPONSE_BYTES as u64 + 1);
limited.read_to_end(&mut bytes).map_err(|_| {
McpError::Transport("MCP HTTP response body read failed".to_string())
})?;
if bytes.len() > MAX_MCP_HTTP_RESPONSE_BYTES {
return Err(McpError::Transport(format!(
"MCP HTTP response body from {} exceeded {} bytes",
self.display_url, MAX_MCP_HTTP_RESPONSE_BYTES
)));
}
let value: Value = serde_json::from_slice(&bytes).map_err(McpError::transport)?;
return self.value_from_jsonrpc_message(value, request_id);
}
Err(McpError::Transport(format!(
"MCP HTTP unsupported content type '{}' from {}",
content_type, self.display_url
)))
}
fn read_sse_response(
&self,
response: Response,
request_id: Option<&RequestId>,
) -> McpResult<Option<Value>> {
let mut reader = BufReader::new(response);
loop {
let event = match read_sse_event(&mut reader, DEFAULT_MAX_SSE_EVENT_BYTES)? {
Some(event) => event,
None => return Ok(None),
};
if let Some(id) = &event.id {
let _ = self
.last_event_id
.lock()
.map(|mut last| *last = Some(id.clone()));
}
let Some(message) = jsonrpc_message_from_event(&event)? else {
continue;
};
if let Some(value) = self.handle_stream_message(message, request_id)? {
return Ok(Some(value));
}
}
}
fn value_from_jsonrpc_message(
&self,
value: Value,
request_id: Option<&RequestId>,
) -> McpResult<Option<Value>> {
let message: JsonRpcMessage = serde_json::from_value(value).map_err(McpError::transport)?;
self.handle_stream_message(message, request_id)
}
fn handle_stream_message(
&self,
message: JsonRpcMessage,
request_id: Option<&RequestId>,
) -> McpResult<Option<Value>> {
match message {
JsonRpcMessage::Response(response) => {
if request_id.is_none_or(|id| *id == response.id) {
Ok(Some(response.result))
} else {
dispatch_pending(&self.pending, response.id, Ok(response.result));
Ok(None)
}
}
JsonRpcMessage::Error(error) => {
if request_id.is_none_or(|id| Some(id) == error.id.as_ref()) {
let data = error.error;
Err(McpError::Protocol {
code: data.code,
message: data.message,
})
} else if let Some(id) = error.id {
let data = error.error;
dispatch_pending(
&self.pending,
id,
Err(McpError::Protocol {
code: data.code,
message: data.message,
}),
);
Ok(None)
} else {
Ok(None)
}
}
JsonRpcMessage::Notification(_) | JsonRpcMessage::Request(_) => Ok(None),
}
}
fn wait_for_pending(
&self,
id: RequestId,
receiver: mpsc::Receiver<McpResult<Value>>,
cancellation: Option<&AgentCancellation>,
) -> McpResult<Value> {
let started = Instant::now();
loop {
if let Some(cancellation) = cancellation
&& let Err(error) = cancellation.check()
{
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
return Err(McpError::Transport(error.to_string()));
}
if let Some(reason) = self
.notification_stream_unavailable
.lock()
.ok()
.and_then(|reason| reason.clone())
{
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
return Err(McpError::Transport(format!(
"MCP HTTP GET SSE unavailable while waiting for response: {reason}"
)));
}
let remaining = self.timeout.saturating_sub(started.elapsed());
if remaining.is_zero() {
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
return Err(McpError::Timeout {
seconds: self.timeout.as_secs(),
});
}
match receiver.recv_timeout(remaining.min(Duration::from_millis(100))) {
Ok(result) => return result,
Err(mpsc::RecvTimeoutError::Timeout) => continue,
Err(mpsc::RecvTimeoutError::Disconnected) => {
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
return Err(McpError::Transport(
"MCP HTTP pending response channel disconnected".to_string(),
));
}
}
}
}
fn send_post_json(&self, body: Value, force_refresh: bool) -> McpResult<Response> {
let builder = self
.request_with_session_headers(self.client.post(&self.url), force_refresh)?
.header(CONTENT_TYPE, "application/json")
.header(ACCEPT, ACCEPT_POST)
.json(&body);
builder
.send()
.map_err(|error| sanitize_reqwest_error_with_timeout(error, self.timeout))
}
fn request_with_session_headers(
&self,
builder: reqwest::blocking::RequestBuilder,
force_refresh: bool,
) -> McpResult<reqwest::blocking::RequestBuilder> {
let mut builder = builder.headers(self.headers.clone());
if let Some(provider) = &self.token_provider {
let token = if force_refresh {
provider.force_refresh_access_token()?
} else {
provider.access_token()?
};
builder = builder.header(AUTHORIZATION, format!("Bearer {token}"));
}
if let Some(session_id) = self.session_id.lock().ok().and_then(|id| id.clone()) {
builder = builder.header(HEADER_MCP_SESSION_ID, session_id);
}
if let Some(protocol_version) = self.protocol_version.lock().ok().and_then(|v| v.clone()) {
builder = builder.header(HEADER_MCP_PROTOCOL_VERSION, protocol_version);
}
Ok(builder)
}
fn capture_session_id(&self, headers: &HeaderMap) {
if let Some(value) = headers
.get(HEADER_MCP_SESSION_ID)
.and_then(|value| value.to_str().ok())
&& let Ok(mut session_id) = self.session_id.lock()
{
*session_id = Some(value.to_string());
}
}
fn capture_initialize_result(&self, value: &Value) {
if let Some(protocol_version) = value.get("protocolVersion").and_then(Value::as_str)
&& protocol_version == PROTOCOL_VERSION
&& let Ok(mut stored) = self.protocol_version.lock()
{
*stored = Some(protocol_version.to_string());
}
}
fn start_notification_stream(&self) {
if self
.get_thread
.lock()
.map_or(true, |handle| handle.is_some())
{
return;
}
let url = self.url.clone();
let display_url = self.display_url.clone();
let headers = self.headers.clone();
let token_provider = self.token_provider.clone();
let session_id = Arc::clone(&self.session_id);
let protocol_version = Arc::clone(&self.protocol_version);
let pending = Arc::clone(&self.pending);
let unavailable = Arc::clone(&self.notification_stream_unavailable);
let timeout = self.timeout;
let last_event_id = Arc::clone(&self.last_event_id);
let shutdown = self.shutdown_sender.subscribe();
let handle = thread::spawn(move || {
let runtime = match tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
{
Ok(runtime) => runtime,
Err(error) => {
record_unavailable(
&unavailable,
format!("MCP HTTP GET SSE runtime failed: {error}"),
);
return;
}
};
let client = match reqwest::Client::builder().connect_timeout(timeout).build() {
Ok(client) => client,
Err(error) => {
record_unavailable(&unavailable, sanitize_reqwest_error(error).to_string());
return;
}
};
runtime.block_on(run_get_stream(GetStreamContext {
client,
url,
display_url,
headers,
token_provider,
session_id,
protocol_version,
pending,
unavailable,
last_event_id,
shutdown,
}));
});
if let Ok(mut slot) = self.get_thread.lock() {
*slot = Some(handle);
}
}
fn delete_session_best_effort(&self) {
let Some(session_id) = self.session_id.lock().ok().and_then(|id| id.clone()) else {
return;
};
let mut builder = self
.client
.delete(&self.url)
.headers(self.headers.clone())
.header(HEADER_MCP_SESSION_ID, session_id)
.timeout(HTTP_SHUTDOWN_REQUEST_TIMEOUT);
if let Some(provider) = &self.token_provider
&& let Ok(token) = provider.access_token()
{
builder = builder.header(AUTHORIZATION, format!("Bearer {token}"));
}
if let Some(protocol_version) = self.protocol_version.lock().ok().and_then(|v| v.clone()) {
builder = builder.header(HEADER_MCP_PROTOCOL_VERSION, protocol_version);
}
let _ = builder.send();
}
}
impl Drop for HttpConnection {
fn drop(&mut self) {
self.shutdown();
}
}
fn has_auth_header(headers: &HeaderMap) -> bool {
headers.keys().any(|name| {
name.as_str().eq_ignore_ascii_case("authorization")
|| name.as_str().eq_ignore_ascii_case("proxy-authorization")
})
}
fn dispatch_pending(pending: &PendingMap, id: RequestId, result: McpResult<Value>) {
if let Ok(mut pending) = pending.lock()
&& let Some(sender) = pending.remove(&id)
{
let _ = sender.send(result);
}
}
struct GetStreamContext {
client: reqwest::Client,
url: String,
display_url: String,
headers: HeaderMap,
token_provider: Option<TokenProvider>,
session_id: Arc<Mutex<Option<String>>>,
protocol_version: Arc<Mutex<Option<String>>>,
pending: PendingMap,
unavailable: Arc<Mutex<Option<String>>>,
last_event_id: Arc<Mutex<Option<String>>>,
shutdown: watch::Receiver<bool>,
}
async fn run_get_stream(context: GetStreamContext) {
let GetStreamContext {
client,
url,
display_url,
headers,
token_provider,
session_id,
protocol_version,
pending,
unavailable,
last_event_id,
mut shutdown,
} = context;
let mut builder = client.get(url).headers(headers).header(ACCEPT, ACCEPT_GET);
if let Some(provider) = &token_provider {
match provider.access_token() {
Ok(token) => builder = builder.header(AUTHORIZATION, format!("Bearer {token}")),
Err(error) => {
record_unavailable(&unavailable, error.to_string());
return;
}
}
}
if let Some(id) = session_id.lock().ok().and_then(|id| id.clone()) {
builder = builder.header(HEADER_MCP_SESSION_ID, id);
}
if let Some(version) = protocol_version.lock().ok().and_then(|v| v.clone()) {
builder = builder.header(HEADER_MCP_PROTOCOL_VERSION, version);
}
let response = tokio::select! {
result = builder.send() => match result {
Ok(response) => response,
Err(error) => {
if !*shutdown.borrow() {
record_unavailable(&unavailable, sanitize_reqwest_error(error).to_string());
}
return;
}
},
_ = shutdown.changed() => return,
};
if *shutdown.borrow() {
return;
}
if !response.status().is_success() {
record_unavailable(
&unavailable,
format!(
"MCP HTTP GET SSE unavailable: status {} from {display_url}",
response.status()
),
);
return;
}
let content_type = response
.headers()
.get(CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or("")
.to_ascii_lowercase();
if !content_type.contains("text/event-stream") {
record_unavailable(
&unavailable,
format!(
"MCP HTTP GET SSE unsupported content type '{content_type}' from {display_url}"
),
);
return;
}
let mut response = response;
let mut decoder = SseDecoder::default();
loop {
let chunk = tokio::select! {
result = response.chunk() => match result {
Ok(Some(chunk)) => chunk,
Ok(None) => {
match decoder.finish(DEFAULT_MAX_SSE_EVENT_BYTES) {
Ok(Some(event)) => {
if handle_get_event(event, &pending, &unavailable, &last_event_id) {
return;
}
}
Ok(None) => {}
Err(error) => {
record_unavailable(&unavailable, error.to_string());
return;
}
}
record_unavailable(&unavailable, "MCP HTTP GET SSE stream closed".to_string());
return;
}
Err(error) => {
if !*shutdown.borrow() {
record_unavailable(&unavailable, sanitize_reqwest_error(error).to_string());
}
return;
}
},
_ = shutdown.changed() => return,
};
match decoder.push_chunk(&chunk, DEFAULT_MAX_SSE_EVENT_BYTES) {
Ok(events) => {
for event in events {
if handle_get_event(event, &pending, &unavailable, &last_event_id) {
return;
}
}
}
Err(error) => {
record_unavailable(&unavailable, error.to_string());
return;
}
}
}
}
fn handle_get_event(
event: crate::mcp::sse::SseEvent,
pending: &PendingMap,
unavailable: &Arc<Mutex<Option<String>>>,
last_event_id: &Arc<Mutex<Option<String>>>,
) -> bool {
if let Some(id) = &event.id {
let _ = last_event_id
.lock()
.map(|mut last| *last = Some(id.clone()));
}
match jsonrpc_message_from_event(&event) {
Ok(Some(JsonRpcMessage::Response(response))) => {
dispatch_pending(pending, response.id, Ok(response.result));
false
}
Ok(Some(JsonRpcMessage::Error(error))) => {
if let Some(id) = error.id {
let data = error.error;
dispatch_pending(
pending,
id,
Err(McpError::Protocol {
code: data.code,
message: data.message,
}),
);
}
false
}
Ok(Some(JsonRpcMessage::Notification(_)) | Some(JsonRpcMessage::Request(_)) | None) => {
false
}
Err(error) => {
record_unavailable(unavailable, error.to_string());
true
}
}
}
fn fail_pending(pending: &PendingMap, error: McpError) {
if let Ok(mut pending) = pending.lock() {
for (_, sender) in pending.drain() {
let _ = sender.send(Err(error.clone()));
}
}
}
fn record_unavailable(slot: &Arc<Mutex<Option<String>>>, reason: String) {
let bounded = bound_text(&reason, MAX_MCP_HTTP_ERROR_BODY_BYTES);
if let Ok(mut slot) = slot.lock() {
*slot = Some(bounded);
}
}
fn http_status_error(
status: StatusCode,
mut response: Response,
display_url: &str,
headers: &HeaderMap,
server_name: Option<&str>,
) -> McpError {
if status == StatusCode::UNAUTHORIZED {
if let Some(value) = response
.headers()
.get(reqwest::header::WWW_AUTHENTICATE)
.and_then(|value| value.to_str().ok())
&& parse_www_authenticate_resource_metadata(value).is_some()
{
let server = server_name.unwrap_or("<server>");
return McpError::Transport(format!(
"server appears to require OAuth; add mcp_servers.{server}.oauth to config and run magi-code mcp login {server}"
));
}
return McpError::Transport(format!("MCP HTTP status {status} from {display_url}"));
}
if status == StatusCode::FORBIDDEN {
return McpError::Transport(format!("MCP HTTP status {status} from {display_url}"));
}
let mut bytes = Vec::new();
let mut limited = response
.by_ref()
.take(MAX_MCP_HTTP_ERROR_BODY_BYTES as u64 + 1);
let _ = limited.read_to_end(&mut bytes);
let truncated = bytes.len() > MAX_MCP_HTTP_ERROR_BODY_BYTES;
bytes.truncate(MAX_MCP_HTTP_ERROR_BODY_BYTES);
let raw_preview = String::from_utf8_lossy(&bytes).replace(['\r', '\n'], " ");
let mut preview = redact_http_error_preview(&raw_preview, headers);
if truncated {
preview.push_str("...");
}
if preview.is_empty() {
McpError::Transport(format!("MCP HTTP status {status} from {display_url}"))
} else {
McpError::Transport(format!(
"MCP HTTP status {status} from {display_url}: {preview}"
))
}
}
fn redact_http_error_preview(preview: &str, headers: &HeaderMap) -> String {
let mut redacted = preview.to_string();
for value in headers.values().filter_map(|value| value.to_str().ok()) {
if value.len() >= 4 {
redacted = redacted.replace(value, "[REDACTED]");
}
}
let words = redacted
.split_whitespace()
.map(redact_secret_word)
.collect::<Vec<_>>();
redact_bearer_tokens(&words.join(" "))
}
fn redact_bearer_tokens(text: &str) -> String {
let mut out = Vec::new();
let mut redact_next = false;
for word in text.split_whitespace() {
if redact_next {
out.push("[REDACTED]".to_string());
redact_next = false;
continue;
}
out.push(word.to_string());
if word.trim_end_matches(':').eq_ignore_ascii_case("bearer") {
redact_next = true;
}
}
out.join(" ")
}
fn redact_secret_word(word: &str) -> String {
let trimmed =
word.trim_matches(|c: char| !c.is_ascii_alphanumeric() && c != '-' && c != '_' && c != '.');
if looks_like_secret_token(trimmed) {
word.replace(trimmed, "[REDACTED]")
} else {
word.to_string()
}
}
fn looks_like_secret_token(token: &str) -> bool {
let lower = token.to_ascii_lowercase();
lower.starts_with("sk-")
|| lower.starts_with("key-")
|| lower.starts_with("api-key-")
|| (token.len() >= 32
&& token.bytes().all(|byte| {
byte.is_ascii_alphanumeric()
|| matches!(byte, b'-' | b'_' | b'.' | b'+' | b'/' | b'=')
}))
}
fn sanitize_reqwest_error_with_timeout(error: reqwest::Error, timeout: Duration) -> McpError {
if error.is_timeout() {
McpError::Timeout {
seconds: timeout.as_secs(),
}
} else {
sanitize_reqwest_error(error)
}
}
fn sanitize_reqwest_error(error: reqwest::Error) -> McpError {
let kind = if error.is_timeout() {
"request timed out"
} else if error.is_connect() {
"connection failed"
} else if error.is_decode() {
"decode failed"
} else if error.is_body() {
"body read failed"
} else {
"request failed"
};
McpError::Transport(format!("MCP HTTP {kind}"))
}
fn sanitize_url(url: &str) -> String {
match reqwest::Url::parse(url) {
Ok(parsed) => {
let host = parsed.host_str().unwrap_or("<unknown>");
let port = parsed
.port()
.map(|port| format!(":{port}"))
.unwrap_or_default();
format!("{}://{}{}{}", parsed.scheme(), host, port, parsed.path())
}
Err(_) => "<invalid-url>".to_string(),
}
}
fn bound_text(text: &str, max: usize) -> String {
if text.len() <= max {
return text.to_string();
}
let end = text
.char_indices()
.rev()
.find(|(i, _)| *i <= max)
.map(|(i, _)| i)
.unwrap_or(0);
let mut bounded = text[..end].to_string();
bounded.push_str("...");
bounded
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
agent::cancellation::AgentCancellation,
config::McpServerConfig,
mcp::{ContentBlock, McpClient},
};
use serde_json::json;
use std::{
collections::BTreeMap,
io::{Read, Write},
net::{TcpListener, TcpStream},
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
},
};
#[derive(Clone, Copy)]
enum Mode {
NormalJson,
PostSse,
GetUnsupported,
StatusError,
StatusBearerSecret,
Unauthorized,
UnauthorizedEchoAuthorization,
SsePriming,
Timeout,
MalformedSse,
SessionMismatch,
OAuthRequireNewToken,
OAuth401ThenOk,
AsyncAcceptedGetClosed,
AuthRequiredMetadata,
HangingGetSse,
HangingGetHeaders,
OversizedJsonResponse,
}
#[derive(Default)]
struct Seen {
session_on_list: bool,
protocol_on_list: bool,
deleted: bool,
get_open: bool,
get_peer_closed: bool,
active_gets: usize,
peer_closed_gets: usize,
request_ids: Vec<Value>,
post_count: usize,
authorization_headers: Vec<String>,
delete_authorization_headers: Vec<String>,
}
struct MockServer {
url: String,
seen: Arc<Mutex<Seen>>,
stop: Arc<AtomicBool>,
handle: Option<JoinHandle<()>>,
}
impl MockServer {
fn start(mode: Mode) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let url = format!("http://{}/mcp", listener.local_addr().unwrap());
let seen = Arc::new(Mutex::new(Seen::default()));
let thread_seen = Arc::clone(&seen);
let stop = Arc::new(AtomicBool::new(false));
let thread_stop = Arc::clone(&stop);
listener.set_nonblocking(true).unwrap();
let handle = thread::spawn(move || {
while !thread_stop.load(Ordering::SeqCst) {
match listener.accept() {
Ok((mut stream, _)) => handle_request(&mut stream, mode, &thread_seen),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(10));
}
Err(_) => break,
}
}
});
Self {
url,
seen,
stop,
handle: Some(handle),
}
}
fn config(&self) -> McpServerConfig {
self.config_with_headers(BTreeMap::new())
}
fn config_with_headers(&self, headers: BTreeMap<String, String>) -> McpServerConfig {
McpServerConfig::Http(McpHttpServerConfig {
url: self.url.clone(),
headers,
oauth: None,
enabled: true,
timeout: Some(1),
})
}
}
impl Drop for MockServer {
fn drop(&mut self) {
self.stop.store(true, Ordering::SeqCst);
let _ = TcpStream::connect(
self.url
.strip_prefix("http://")
.unwrap()
.strip_suffix("/mcp")
.unwrap(),
);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
fn read_http_request(stream: &mut TcpStream) -> String {
let _ = stream.set_read_timeout(Some(Duration::from_millis(100)));
let deadline = Instant::now() + Duration::from_secs(10);
let mut bytes = Vec::new();
let mut buf = [0u8; 8192];
loop {
let n = match stream.read(&mut buf) {
Ok(0) => break,
Ok(n) => n,
Err(error)
if matches!(
error.kind(),
std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut
) && Instant::now() < deadline =>
{
continue;
}
Err(error) => panic!("read mock MCP request: {error}"),
};
bytes.extend_from_slice(&buf[..n]);
let Some(header_end) = bytes.windows(4).position(|window| window == b"\r\n\r\n") else {
continue;
};
let head = String::from_utf8_lossy(&bytes[..header_end]).to_ascii_lowercase();
let content_length = head
.lines()
.find_map(|line| line.strip_prefix("content-length:"))
.and_then(|value| value.trim().parse::<usize>().ok())
.unwrap_or(0);
if bytes.len() >= header_end + 4 + content_length {
break;
}
}
String::from_utf8_lossy(&bytes).to_string()
}
fn handle_request(stream: &mut TcpStream, mode: Mode, seen: &Arc<Mutex<Seen>>) {
if matches!(mode, Mode::Timeout) {
thread::sleep(Duration::from_millis(1500));
return;
}
let request = read_http_request(stream);
if request.is_empty() {
return;
}
let mut parts = request.split("\r\n\r\n");
let head = parts.next().unwrap_or("");
let body = parts.next().unwrap_or("");
let request_line = head.lines().next().unwrap_or("");
if request_line.starts_with("POST ") {
let authorization = head
.lines()
.find(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.unwrap_or("")
.to_string();
let mut seen = seen.lock().unwrap();
seen.post_count += 1;
seen.authorization_headers.push(authorization.clone());
if matches!(mode, Mode::OAuthRequireNewToken)
&& !authorization.contains("Bearer new-access-token")
{
drop(seen);
write_response(stream, "401 Unauthorized", "text/plain", "unauthorized");
return;
}
if matches!(mode, Mode::OAuth401ThenOk) {
let count = seen.post_count;
drop(seen);
if count == 1 {
write_response(stream, "401 Unauthorized", "text/plain", "unauthorized");
return;
}
if !authorization.contains("Bearer new-access-token") {
write_response(stream, "401 Unauthorized", "text/plain", "unauthorized");
return;
}
}
}
if request_line.starts_with("GET ") {
match mode {
Mode::SsePriming => write_response(
stream,
"200 OK",
"text/event-stream",
"id: priming-1\nretry: 1000\ndata:\n\ndata: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/tools/list_changed\",\"params\":{}}\n\n",
),
Mode::AsyncAcceptedGetClosed => {
write_response(stream, "200 OK", "text/event-stream", "")
}
Mode::HangingGetSse | Mode::HangingGetHeaders => {
let mut state = seen.lock().unwrap();
state.get_open = true;
state.active_gets += 1;
drop(state);
if matches!(mode, Mode::HangingGetSse) {
hold_sse_stream_open(stream, seen);
} else {
hold_get_headers_open(stream, seen);
}
}
_ => write_response(stream, "405 Method Not Allowed", "text/plain", "no"),
}
return;
}
if request_line.starts_with("DELETE ") {
let authorization = head
.lines()
.find(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.unwrap_or("")
.to_string();
let mut seen = seen.lock().unwrap();
seen.deleted = true;
seen.delete_authorization_headers.push(authorization);
drop(seen);
write_response(stream, "204 No Content", "text/plain", "");
return;
}
if matches!(mode, Mode::StatusError) {
write_response(
stream,
"500 Internal Server Error",
"text/plain",
&"x".repeat(2048),
);
return;
}
if matches!(mode, Mode::StatusBearerSecret) {
write_response(
stream,
"500 Internal Server Error",
"text/plain",
"upstream failed with Bearer sk-test-not-a-real-secret-token and key-test-not-real",
);
return;
}
if matches!(mode, Mode::UnauthorizedEchoAuthorization) {
let authorization = head
.lines()
.find(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.unwrap_or("authorization: <missing>");
write_response(stream, "401 Unauthorized", "text/html", authorization);
return;
}
if matches!(mode, Mode::Unauthorized) {
write_response(stream, "401 Unauthorized", "text/html", &"u".repeat(2048));
return;
}
if matches!(mode, Mode::AuthRequiredMetadata) {
write_response_with_headers(
stream,
"401 Unauthorized",
"text/plain",
"oauth required",
&[
"WWW-Authenticate: Bearer resource_metadata=\"http://127.0.0.1/.well-known/oauth-protected-resource\"",
],
);
return;
}
if matches!(mode, Mode::OversizedJsonResponse) {
write_response(
stream,
"200 OK",
"application/json",
&"{".repeat(MAX_MCP_HTTP_RESPONSE_BYTES + 1),
);
return;
}
let value: Value = serde_json::from_str(body).unwrap_or_else(|_| json!({}));
let method = value.get("method").and_then(Value::as_str).unwrap_or("");
let id = value.get("id").cloned().unwrap_or(json!(1));
if value.get("id").is_some() {
seen.lock().unwrap().request_ids.push(id.clone());
}
if matches!(mode, Mode::SessionMismatch)
&& method != "initialize"
&& method != "notifications/initialized"
&& !head
.to_ascii_lowercase()
.contains("mcp-session-id: session-1")
{
write_response(stream, "404 Not Found", "text/plain", "missing session");
return;
}
match method {
"initialize" if matches!(mode, Mode::SessionMismatch) => write_response_with_headers(
stream,
"200 OK",
"application/json",
&format!(
"{{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"protocolVersion\":\"2025-03-26\",\"capabilities\":{{\"tools\":{{}}}},\"serverInfo\":{{\"name\":\"mock-http\",\"version\":\"1\"}}}}}}"
),
&["Mcp-Session-Id: session-1"],
),
"initialize" if matches!(mode, Mode::AsyncAcceptedGetClosed) => {
write_response_with_headers(
stream,
"200 OK",
"application/json",
&format!(
"{{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"protocolVersion\":\"2025-03-26\",\"capabilities\":{{\"tools\":{{}}}},\"serverInfo\":{{\"name\":\"mock-http\",\"version\":\"1\"}}}}}}"
),
&[],
)
}
"initialize" => write_response_with_headers(
stream,
"200 OK",
"application/json",
&format!(
"{{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"protocolVersion\":\"2025-03-26\",\"capabilities\":{{\"tools\":{{}}}},\"serverInfo\":{{\"name\":\"mock-http\",\"version\":\"1\"}}}}}}"
),
&["Mcp-Session-Id: session-1"],
),
"notifications/initialized" => {
write_response(stream, "204 No Content", "text/plain", "")
}
"tools/list" => {
let lower = head.to_ascii_lowercase();
let mut seen = seen.lock().unwrap();
seen.session_on_list = lower.contains("mcp-session-id: session-1");
seen.protocol_on_list = lower.contains("mcp-protocol-version: 2025-03-26");
write_response(
stream,
"200 OK",
"application/json",
&format!(
"{{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"tools\":[{{\"name\":\"echo\",\"description\":\"Echo\",\"inputSchema\":{{\"type\":\"object\"}}}}]}}}}"
),
);
}
"tools/call" if matches!(mode, Mode::AsyncAcceptedGetClosed) => {
write_response(stream, "202 Accepted", "text/plain", "")
}
"tools/call" if matches!(mode, Mode::PostSse) => write_response(
stream,
"200 OK",
"text/event-stream",
&format!(
"data: {{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"content\":[{{\"type\":\"text\",\"text\":\"ok-sse\"}}],\"isError\":false}}}}\n\n"
),
),
"tools/call" if matches!(mode, Mode::MalformedSse) => write_response(
stream,
"200 OK",
"text/event-stream",
"data: {not json}\n\n",
),
"tools/call" => write_response(
stream,
"200 OK",
"application/json",
&format!(
"{{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"content\":[{{\"type\":\"text\",\"text\":\"ok\"}}],\"isError\":false}}}}"
),
),
_ => write_response(stream, "404 Not Found", "text/plain", "missing"),
}
}
fn hold_get_headers_open(stream: &mut TcpStream, seen: &Arc<Mutex<Seen>>) {
let _ = stream.set_read_timeout(Some(Duration::from_secs(3)));
let mut byte = [0u8; 1];
if matches!(stream.read(&mut byte), Ok(0) | Err(_)) {
let mut seen = seen.lock().unwrap();
seen.get_peer_closed = true;
seen.active_gets = seen.active_gets.saturating_sub(1);
seen.peer_closed_gets += 1;
}
}
fn hold_sse_stream_open(stream: &mut TcpStream, seen: &Arc<Mutex<Seen>>) {
let response =
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: keep-alive\r\n\r\n";
stream.write_all(response.as_bytes()).unwrap();
stream.flush().unwrap();
let _ = stream.set_read_timeout(Some(Duration::from_secs(3)));
let mut byte = [0u8; 1];
if matches!(stream.read(&mut byte), Ok(0) | Err(_)) {
let mut seen = seen.lock().unwrap();
seen.get_peer_closed = true;
seen.active_gets = seen.active_gets.saturating_sub(1);
seen.peer_closed_gets += 1;
}
}
fn write_response(stream: &mut TcpStream, status: &str, content_type: &str, body: &str) {
write_response_with_headers(stream, status, content_type, body, &[]);
}
fn write_response_with_headers(
stream: &mut TcpStream,
status: &str,
content_type: &str,
body: &str,
extra: &[&str],
) {
let mut response = format!(
"HTTP/1.1 {status}\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\nConnection: close\r\n",
body.len()
);
for header in extra {
response.push_str(header);
response.push_str("\r\n");
}
response.push_str("\r\n");
response.push_str(body);
let _ = stream.write_all(response.as_bytes());
}
#[test]
fn http_initializes_lists_calls_and_deletes_session() {
let server = MockServer::start(Mode::NormalJson);
let mut client = McpClient::connect(&server.config()).unwrap();
let init = client.initialize().unwrap();
assert_eq!(init.server_info.name, "mock-http");
let tools = client.list_tools().unwrap();
assert_eq!(tools[0].name, "echo");
let result = client
.call_tool("echo", Some(json!({"text":"hi"})))
.unwrap();
assert!(matches!(&result.content[0], ContentBlock::Text { text } if text == "ok"));
client.shutdown();
let seen = server.seen.lock().unwrap();
assert!(seen.session_on_list);
assert!(seen.protocol_on_list);
assert!(seen.deleted);
}
#[test]
fn http_parses_post_sse_response() {
let server = MockServer::start(Mode::PostSse);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
let result = client.call_tool("echo", None).unwrap();
assert!(matches!(&result.content[0], ContentBlock::Text { text } if text == "ok-sse"));
client.shutdown();
}
#[test]
fn http_get_sse_unsupported_does_not_break_post_flow() {
let server = MockServer::start(Mode::GetUnsupported);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
assert_eq!(client.list_tools().unwrap().len(), 1);
client.shutdown();
}
#[test]
fn http_status_error_body_is_capped() {
let server = MockServer::start(Mode::StatusError);
let client = McpClient::connect(&server.config()).unwrap();
let error = client.initialize().unwrap_err().to_string();
assert!(error.contains("500"), "{error}");
assert!(error.len() < 1200, "{}", error.len());
}
#[test]
fn http_success_json_body_is_capped() {
let server = MockServer::start(Mode::OversizedJsonResponse);
let client = McpClient::connect(&server.config()).unwrap();
let error = client.initialize().unwrap_err().to_string();
assert!(error.contains("response body"), "{error}");
assert!(error.contains("exceeded"), "{error}");
}
#[test]
fn http_get_sse_priming_does_not_break_post_flow() {
let server = MockServer::start(Mode::SsePriming);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
assert_eq!(client.list_tools().unwrap().len(), 1);
client.shutdown();
}
#[test]
fn http_pending_async_response_fails_when_get_sse_closes() {
let server = MockServer::start(Mode::AsyncAcceptedGetClosed);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
let started = std::time::Instant::now();
let error = client.call_tool("echo", None).unwrap_err().to_string();
assert!(error.contains("GET SSE unavailable"), "{error}");
assert!(
started.elapsed() < Duration::from_secs(1),
"elapsed: {:?}",
started.elapsed()
);
client.shutdown();
}
#[test]
fn http_shutdown_is_bounded_when_get_sse_read_stalls() {
let server = MockServer::start(Mode::HangingGetSse);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
let deadline = std::time::Instant::now() + Duration::from_secs(1);
while !server.seen.lock().unwrap().get_open && std::time::Instant::now() < deadline {
thread::yield_now();
}
let started = std::time::Instant::now();
client.shutdown();
assert!(started.elapsed() < Duration::from_millis(1500));
let seen = server.seen.lock().unwrap();
assert!(seen.get_open);
assert!(seen.get_peer_closed);
assert!(seen.deleted);
drop(seen);
client.shutdown();
}
#[test]
fn http_timeout_is_bounded() {
let server = MockServer::start(Mode::Timeout);
let client = McpClient::connect(&server.config()).unwrap();
let error = client.initialize().unwrap_err().to_string();
assert!(error.contains("timed out"), "{error}");
assert!(!error.contains(&server.url), "{error}");
}
#[test]
fn http_malformed_sse_returns_bounded_error() {
let server = MockServer::start(Mode::MalformedSse);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
let error = client.call_tool("echo", None).unwrap_err().to_string();
assert!(error.contains("SSE JSON parse failed"), "{error}");
assert!(!error.contains("{not json}"), "{error}");
client.shutdown();
}
#[test]
fn http_session_tracking_and_multiple_request_ids_work() {
let server = MockServer::start(Mode::SessionMismatch);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
client.list_tools().unwrap();
client.call_tool("echo", None).unwrap();
client.shutdown();
let seen = server.seen.lock().unwrap();
assert!(seen.session_on_list);
assert!(seen.protocol_on_list);
assert!(seen.request_ids.windows(2).all(|pair| pair[0] != pair[1]));
}
#[test]
fn http_401_error_body_is_capped() {
let server = MockServer::start(Mode::Unauthorized);
let client = McpClient::connect(&server.config()).unwrap();
let error = client.initialize().unwrap_err().to_string();
assert!(error.contains("401"), "{error}");
assert!(error.len() < 1200, "{}", error.len());
}
#[test]
fn http_401_error_omits_body_echoing_resolved_header_secret() {
let env = crate::test_support::env::env_lock();
env.set_var("MCP_HTTP_AUTH_ECHO_TOKEN", "test-token-from-env-123");
let server = MockServer::start(Mode::UnauthorizedEchoAuthorization);
let mut headers = BTreeMap::new();
headers.insert(
"Authorization".to_string(),
"{env:MCP_HTTP_AUTH_ECHO_TOKEN}".to_string(),
);
let client = McpClient::connect(&server.config_with_headers(headers)).unwrap();
let error = client.initialize().unwrap_err().to_string();
assert!(error.contains("401"), "{error}");
assert!(!error.contains("test-token-from-env-123"), "{error}");
assert!(!error.contains("authorization:"), "{error}");
env.remove_var("MCP_HTTP_AUTH_ECHO_TOKEN");
}
#[test]
fn http_500_error_redacts_common_secret_patterns() {
let server = MockServer::start(Mode::StatusBearerSecret);
let client = McpClient::connect(&server.config()).unwrap();
let error = client.initialize().unwrap_err().to_string();
assert!(error.contains("500"), "{error}");
assert!(error.contains("[REDACTED]"), "{error}");
assert!(
!error.contains("sk-test-not-a-real-secret-token"),
"{error}"
);
assert!(!error.contains("key-test-not-real"), "{error}");
}
#[test]
fn http_manager_status_does_not_leak_401_echoed_header_secret() {
let env = crate::test_support::env::env_lock();
env.set_var("MCP_HTTP_MANAGER_TOKEN", "manager-token-from-env-123");
let server = MockServer::start(Mode::UnauthorizedEchoAuthorization);
let mut headers = BTreeMap::new();
headers.insert(
"Authorization".to_string(),
"{env:MCP_HTTP_MANAGER_TOKEN}".to_string(),
);
let mut settings = crate::config::McpServersSettings::new();
settings.insert(
"remote".to_string(),
match server.config_with_headers(headers) {
McpServerConfig::Http(config) => McpServerConfig::Http(config),
McpServerConfig::Stdio(_) => unreachable!(),
},
);
let manager = crate::mcp::manager::McpManager::from_settings(&settings);
let Some(crate::mcp::manager::McpServerStatus::Failed { error, .. }) =
manager.statuses().get("remote")
else {
panic!("expected failed MCP status")
};
assert!(error.contains("401"), "{error}");
assert!(!error.contains("manager-token-from-env-123"), "{error}");
env.remove_var("MCP_HTTP_MANAGER_TOKEN");
}
#[test]
fn http_401_with_resource_metadata_suggests_oauth_without_leaking_header() {
let server = MockServer::start(Mode::AuthRequiredMetadata);
let mut settings = crate::config::McpServersSettings::new();
settings.insert(
"remote".to_string(),
match server.config() {
McpServerConfig::Http(config) => McpServerConfig::Http(config),
McpServerConfig::Stdio(_) => unreachable!(),
},
);
let manager = crate::mcp::manager::McpManager::from_settings(&settings);
let Some(crate::mcp::manager::McpServerStatus::Failed { error, .. }) =
manager.statuses().get("remote")
else {
panic!("expected failed MCP status")
};
assert!(error.contains("server appears to require OAuth"), "{error}");
assert!(error.contains("mcp_servers.remote.oauth"), "{error}");
assert!(error.contains("magi-code mcp login remote"), "{error}");
assert!(!error.contains("resource_metadata"), "{error}");
}
#[test]
fn http_precanceled_request_returns_without_posting() {
let server = MockServer::start(Mode::NormalJson);
let client = McpClient::connect(&server.config()).unwrap();
let cancel_flag = Arc::new(AtomicBool::new(true));
let cancellation = AgentCancellation::new(cancel_flag);
let error = client
.send_request_cancellable("tools/list", None, &cancellation)
.unwrap_err()
.to_string();
assert!(error.contains("prompt canceled"), "{error}");
assert!(server.seen.lock().unwrap().request_ids.is_empty());
}
#[test]
fn http_redirect_does_not_send_headers_to_second_origin() {
let env = crate::test_support::env::env_lock();
env.set_var("MCP_HTTP_REDIRECT_TOKEN", "Bearer redirect-test-token");
let target = TcpListener::bind("127.0.0.1:0").unwrap();
target.set_nonblocking(true).unwrap();
let target_url = format!("http://{}/mcp", target.local_addr().unwrap());
let origin = TcpListener::bind("127.0.0.1:0").unwrap();
origin.set_nonblocking(true).unwrap();
let origin_url = format!("http://{}/mcp", origin.local_addr().unwrap());
let origin_thread = thread::spawn(move || {
let deadline = Instant::now() + Duration::from_secs(2);
loop {
match origin.accept() {
Ok((mut stream, _)) => {
let request = read_http_request(&mut stream).to_ascii_lowercase();
assert!(request.contains("x-team: platform"), "{request}");
assert!(
request.contains("authorization: bearer redirect-test-token"),
"{request}"
);
write_response_with_headers(
&mut stream,
"302 Found",
"text/plain",
"redirect",
&[&format!("Location: {target_url}")],
);
break;
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
if Instant::now() >= deadline {
panic!("redirect origin received no request");
}
thread::sleep(Duration::from_millis(10));
}
Err(error) => panic!("redirect origin accept failed: {error}"),
}
}
});
let mut headers = BTreeMap::new();
headers.insert("X-Team".to_string(), "platform".to_string());
headers.insert(
"Authorization".to_string(),
"{env:MCP_HTTP_REDIRECT_TOKEN}".to_string(),
);
let config = McpHttpServerConfig {
url: origin_url,
headers,
oauth: None,
enabled: true,
timeout: Some(1),
};
let client = HttpConnection::connect(&config).unwrap();
let error = client
.send_request(RequestId::Number(1), METHOD_INITIALIZE, None, None)
.unwrap_err()
.to_string();
assert!(
error.contains("status 302") || error.contains("unexpected status"),
"{error}"
);
env.remove_var("MCP_HTTP_REDIRECT_TOKEN");
origin_thread.join().unwrap();
thread::sleep(Duration::from_millis(100));
assert!(
matches!(target.accept(), Err(error) if error.kind() == std::io::ErrorKind::WouldBlock)
);
}
#[test]
fn http_debug_output_redacts_header_values() {
let mut headers = BTreeMap::new();
headers.insert("X-Team".to_string(), "literal-secret-value".to_string());
let config = McpHttpServerConfig {
url: "http://localhost/mcp?token=url-secret#frag".to_string(),
headers,
oauth: None,
enabled: true,
timeout: Some(1),
};
let connection = HttpConnection::connect(&config).unwrap();
let debug = format!("{connection:?}");
assert!(debug.contains("http://localhost/mcp"), "{debug}");
assert!(!debug.contains("literal-secret-value"), "{debug}");
assert!(!debug.contains("url-secret"), "{debug}");
assert!(!debug.contains("token="), "{debug}");
}
#[test]
fn http_error_strips_url_query_and_header_values() {
let server = MockServer::start(Mode::StatusError);
let mut headers = BTreeMap::new();
headers.insert("X-Team".to_string(), "literal-secret-value".to_string());
let config = McpHttpServerConfig {
url: format!("{}?token=url-secret#frag", server.url),
headers,
oauth: None,
enabled: true,
timeout: Some(1),
};
let client = McpClient::connect(&McpServerConfig::Http(config)).unwrap();
let error = client.initialize().unwrap_err().to_string();
assert!(error.contains("500"), "{error}");
assert!(!error.contains("literal-secret-value"), "{error}");
assert!(!error.contains("url-secret"), "{error}");
assert!(!error.contains("token="), "{error}");
}
#[test]
fn oauth_runtime_rejects_static_authorization_with_oauth() {
let temp = tempfile::TempDir::new().unwrap();
let server = MockServer::start(Mode::NormalJson);
let mut config = match oauth_config(&server.url) {
McpServerConfig::Http(config) => config,
McpServerConfig::Stdio(_) => unreachable!(),
};
config.headers.insert(
"Authorization".to_string(),
"Bearer static-token".to_string(),
);
let error = HttpConnection::connect_named(Some("remote"), &config, Some(temp.path()))
.unwrap_err()
.to_string();
assert!(
error.contains("cannot be combined") || error.contains("Authorization"),
"{error}"
);
}
#[test]
fn oauth_http_injects_authorization_header() {
let temp = tempfile::TempDir::new().unwrap();
let server = MockServer::start(Mode::OAuthRequireNewToken);
write_oauth_token(
temp.path(),
"remote",
&server.url,
"new-access-token",
None,
3600,
);
let config = oauth_config(&server.url);
let mut client =
McpClient::connect_named(Some("remote"), &config, Some(temp.path())).unwrap();
client.initialize().unwrap();
client.shutdown();
let seen = server.seen.lock().unwrap();
assert!(
seen.authorization_headers
.iter()
.any(|h| h.contains("Bearer new-access-token"))
);
assert!(!format!("{:?}", seen.authorization_headers).contains("refresh-token"));
}
#[test]
fn oauth_http_injects_authorization_header_on_delete() {
let _env = crate::test_support::env::env_lock();
let temp = tempfile::TempDir::new().unwrap();
let server = MockServer::start(Mode::OAuthRequireNewToken);
write_oauth_token(
temp.path(),
"remote",
&server.url,
"new-access-token",
None,
3600,
);
let config = oauth_config(&server.url);
let mut client =
McpClient::connect_named(Some("remote"), &config, Some(temp.path())).unwrap();
client.initialize().unwrap();
client.shutdown();
let seen = server.seen.lock().unwrap();
assert!(seen.deleted);
assert!(
seen.delete_authorization_headers
.iter()
.any(|h| h.contains("Bearer new-access-token")),
"{:?}",
seen.delete_authorization_headers
);
}
#[test]
fn oauth_http_proactively_refreshes_expired_token() {
let temp = tempfile::TempDir::new().unwrap();
let server = MockServer::start(Mode::OAuthRequireNewToken);
let token_endpoint = serve_token_endpoint("new-access-token", None);
write_oauth_token(
temp.path(),
"remote",
&server.url,
"old-access-token",
Some(&token_endpoint),
-60,
);
let config = oauth_config(&server.url);
let mut client =
McpClient::connect_named(Some("remote"), &config, Some(temp.path())).unwrap();
client.initialize().unwrap();
client.shutdown();
let stored = crate::mcp::oauth::read_token(temp.path(), "remote")
.unwrap()
.unwrap();
assert_eq!(stored.access_token, "new-access-token");
}
#[test]
fn oauth_http_retries_401_once_after_refresh() {
let temp = tempfile::TempDir::new().unwrap();
let server = MockServer::start(Mode::OAuth401ThenOk);
let token_endpoint = serve_token_endpoint("new-access-token", None);
write_oauth_token(
temp.path(),
"remote",
&server.url,
"old-access-token",
Some(&token_endpoint),
3600,
);
let config = oauth_config(&server.url);
let mut client =
McpClient::connect_named(Some("remote"), &config, Some(temp.path())).unwrap();
client.initialize().unwrap();
client.shutdown();
assert!(server.seen.lock().unwrap().post_count >= 2);
}
fn oauth_config(url: &str) -> McpServerConfig {
McpServerConfig::Http(McpHttpServerConfig {
url: url.to_string(),
headers: BTreeMap::new(),
oauth: Some(crate::config::McpOAuthConfig {
client_id: Some("client-1".to_string()),
scopes: Vec::new(),
authorization_server: None,
}),
enabled: true,
timeout: Some(2),
})
}
fn write_oauth_token(
mc_home: &std::path::Path,
server_name: &str,
server_url: &str,
access_token: &str,
token_endpoint: Option<&str>,
expires_offset_seconds: i64,
) {
crate::mcp::oauth::write_token(
mc_home,
server_name,
&crate::mcp::oauth::StoredToken {
client_id: "client-1".to_string(),
access_token: access_token.to_string(),
refresh_token: Some("refresh-token".to_string()),
expires_at: Some(chrono::Utc::now().timestamp() + expires_offset_seconds),
granted_scopes: Vec::new(),
client_secret: None,
authorization_server: None,
issuer: None,
token_endpoint: token_endpoint.map(ToString::to_string),
resource: None,
server_url: server_url.to_string(),
token_received_at: chrono::Utc::now().timestamp(),
},
)
.unwrap();
}
fn serve_token_endpoint(
access_token: &'static str,
refresh_token: Option<&'static str>,
) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut buf = [0u8; 4096];
let n = stream.read(&mut buf).unwrap();
let request = String::from_utf8_lossy(&buf[..n]);
assert!(request.contains("grant_type=refresh_token"), "{request}");
assert!(!request.contains("new-access-token"), "{request}");
let refresh = refresh_token
.map(|token| format!(",\"refresh_token\":\"{token}\""))
.unwrap_or_default();
let body = format!(
"{{\"access_token\":\"{access_token}\",\"token_type\":\"Bearer\",\"expires_in\":3600{refresh}}}"
);
write_response(&mut stream, "200 OK", "application/json", &body);
});
format!("http://{addr}/token")
}
#[test]
fn http_shutdown_cancels_get_before_headers() {
let server = MockServer::start(Mode::HangingGetHeaders);
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
let deadline = Instant::now() + Duration::from_secs(1);
while !server.seen.lock().unwrap().get_open && Instant::now() < deadline {
thread::yield_now();
}
client.shutdown();
let seen = server.seen.lock().unwrap();
assert_eq!(seen.active_gets, 0);
assert_eq!(seen.peer_closed_gets, 1);
}
#[test]
fn http_repeated_shutdown_lifecycles_close_every_get_peer() {
let server = MockServer::start(Mode::HangingGetSse);
for cycle in 0..3 {
let mut client = McpClient::connect(&server.config()).unwrap();
client.initialize().unwrap();
let deadline = Instant::now() + Duration::from_secs(1);
while server.seen.lock().unwrap().active_gets == 0 && Instant::now() < deadline {
thread::yield_now();
}
client.shutdown();
let seen = server.seen.lock().unwrap();
assert_eq!(seen.active_gets, 0, "cycle {cycle}");
assert_eq!(seen.peer_closed_gets, cycle + 1, "cycle {cycle}");
}
}
#[test]
fn http_pending_tools_call_gets_one_closed_result_across_repeated_shutdown() {
let pending: PendingMap = Arc::new(Mutex::new(HashMap::new()));
let (sender, receiver) = mpsc::sync_channel(1);
pending.lock().unwrap().insert(RequestId::Number(1), sender);
fail_pending(
&pending,
McpError::Transport("MCP HTTP server closed".to_string()),
);
fail_pending(
&pending,
McpError::Transport("MCP HTTP server closed".to_string()),
);
assert_eq!(
receiver.recv().unwrap().unwrap_err().to_string(),
"MCP transport error: MCP HTTP server closed",
);
assert!(receiver.try_recv().is_err());
assert!(pending.lock().unwrap().is_empty());
}
#[test]
fn bound_text_truncates_at_utf8_char_boundary() {
assert_eq!(bound_text("éx", 1), "...");
assert_eq!(bound_text("éx", 2), "é...");
}
#[test]
fn sanitize_url_strips_query_and_fragment() {
assert_eq!(
sanitize_url("https://example.test/mcp?token=secret#frag"),
"https://example.test/mcp"
);
}
}