use async_trait::async_trait;
use futures::Stream;
use reqwest::{Client, Response};
use serde_json::{Deserializer, Value};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use url::Url;
use crate::error::{McpClientResult, TransportError};
use crate::transport::{
ConnectionInfo, EventReceiver, ServerEvent, Transport, TransportCapabilities,
TransportStatistics, TransportType,
};
const MCP_POST_ACCEPT: &str = "application/json, text/event-stream";
#[derive(Debug)]
pub struct HttpTransport {
client: Client,
endpoint: Url,
connected: AtomicBool,
request_counter: AtomicU64,
stats: Arc<parking_lot::Mutex<TransportStatistics>>,
event_sender: parking_lot::Mutex<Option<mpsc::UnboundedSender<ServerEvent>>>,
queued_events: Arc<parking_lot::Mutex<Vec<ServerEvent>>>,
session_id: Arc<parking_lot::Mutex<Option<String>>>,
auth_override: Arc<parking_lot::RwLock<Option<String>>>,
protocol_version: Arc<parking_lot::RwLock<String>>,
}
impl HttpTransport {
pub fn new(endpoint: &str) -> McpClientResult<Self> {
let url = Url::parse(endpoint)
.map_err(|e| TransportError::ConnectionFailed(format!("Invalid URL: {}", e)))?;
if !matches!(url.scheme(), "http" | "https") {
return Err(TransportError::ConnectionFailed(format!(
"Invalid scheme for HTTP transport: {}",
url.scheme()
))
.into());
}
let client = Client::builder()
.timeout(Duration::from_secs(30))
.user_agent("mcp-client/0.1.0")
.http2_keep_alive_interval(Duration::from_secs(30))
.http2_keep_alive_timeout(Duration::from_secs(10))
.http2_keep_alive_while_idle(true)
.build()
.map_err(|e| TransportError::Http(format!("Failed to create HTTP client: {}", e)))?;
Ok(Self {
client,
endpoint: url,
connected: AtomicBool::new(false),
request_counter: AtomicU64::new(0),
stats: Arc::new(parking_lot::Mutex::new(TransportStatistics::default())),
event_sender: parking_lot::Mutex::new(None),
queued_events: Arc::new(parking_lot::Mutex::new(Vec::new())),
session_id: Arc::new(parking_lot::Mutex::new(None)),
auth_override: Arc::new(parking_lot::RwLock::new(None)),
protocol_version: Arc::new(parking_lot::RwLock::new("2025-11-25".to_string())),
})
}
pub fn with_config(
endpoint: &str,
config: &crate::config::ConnectionConfig,
) -> McpClientResult<Self> {
let url = Url::parse(endpoint)
.map_err(|e| TransportError::ConnectionFailed(format!("Invalid URL: {}", e)))?;
if !matches!(url.scheme(), "http" | "https") {
return Err(TransportError::ConnectionFailed(format!(
"Invalid scheme for HTTP transport: {}",
url.scheme()
))
.into());
}
let user_agent = config.user_agent.as_deref().unwrap_or("mcp-client/0.1.0");
let mut builder = Client::builder()
.timeout(Duration::from_secs(30))
.user_agent(user_agent)
.pool_max_idle_per_host(config.pool_settings.max_idle_per_host as usize)
.pool_idle_timeout(config.pool_settings.idle_timeout)
.http2_keep_alive_interval(Duration::from_secs(30))
.http2_keep_alive_timeout(Duration::from_secs(10))
.http2_keep_alive_while_idle(true);
builder = if config.follow_redirects {
builder.redirect(reqwest::redirect::Policy::limited(
config.max_redirects as usize,
))
} else {
builder.redirect(reqwest::redirect::Policy::none())
};
if let Some(ref headers) = config.headers {
let mut header_map = reqwest::header::HeaderMap::new();
for (k, v) in headers {
if let (Ok(name), Ok(value)) = (
reqwest::header::HeaderName::from_bytes(k.as_bytes()),
reqwest::header::HeaderValue::from_str(v),
) {
header_map.insert(name, value);
}
}
builder = builder.default_headers(header_map);
}
let client = builder
.build()
.map_err(|e| TransportError::Http(format!("Failed to create HTTP client: {}", e)))?;
Ok(Self {
client,
endpoint: url,
connected: AtomicBool::new(false),
request_counter: AtomicU64::new(0),
stats: Arc::new(parking_lot::Mutex::new(TransportStatistics::default())),
event_sender: parking_lot::Mutex::new(None),
queued_events: Arc::new(parking_lot::Mutex::new(Vec::new())),
session_id: Arc::new(parking_lot::Mutex::new(None)),
auth_override: Arc::new(parking_lot::RwLock::new(None)),
protocol_version: Arc::new(parking_lot::RwLock::new("2025-11-25".to_string())),
})
}
pub fn with_client(endpoint: &str, client: Client) -> McpClientResult<Self> {
let url = Url::parse(endpoint)
.map_err(|e| TransportError::ConnectionFailed(format!("Invalid URL: {}", e)))?;
Ok(Self {
client,
endpoint: url,
connected: AtomicBool::new(false),
request_counter: AtomicU64::new(0),
stats: Arc::new(parking_lot::Mutex::new(TransportStatistics::default())),
event_sender: parking_lot::Mutex::new(None),
queued_events: Arc::new(parking_lot::Mutex::new(Vec::new())),
session_id: Arc::new(parking_lot::Mutex::new(None)),
auth_override: Arc::new(parking_lot::RwLock::new(None)),
protocol_version: Arc::new(parking_lot::RwLock::new("2025-11-25".to_string())),
})
}
pub fn set_session_id(&self, session_id: String) {
debug!("Setting session ID: {}", session_id);
*self.session_id.lock() = Some(session_id);
}
pub fn set_protocol_version(&self, version: &str) {
*self.protocol_version.write() = version.to_string();
if version == "2026-07-28" {
*self.session_id.lock() = None;
}
}
fn uses_session_header(&self) -> bool {
*self.protocol_version.read() != "2026-07-28"
}
pub fn clear_session_id(&self) {
debug!("Clearing session ID for re-initialization");
*self.session_id.lock() = None;
}
fn next_request_id(&self) -> String {
let counter = self.request_counter.fetch_add(1, Ordering::SeqCst);
format!("req_{}", counter)
}
fn apply_request_metadata_headers(
&self,
mut builder: reqwest::RequestBuilder,
message: &Value,
) -> reqwest::RequestBuilder {
if *self.protocol_version.read() != "2026-07-28" {
return builder;
}
if let Some(method) = message.get("method").and_then(|m| m.as_str()) {
builder = builder.header("Mcp-Method", method);
let name = match method {
"tools/call" | "prompts/get" => {
message.pointer("/params/name").and_then(|v| v.as_str())
}
"resources/read" => message.pointer("/params/uri").and_then(|v| v.as_str()),
_ => None,
};
if let Some(name) = name {
#[cfg(any(feature = "client-bilingual", feature = "client-2026-07-28-only"))]
let header_value = turul_mcp_protocol_2026_07_28::headers::encode_param_value(
&Value::String(name.to_string()),
)
.unwrap_or_else(|| name.to_string());
#[cfg(not(any(feature = "client-bilingual", feature = "client-2026-07-28-only")))]
let header_value = name.to_string();
builder = builder.header("Mcp-Name", header_value);
}
}
builder
}
fn apply_auth_override(&self, builder: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
if let Some(value) = self.auth_override.read().as_ref() {
builder.header(reqwest::header::AUTHORIZATION, value)
} else {
builder
}
}
fn update_stats<F>(&self, update_fn: F)
where
F: FnOnce(&mut TransportStatistics),
{
let mut stats = self.stats.lock();
update_fn(&mut stats);
}
#[doc(hidden)]
pub async fn test_handle_byte_stream<S, B, E>(&self, stream: S) -> McpClientResult<Value>
where
S: Stream<Item = Result<B, E>> + Unpin,
B: AsRef<[u8]>,
E: std::error::Error + Send + Sync + 'static,
{
self.handle_byte_stream(stream).await
}
async fn handle_sse_stream(&self, response: Response) -> McpClientResult<Value> {
use futures::StreamExt;
use tokio::io::AsyncBufReadExt;
let byte_stream = response.bytes_stream();
let reader = tokio_util::io::StreamReader::new(
byte_stream.map(|r| r.map_err(std::io::Error::other)),
);
let mut lines = tokio::io::BufReader::new(reader).lines();
let sender_snapshot = self.event_sender.lock().clone();
parse_sse_lines(
&mut lines,
sender_snapshot,
&self.queued_events,
&self.stats,
)
.await
}
#[doc(hidden)]
pub async fn test_handle_sse_stream<S, B, E>(&self, stream: S) -> McpClientResult<Value>
where
S: Stream<Item = Result<B, E>> + Unpin,
B: AsRef<[u8]>,
E: std::error::Error + Send + Sync + 'static,
{
use futures::StreamExt;
use tokio::io::AsyncBufReadExt;
let mut buffer = Vec::new();
let mut stream = stream;
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(|e| TransportError::Http(format!("Stream error: {}", e)))?;
buffer.extend_from_slice(chunk.as_ref());
}
let cursor = std::io::Cursor::new(buffer);
let mut lines = tokio::io::BufReader::new(cursor).lines();
let sender_snapshot = self.event_sender.lock().clone();
parse_sse_lines(
&mut lines,
sender_snapshot,
&self.queued_events,
&self.stats,
)
.await
}
fn rescue_400_jsonrpc_envelope(result: McpClientResult<Value>) -> McpClientResult<Value> {
match result {
Err(crate::error::McpClientError::Transport(TransportError::HttpStatus {
status: 400,
message,
})) => match serde_json::from_str::<Value>(&message) {
Ok(body) if body.pointer("/error/code").is_some() => Ok(body),
_ => Err(TransportError::HttpStatus {
status: 400,
message,
}
.into()),
},
other => other,
}
}
fn classify_non_2xx(status: u16, message: String) -> crate::error::McpClientError {
if status == 400
&& let Ok(body) = serde_json::from_str::<Value>(&message)
&& let Some(code) = body.pointer("/error/code").and_then(|c| c.as_i64())
{
let msg = body
.pointer("/error/message")
.and_then(|v| v.as_str())
.unwrap_or("server error")
.to_string();
let data = body.pointer("/error/data").cloned();
return crate::error::McpClientError::server_error(code as i32, msg, data);
}
TransportError::HttpStatus { status, message }.into()
}
async fn handle_response(&self, response: Response) -> McpClientResult<Value> {
let status = response.status();
if !status.is_success() {
let error_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
self.update_stats(|stats| {
stats.errors += 1;
stats.last_error = Some(format!("HTTP {}: {}", status, error_text));
});
return Err(TransportError::HttpStatus {
status: status.as_u16(),
message: error_text,
}
.into());
}
if *self.protocol_version.read() != "2026-07-28"
&& let Some(session_header) = response.headers().get("mcp-session-id")
&& let Ok(session_str) = session_header.to_str()
{
debug!("Captured session ID from response: {}", session_str);
*self.session_id.lock() = Some(session_str.to_owned());
}
let content_type = response
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if content_type.contains("application/json") {
self.handle_json_stream(response).await
} else if content_type.contains("text/event-stream") {
self.handle_sse_stream(response).await
} else {
Err(TransportError::Http(format!("Unsupported content type: {}", content_type)).into())
}
}
async fn handle_json_stream(&self, response: Response) -> McpClientResult<Value> {
use futures::StreamExt;
let stream = response
.bytes_stream()
.map(|result| result.map_err(std::io::Error::other));
self.handle_byte_stream(stream).await
}
async fn handle_byte_stream<S, B, E>(&self, mut stream: S) -> McpClientResult<Value>
where
S: Stream<Item = Result<B, E>> + Unpin,
B: AsRef<[u8]>,
E: std::error::Error + Send + Sync + 'static,
{
use futures::StreamExt;
{
let mut guard = self.event_sender.lock();
if guard.is_none() {
let (tx, _rx) = mpsc::unbounded_channel();
*guard = Some(tx);
}
}
let mut buffer = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(|e| TransportError::Http(format!("Stream error: {}", e)))?;
buffer.extend_from_slice(chunk.as_ref());
let mut processed_bytes = 0;
loop {
if processed_bytes >= buffer.len() {
break;
}
let remaining_buffer = &buffer[processed_bytes..];
let mut stream = Deserializer::from_slice(remaining_buffer).into_iter::<Value>();
match stream.next() {
Some(Ok(json)) => {
let consumed_bytes = stream.byte_offset();
debug!(
"Parsed JSON frame ({} bytes): {}",
consumed_bytes,
serde_json::to_string(&json).unwrap_or_default()
);
if let Some(_id) = json.get("id") {
if json.get("result").is_some() {
self.update_stats(|stats| stats.responses_received += 1);
return Ok(json); } else if json.get("error").is_some() {
self.update_stats(|stats| {
stats.errors += 1;
stats.last_error =
json["error"]["message"].as_str().map(|s| s.to_string());
});
return Ok(json); }
}
if json.get("method").is_some() {
let event = if json.get("id").is_some() && !json["id"].is_null() {
ServerEvent::Request(json.clone())
} else {
ServerEvent::Notification(json.clone())
};
let sender_snapshot = self.event_sender.lock().clone();
if let Some(sender) = sender_snapshot {
if sender.send(event.clone()).is_err() {
debug!("Event channel closed, queuing notification");
self.queued_events.lock().push(event);
} else {
debug!(
"Forwarded progress notification via active event channel"
);
}
} else {
debug!("No event listener active, queuing progress notification");
self.queued_events.lock().push(event);
}
}
processed_bytes += consumed_bytes;
}
Some(Err(e)) => {
debug!(
"Incomplete JSON in buffer, waiting for more data. Parse error: {}",
e
);
break;
}
None => {
break;
}
}
}
if processed_bytes > 0 {
buffer.drain(..processed_bytes);
}
}
if !buffer.is_empty()
&& let Ok(json) = serde_json::from_slice::<Value>(&buffer)
&& json.get("id").is_some()
&& (json.get("result").is_some() || json.get("error").is_some())
{
self.update_stats(|stats| stats.responses_received += 1);
return Ok(json);
}
Err(TransportError::Http("Stream ended without final result".to_string()).into())
}
async fn handle_response_with_headers(
&self,
response: Response,
) -> McpClientResult<crate::transport::TransportResponse> {
let status = response.status();
if !status.is_success() {
let error_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
self.update_stats(|stats| {
stats.errors += 1;
stats.last_error = Some(format!("HTTP {}: {}", status, error_text));
});
return Err(TransportError::HttpStatus {
status: status.as_u16(),
message: error_text,
}
.into());
}
if let Some(session_header) = response.headers().get("mcp-session-id")
&& let Ok(session_str) = session_header.to_str()
{
debug!("Captured session ID from response headers: {}", session_str);
*self.session_id.lock() = Some(session_str.to_owned());
}
let mut headers = std::collections::HashMap::new();
for (name, value) in response.headers() {
if let Ok(value_str) = value.to_str() {
headers.insert(name.to_string(), value_str.to_string());
}
}
let content_type = response
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let response_json = if content_type.contains("application/json") {
self.handle_json_stream(response).await?
} else if content_type.contains("text/event-stream") {
self.handle_sse_stream(response).await?
} else {
return Err(TransportError::Http(format!(
"Unsupported content type: {}",
content_type
))
.into());
};
Ok(crate::transport::TransportResponse::new(
response_json,
headers,
))
}
}
#[async_trait]
impl Transport for HttpTransport {
fn transport_type(&self) -> TransportType {
TransportType::Http
}
fn capabilities(&self) -> TransportCapabilities {
TransportCapabilities {
streaming: true,
bidirectional: false,
server_events: true,
max_message_size: None,
persistent: false,
}
}
async fn connect(&self) -> McpClientResult<()> {
debug!(endpoint = %self.endpoint, "Connecting to HTTP endpoint");
self.connected.store(true, Ordering::SeqCst);
info!(endpoint = %self.endpoint, "HTTP transport connected");
Ok(())
}
async fn disconnect(&self) -> McpClientResult<()> {
debug!("Disconnecting HTTP transport");
self.connected.store(false, Ordering::SeqCst);
if let Some(sender) = self.event_sender.lock().take() {
sender.send(ServerEvent::ConnectionLost).ok();
}
info!("HTTP transport disconnected");
Ok(())
}
fn is_connected(&self) -> bool {
self.connected.load(Ordering::SeqCst)
}
async fn send_request(&self, request: Value) -> McpClientResult<Value> {
if !self.is_connected() {
return Err(TransportError::ConnectionFailed("Not connected".to_string()).into());
}
let start_time = Instant::now();
let mut request = request;
if request.get("id").is_none() {
request["id"] = Value::String(self.next_request_id());
}
debug!(
method = request.get("method").and_then(|v| v.as_str()),
id = request.get("id").and_then(|v| v.as_str()),
"Sending HTTP request"
);
self.update_stats(|stats| stats.requests_sent += 1);
let mut req_builder = self
.client
.post(self.endpoint.clone())
.header("Content-Type", "application/json")
.header("Accept", MCP_POST_ACCEPT)
.header("MCP-Protocol-Version", self.protocol_version.read().clone());
req_builder = self.apply_request_metadata_headers(req_builder, &request);
req_builder = self.apply_auth_override(req_builder);
if let Some(ref session_id) = *self.session_id.lock() {
debug!("HTTP request using session ID: {}", session_id);
req_builder = req_builder.header("Mcp-Session-Id", session_id);
} else if self.uses_session_header() {
warn!("HTTP request attempted without session ID - server may reject");
}
let response = req_builder
.json(&request)
.send()
.await
.map_err(|e| TransportError::Http(format!("Failed to send request: {}", e)))?;
let result = Self::rescue_400_jsonrpc_envelope(self.handle_response(response).await)?;
let elapsed = start_time.elapsed();
self.update_stats(|stats| {
let new_avg = if stats.responses_received > 0 {
(stats.avg_response_time_ms * (stats.responses_received - 1) as f64
+ elapsed.as_millis() as f64)
/ stats.responses_received as f64
} else {
elapsed.as_millis() as f64
};
stats.avg_response_time_ms = new_avg;
});
debug!(elapsed_ms = elapsed.as_millis(), "HTTP request completed");
Ok(result)
}
async fn send_request_streaming(
&self,
request: Value,
) -> McpClientResult<tokio::sync::mpsc::UnboundedReceiver<Value>> {
if !self.is_connected() {
return Err(TransportError::ConnectionFailed("Not connected".to_string()).into());
}
let mut request = request;
if request.get("id").is_none() {
request["id"] = Value::String(self.next_request_id());
}
let mut req_builder = self
.client
.post(self.endpoint.clone())
.header("Content-Type", "application/json")
.header("Accept", MCP_POST_ACCEPT)
.header("MCP-Protocol-Version", self.protocol_version.read().clone());
req_builder = self.apply_request_metadata_headers(req_builder, &request);
req_builder = self.apply_auth_override(req_builder);
let response = req_builder
.json(&request)
.send()
.await
.map_err(|e| TransportError::Http(format!("Failed to send request: {}", e)))?;
if !response.status().is_success() {
let status = response.status().as_u16();
let message = response.text().await.unwrap_or_default();
return Err(Self::classify_non_2xx(status, message));
}
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tokio::spawn(async move {
use futures::StreamExt;
let mut stream = response.bytes_stream();
let mut buffer = String::new();
while let Some(chunk) = stream.next().await {
let Ok(chunk) = chunk else { break };
buffer.push_str(&String::from_utf8_lossy(&chunk));
while let Some(pos) = buffer.find("\n\n") {
let event: String = buffer.drain(..pos + 2).collect();
for line in event.lines() {
if let Some(data) = line.strip_prefix("data: ")
&& let Ok(json) = serde_json::from_str::<Value>(data)
&& tx.send(json).is_err()
{
return;
}
}
}
}
});
Ok(rx)
}
async fn send_request_with_extra_headers(
&self,
request: Value,
extra_headers: &[(String, String)],
) -> McpClientResult<Value> {
if !self.is_connected() {
return Err(TransportError::ConnectionFailed("Not connected".to_string()).into());
}
let mut request = request;
if request.get("id").is_none() {
request["id"] = Value::String(self.next_request_id());
}
self.update_stats(|stats| stats.requests_sent += 1);
let mut req_builder = self
.client
.post(self.endpoint.clone())
.header("Content-Type", "application/json")
.header("Accept", MCP_POST_ACCEPT)
.header("MCP-Protocol-Version", self.protocol_version.read().clone());
req_builder = self.apply_request_metadata_headers(req_builder, &request);
for (name, value) in extra_headers {
req_builder = req_builder.header(name, value);
}
req_builder = self.apply_auth_override(req_builder);
if let Some(ref session_id) = *self.session_id.lock() {
req_builder = req_builder.header("Mcp-Session-Id", session_id);
}
let response = req_builder
.json(&request)
.send()
.await
.map_err(|e| TransportError::Http(format!("Failed to send request: {}", e)))?;
Self::rescue_400_jsonrpc_envelope(self.handle_response(response).await)
}
async fn send_request_with_headers(
&self,
request: Value,
) -> McpClientResult<crate::transport::TransportResponse> {
if !self.is_connected() {
return Err(TransportError::ConnectionFailed("Not connected".to_string()).into());
}
let start_time = Instant::now();
let mut request = request;
if request.get("id").is_none() {
request["id"] = Value::String(self.next_request_id());
}
debug!(
method = request.get("method").and_then(|v| v.as_str()),
id = request.get("id").and_then(|v| v.as_str()),
"Sending HTTP request with header extraction"
);
self.update_stats(|stats| stats.requests_sent += 1);
let mut req_builder = self
.client
.post(self.endpoint.clone())
.header("Content-Type", "application/json")
.header("Accept", MCP_POST_ACCEPT)
.header("MCP-Protocol-Version", self.protocol_version.read().clone());
req_builder = self.apply_request_metadata_headers(req_builder, &request);
req_builder = self.apply_auth_override(req_builder);
if let Some(ref session_id) = *self.session_id.lock() {
debug!("HTTP request with headers using session ID: {}", session_id);
req_builder = req_builder.header("Mcp-Session-Id", session_id);
} else {
debug!("HTTP request without session ID (expected for initialize)");
}
let response = req_builder
.json(&request)
.send()
.await
.map_err(|e| TransportError::Http(format!("Failed to send request: {}", e)))?;
let result = self.handle_response_with_headers(response).await?;
let elapsed = start_time.elapsed();
self.update_stats(|stats| {
let new_avg = if stats.responses_received > 0 {
(stats.avg_response_time_ms * (stats.responses_received - 1) as f64
+ elapsed.as_millis() as f64)
/ stats.responses_received as f64
} else {
elapsed.as_millis() as f64
};
stats.avg_response_time_ms = new_avg;
});
debug!(
elapsed_ms = elapsed.as_millis(),
"HTTP request with headers completed"
);
Ok(result)
}
async fn send_notification(&self, notification: Value) -> McpClientResult<()> {
if !self.is_connected() {
return Err(TransportError::ConnectionFailed("Not connected".to_string()).into());
}
debug!(
method = notification.get("method").and_then(|v| v.as_str()),
"Sending HTTP notification"
);
self.update_stats(|stats| stats.notifications_sent += 1);
let mut req_builder = self
.client
.post(self.endpoint.clone())
.header("Accept", MCP_POST_ACCEPT)
.header("Content-Type", "application/json")
.header("MCP-Protocol-Version", self.protocol_version.read().clone());
req_builder = self.apply_request_metadata_headers(req_builder, ¬ification);
req_builder = self.apply_auth_override(req_builder);
if let Some(ref session_id) = *self.session_id.lock() {
debug!("HTTP notification using session ID: {}", session_id);
req_builder = req_builder.header("Mcp-Session-Id", session_id);
} else if self.uses_session_header() {
warn!("HTTP notification attempted without session ID - server may reject");
}
let response = req_builder
.json(¬ification)
.send()
.await
.map_err(|e| TransportError::Http(format!("Failed to send notification: {}", e)))?;
if response.status().is_success() {
debug!("HTTP notification sent successfully");
Ok(())
} else {
let status = response.status();
let error_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
Err(TransportError::HttpStatus {
status: status.as_u16(),
message: error_text,
}
.into())
}
}
async fn send_delete(&self, session_id: &str) -> McpClientResult<()> {
if !self.is_connected() {
return Err(TransportError::ConnectionFailed("Not connected".to_string()).into());
}
info!(
endpoint = %self.endpoint,
session_id = session_id,
"Sending DELETE request for session termination"
);
self.update_stats(|stats| stats.requests_sent += 1);
let start_time = Instant::now();
let req_builder = self
.client
.delete(self.endpoint.clone())
.header("Content-Type", "application/json")
.header("MCP-Protocol-Version", self.protocol_version.read().clone())
.header("Mcp-Session-Id", session_id);
let response = self
.apply_auth_override(req_builder)
.send()
.await
.map_err(|e| TransportError::Http(format!("Failed to send DELETE request: {}", e)))?;
let elapsed = start_time.elapsed();
self.update_stats(|stats| {
stats.responses_received += 1;
let elapsed_ms = elapsed.as_millis() as f64;
if stats.responses_received == 1 {
stats.avg_response_time_ms = elapsed_ms;
} else {
stats.avg_response_time_ms = (stats.avg_response_time_ms
* (stats.responses_received - 1) as f64
+ elapsed_ms)
/ stats.responses_received as f64;
}
});
if response.status().is_success() {
info!(
session_id = session_id,
status = %response.status(),
elapsed_ms = elapsed.as_millis(),
"Session DELETE request completed successfully"
);
Ok(())
} else {
warn!(
session_id = session_id,
status = %response.status(),
elapsed_ms = elapsed.as_millis(),
"DELETE request failed but continuing with cleanup"
);
Ok(())
}
}
async fn start_event_listener(&self) -> McpClientResult<EventReceiver> {
use futures::StreamExt;
let (tx, rx) = mpsc::unbounded_channel();
*self.event_sender.lock() = Some(tx.clone());
let queued_events = {
let mut queue = self.queued_events.lock();
let events = queue.clone();
queue.clear(); events
};
for event in &queued_events {
if tx.send(event.clone()).is_err() {
warn!("Failed to replay queued event - channel already closed");
break;
}
}
if !queued_events.is_empty() {
debug!(
"Replayed {} queued events to new listener",
queued_events.len()
);
}
if !self.is_connected() {
warn!("Not connected - event listener will be inactive");
return Ok(rx);
}
if !self.uses_session_header() {
debug!("GET SSE listener not started: 2026-07-28 has no GET stream");
return Ok(rx);
}
let client = self.client.clone();
let url = self.endpoint.clone();
let session_id = self.session_id.clone();
let auth_override = self.auth_override.clone();
let protocol_version = self.protocol_version.clone();
info!("Starting SSE event listener for GET requests at: {}", url);
tokio::spawn(async move {
loop {
let mut request_builder = client
.get(url.as_str())
.header("Accept", "text/event-stream")
.header("MCP-Protocol-Version", protocol_version.read().clone());
if let Some(value) = auth_override.read().as_ref() {
request_builder = request_builder.header(reqwest::header::AUTHORIZATION, value);
}
let sent_session_id: Option<String> = session_id.lock().clone();
if let Some(ref current_session_id) = sent_session_id {
debug!("SSE request using session ID: {}", current_session_id);
request_builder = request_builder.header("Mcp-Session-Id", current_session_id);
} else {
warn!("SSE request attempted without session ID - server may reject");
}
let response = match request_builder.send().await {
Ok(resp) if resp.status().is_success() => {
debug!("SSE connection established, status: {}", resp.status());
let content_type = resp
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if !content_type.contains("text/event-stream") {
warn!(
"Server returned content-type '{}' instead of 'text/event-stream'",
content_type
);
}
resp
}
Ok(resp) => {
let status = resp.status();
warn!("SSE connection failed with status: {}", status);
if status.is_client_error() {
let mut guard = session_id.lock();
if *guard == sent_session_id {
*guard = None;
}
drop(guard);
let _ = tx.send(ServerEvent::Error(format!(
"SSE GET rejected with HTTP {} — listener exiting",
status
)));
return;
}
if tx
.send(ServerEvent::Error(format!("HTTP {}", status)))
.is_err()
{
debug!("Event channel closed, stopping SSE listener");
return;
}
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
continue;
}
Err(e) => {
warn!("SSE connection error: {}", e);
if tx.send(ServerEvent::Error(e.to_string())).is_err() {
debug!("Event channel closed, stopping SSE listener");
return;
}
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
continue;
}
};
let mut stream = response.bytes_stream();
let mut buffer = String::new();
while let Some(chunk) = stream.next().await {
if tx.is_closed() {
debug!("Event receiver dropped, stopping SSE listener");
return; }
match chunk {
Ok(bytes) => {
buffer.push_str(&String::from_utf8_lossy(&bytes));
while let Some(end) = buffer.find("\n\n") {
let event_text = buffer[..end].to_string();
buffer.drain(..end + 2);
let mut event_type = None;
let mut data = String::new();
let mut id = None;
for line in event_text.lines() {
if let Some(event_value) = line.strip_prefix("event: ") {
event_type = Some(event_value.to_string());
} else if let Some(data_value) = line.strip_prefix("data: ") {
if !data.is_empty() {
data.push('\n');
}
data.push_str(data_value);
} else if let Some(id_value) = line.strip_prefix("id: ") {
id = Some(id_value.to_string());
}
}
if !data.is_empty() {
match serde_json::from_str::<Value>(&data) {
Ok(json) => {
debug!(
"Received SSE event: type={:?}, id={:?}, data={}",
event_type, id, data
);
if json.get("method").is_some()
&& json.get("id").is_some()
{
if tx.send(ServerEvent::Request(json)).is_err() {
debug!("Event channel closed during send");
return;
}
} else if json.get("method").is_some() {
if tx.send(ServerEvent::Notification(json)).is_err()
{
debug!("Event channel closed during send");
return;
}
} else if json.get("id").is_some() {
if tx.send(ServerEvent::Response(json)).is_err() {
debug!("Event channel closed during send");
return;
}
} else {
if tx.send(ServerEvent::Notification(json)).is_err()
{
debug!("Event channel closed during send");
return;
}
}
}
Err(e) => {
warn!(
"Failed to parse SSE data as JSON: {} - data: {}",
e, data
);
}
}
}
}
}
Err(e) => {
warn!("SSE stream error: {}", e);
break;
}
}
}
info!("SSE connection lost, attempting to reconnect...");
if tx.send(ServerEvent::ConnectionLost).is_err() {
debug!("Event channel closed, stopping SSE listener");
return;
}
tokio::time::sleep(std::time::Duration::from_secs(2)).await;
}
});
info!("SSE event listener started successfully");
Ok(rx)
}
fn connection_info(&self) -> ConnectionInfo {
ConnectionInfo {
transport_type: self.transport_type(),
endpoint: self.endpoint.to_string(),
connected: self.is_connected(),
capabilities: self.capabilities(),
metadata: serde_json::json!({
"scheme": self.endpoint.scheme(),
"host": self.endpoint.host_str(),
"port": self.endpoint.port(),
"path": self.endpoint.path()
}),
}
}
async fn health_check(&self) -> McpClientResult<bool> {
if !self.is_connected() {
return Ok(false);
}
let ping_request = serde_json::json!({
"jsonrpc": "2.0",
"id": "health_check",
"method": "ping",
"params": {}
});
match self.send_request(ping_request).await {
Ok(_) => Ok(true),
Err(e) => {
warn!(error = %e, "Health check failed");
Ok(false)
}
}
}
fn set_session_id(&self, session_id: String) {
debug!("HttpTransport: Setting session ID: {}", session_id);
*self.session_id.lock() = Some(session_id);
}
fn clear_session_id(&self) {
debug!("HttpTransport: Clearing session ID for re-initialization");
*self.session_id.lock() = None;
}
async fn update_auth_header(&self, value: Option<String>) {
*self.auth_override.write() = value;
debug!("HttpTransport: Authorization override updated");
}
fn set_protocol_version(&self, version: &str) {
HttpTransport::set_protocol_version(self, version);
}
fn statistics(&self) -> TransportStatistics {
self.stats.lock().clone()
}
}
async fn parse_sse_lines<R: tokio::io::AsyncBufRead + Unpin>(
lines: &mut tokio::io::Lines<R>,
event_sender: Option<mpsc::UnboundedSender<ServerEvent>>,
queued_events: &Arc<parking_lot::Mutex<Vec<ServerEvent>>>,
stats: &Arc<parking_lot::Mutex<TransportStatistics>>,
) -> McpClientResult<Value> {
while let Ok(Some(line)) = lines.next_line().await {
let data = line
.strip_prefix("data: ")
.or_else(|| line.strip_prefix("data:"));
let Some(data) = data else {
continue;
};
let Ok(json) = serde_json::from_str::<Value>(data) else {
continue;
};
if json.get("id").is_some() && (json.get("result").is_some() || json.get("error").is_some())
{
stats.lock().responses_received += 1;
return Ok(json);
}
if json.get("method").is_some() {
let event = if json.get("id").is_some() && !json["id"].is_null() {
ServerEvent::Request(json)
} else {
ServerEvent::Notification(json)
};
if let Some(ref sender) = event_sender {
if sender.send(event.clone()).is_err() {
queued_events.lock().push(event);
}
} else {
queued_events.lock().push(event);
}
}
}
Err(TransportError::Http("SSE stream ended without final result".to_string()).into())
}
use parking_lot;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_http_transport_creation() {
let transport = HttpTransport::new("http://localhost:8080/mcp").unwrap();
assert_eq!(transport.transport_type(), TransportType::Http);
assert!(!transport.is_connected());
}
#[test]
fn test_invalid_url() {
let result = HttpTransport::new("invalid-url");
assert!(result.is_err());
}
#[test]
fn test_invalid_scheme() {
let result = HttpTransport::new("ftp://localhost:8080/mcp");
assert!(result.is_err());
}
#[tokio::test]
async fn test_connect_sets_connected_flag() {
let transport = HttpTransport::new("http://localhost:8080/mcp").unwrap();
assert!(!transport.is_connected());
transport.connect().await.unwrap();
assert!(transport.is_connected());
}
#[tokio::test]
async fn test_connection_info() {
let transport = HttpTransport::new("https://example.com:8443/mcp/endpoint").unwrap();
let info = transport.connection_info();
assert_eq!(info.transport_type, TransportType::Http);
assert_eq!(info.endpoint, "https://example.com:8443/mcp/endpoint");
assert!(!info.connected);
let metadata = info.metadata.as_object().unwrap();
assert_eq!(metadata["scheme"], "https");
assert_eq!(metadata["host"], "example.com");
assert_eq!(metadata["port"], 8443);
assert_eq!(metadata["path"], "/mcp/endpoint");
}
#[test]
fn test_accept_header_contains_both_media_types() {
assert!(MCP_POST_ACCEPT.contains("application/json"));
assert!(MCP_POST_ACCEPT.contains("text/event-stream"));
}
#[test]
fn test_request_id_generation() {
let transport = HttpTransport::new("http://localhost:8080/mcp").unwrap();
let id1 = transport.next_request_id();
let id2 = transport.next_request_id();
assert_ne!(id1, id2);
assert!(id1.starts_with("req_"));
assert!(id2.starts_with("req_"));
}
fn create_test_stream(
data: Vec<Vec<u8>>,
) -> impl Stream<Item = Result<Vec<u8>, std::io::Error>> + Unpin {
futures::stream::iter(data.into_iter().map(Ok))
}
#[tokio::test]
async fn test_handle_byte_stream_single_response() {
let transport = HttpTransport::new("http://localhost:8080/mcp").unwrap();
let response_json = serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"result": {
"tools": []
}
});
let response_bytes = serde_json::to_vec(&response_json).unwrap();
let stream = create_test_stream(vec![response_bytes]);
let result = transport.handle_byte_stream(stream).await.unwrap();
assert_eq!(result["jsonrpc"], "2.0");
assert_eq!(result["id"], 1);
assert!(result.get("result").is_some());
}
#[tokio::test]
async fn test_handle_byte_stream_chunked_response() {
let transport = HttpTransport::new("http://localhost:8080/mcp").unwrap();
let response_json = serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"result": {
"tools": [
{"name": "calculator", "description": "A calculator tool"}
]
}
});
let response_bytes = serde_json::to_vec(&response_json).unwrap();
let chunk1 = response_bytes[..20].to_vec();
let chunk2 = response_bytes[20..40].to_vec();
let chunk3 = response_bytes[40..].to_vec();
let stream = create_test_stream(vec![chunk1, chunk2, chunk3]);
let result = transport.handle_byte_stream(stream).await.unwrap();
assert_eq!(result["jsonrpc"], "2.0");
assert_eq!(result["id"], 2);
assert!(result.get("result").is_some());
}
#[tokio::test]
async fn test_handle_byte_stream_with_progress_notifications() {
let transport = HttpTransport::new("http://localhost:8080/mcp").unwrap();
let progress1 = serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/progress",
"params": {
"progressToken": "test_progress",
"progress": 50,
"total": 100
}
});
let progress2 = serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/progress",
"params": {
"progressToken": "test_progress",
"progress": 100,
"total": 100
}
});
let final_response = serde_json::json!({
"jsonrpc": "2.0",
"id": 3,
"result": {
"success": true
}
});
let progress1_bytes = serde_json::to_vec(&progress1).unwrap();
let progress2_bytes = serde_json::to_vec(&progress2).unwrap();
let final_bytes = serde_json::to_vec(&final_response).unwrap();
let stream = create_test_stream(vec![progress1_bytes, progress2_bytes, final_bytes]);
let mut events = transport.start_event_listener().await.unwrap();
let result = transport.handle_byte_stream(stream).await.unwrap();
assert_eq!(result["jsonrpc"], "2.0");
assert_eq!(result["id"], 3);
assert_eq!(result["result"]["success"], true);
use tokio::time::{Duration, timeout};
let event1 = timeout(Duration::from_millis(100), events.recv()).await;
assert!(event1.is_ok() && event1.unwrap().is_some());
let event2 = timeout(Duration::from_millis(100), events.recv()).await;
assert!(event2.is_ok() && event2.unwrap().is_some());
}
#[tokio::test]
async fn test_progress_event_queue_replay() {
let transport = HttpTransport::new("http://localhost:8080/mcp").unwrap();
let progress1 = serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/progress",
"params": {
"progressToken": "early_progress",
"progress": 25,
"total": 100
}
});
let progress2 = serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/progress",
"params": {
"progressToken": "early_progress",
"progress": 75,
"total": 100
}
});
let final_response = serde_json::json!({
"jsonrpc": "2.0",
"id": 4,
"result": {
"queued_events_test": true
}
});
let progress1_bytes = serde_json::to_vec(&progress1).unwrap();
let progress2_bytes = serde_json::to_vec(&progress2).unwrap();
let final_bytes = serde_json::to_vec(&final_response).unwrap();
let stream = create_test_stream(vec![progress1_bytes, progress2_bytes, final_bytes]);
let result = transport.handle_byte_stream(stream).await.unwrap();
assert_eq!(result["jsonrpc"], "2.0");
assert_eq!(result["id"], 4);
assert_eq!(result["result"]["queued_events_test"], true);
let mut events = transport.start_event_listener().await.unwrap();
use tokio::time::{Duration, timeout};
let event1 = timeout(Duration::from_millis(100), events.recv()).await;
assert!(
event1.is_ok() && event1.unwrap().is_some(),
"Should replay first queued event"
);
let event2 = timeout(Duration::from_millis(100), events.recv()).await;
assert!(
event2.is_ok() && event2.unwrap().is_some(),
"Should replay second queued event"
);
let no_more_events = timeout(Duration::from_millis(50), events.recv()).await;
assert!(
no_more_events.is_err(),
"Should not have any more events after replay"
);
}
#[tokio::test]
async fn test_handle_sse_post_response() {
let sse_body: &[u8] = b"event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":\"req_0\",\"result\":{\"tools\":[]}}\n\n";
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(sse_body.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_sse_stream(stream).await;
assert!(
result.is_ok(),
"SSE POST response should parse successfully"
);
let json = result.unwrap();
assert_eq!(json["id"], "req_0");
assert!(json["result"]["tools"].is_array());
}
#[tokio::test]
async fn test_handle_sse_post_response_error() {
let sse_body: &[u8] = b"data: {\"jsonrpc\":\"2.0\",\"id\":\"req_1\",\"error\":{\"code\":-32600,\"message\":\"Invalid\"}}\n\n";
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(sse_body.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_sse_stream(stream).await;
assert!(result.is_ok());
let json = result.unwrap();
assert_eq!(json["error"]["code"], -32600);
}
#[tokio::test]
async fn test_handle_sse_post_no_final_frame() {
let sse_body: &[u8] =
b"data: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\"}\n\n";
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(sse_body.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_sse_stream(stream).await;
assert!(result.is_err(), "Should fail when no final frame");
}
#[test]
fn test_with_config_custom_headers() {
use std::collections::HashMap;
let mut headers = HashMap::new();
headers.insert("Authorization".to_string(), "Bearer test-token".to_string());
let config = crate::config::ConnectionConfig {
headers: Some(headers),
..Default::default()
};
let transport = HttpTransport::with_config("http://localhost:9999/mcp", &config);
assert!(transport.is_ok());
}
#[test]
fn test_with_config_no_redirects() {
let config = crate::config::ConnectionConfig {
follow_redirects: false,
..Default::default()
};
let transport = HttpTransport::with_config("http://localhost:9999/mcp", &config);
assert!(transport.is_ok());
}
#[test]
fn test_with_config_default() {
let config = crate::config::ConnectionConfig::default();
let transport = HttpTransport::with_config("http://localhost:9999/mcp", &config);
assert!(transport.is_ok());
}
#[test]
fn test_http_transport_advertises_server_events() {
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let caps = transport.capabilities();
assert!(
caps.server_events,
"HttpTransport must advertise server_events"
);
assert!(!caps.bidirectional, "HTTP is not bidirectional");
}
#[tokio::test]
async fn test_server_request_routed_as_request_event() {
let server_request =
br#"{"jsonrpc":"2.0","id":"srv-1","method":"sampling/createMessage","params":{}}"#;
let stream = create_test_stream(vec![server_request.to_vec()]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let _ = transport.test_handle_byte_stream(stream).await;
let events = transport.queued_events.lock();
assert!(!events.is_empty(), "Server request should be queued");
match &events[0] {
crate::transport::ServerEvent::Request(val) => {
assert_eq!(val["method"], "sampling/createMessage");
assert_eq!(val["id"], "srv-1");
}
other => panic!("Expected ServerEvent::Request, got {:?}", other),
}
}
#[tokio::test]
async fn test_jsonrpc_error_preserved_through_transport() {
let error_response = br#"{"jsonrpc":"2.0","id":"req_0","error":{"code":-32602,"message":"Invalid params","data":{"detail":"missing field"}}}"#;
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(error_response.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_byte_stream(stream).await;
assert!(
result.is_ok(),
"JSON-RPC error should pass through transport as Ok(Value)"
);
let json = result.unwrap();
assert_eq!(json["error"]["code"], -32602);
assert_eq!(json["error"]["message"], "Invalid params");
assert_eq!(json["error"]["data"]["detail"], "missing field");
}
#[tokio::test]
async fn test_sse_post_with_server_request_routed_correctly() {
let sse_data = b"data: {\"jsonrpc\":\"2.0\",\"id\":\"srv-99\",\"method\":\"sampling/createMessage\",\"params\":{}}\ndata: {\"jsonrpc\":\"2.0\",\"id\":\"req_0\",\"result\":{\"tools\":[]}}\n\n";
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(sse_data.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_sse_stream(stream).await;
assert!(result.is_ok());
let json = result.unwrap();
assert_eq!(
json["id"], "req_0",
"Should return the final response frame"
);
let events = transport.queued_events.lock();
assert_eq!(events.len(), 1, "Should have exactly one queued event");
match &events[0] {
ServerEvent::Request(val) => {
assert_eq!(val["id"], "srv-99");
assert_eq!(val["method"], "sampling/createMessage");
}
other => panic!("Expected ServerEvent::Request, got {:?}", other),
}
}
#[tokio::test]
async fn test_sse_post_with_notification_routed_correctly() {
let sse_data = b"data: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\",\"params\":{\"progress\":50}}\ndata: {\"jsonrpc\":\"2.0\",\"id\":\"req_0\",\"result\":{\"tools\":[]}}\n\n";
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(sse_data.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_sse_stream(stream).await;
assert!(result.is_ok());
let events = transport.queued_events.lock();
assert_eq!(events.len(), 1);
match &events[0] {
ServerEvent::Notification(val) => {
assert_eq!(val["method"], "notifications/progress");
}
other => panic!("Expected ServerEvent::Notification, got {:?}", other),
}
}
#[tokio::test]
async fn test_sse_data_field_without_space_after_colon() {
let sse_body: &[u8] = b"data:{\"jsonrpc\":\"2.0\",\"id\":\"req_0\",\"result\":{}}\n\n";
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(sse_body.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_sse_stream(stream).await;
assert!(
result.is_ok(),
"SSE data field without space should be parseable: {:?}",
result.err()
);
}
#[tokio::test]
async fn test_sse_colon_comment_line_is_ignored() {
let sse_data = b": this is a comment, ignore me\ndata: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\",\"params\":{\"progress\":50}}\n: another comment\ndata: {\"jsonrpc\":\"2.0\",\"id\":\"req_0\",\"result\":{\"tools\":[]}}\n\n";
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(sse_data.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_sse_stream(stream).await;
assert!(
result.is_ok(),
"comment lines must not be treated as malformed input: {:?}",
result.err()
);
assert_eq!(
result.unwrap()["id"],
"req_0",
"final result frame must still be parsed past the comment lines"
);
let events = transport.queued_events.lock();
assert_eq!(
events.len(),
1,
"exactly one notification event, comment lines produce none"
);
match &events[0] {
ServerEvent::Notification(val) => {
assert_eq!(val["method"], "notifications/progress");
}
other => panic!("Expected ServerEvent::Notification, got {:?}", other),
}
}
#[tokio::test]
async fn test_sse_post_no_duplicate_event_delivery() {
use tokio::io::AsyncBufReadExt;
let (tx, mut rx) = mpsc::unbounded_channel();
let queued = Arc::new(parking_lot::Mutex::new(Vec::new()));
let stats = Arc::new(parking_lot::Mutex::new(TransportStatistics::default()));
let sse_data = b"data: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\",\"params\":{}}\ndata: {\"jsonrpc\":\"2.0\",\"id\":\"req_0\",\"result\":{}}\n\n";
let cursor = std::io::Cursor::new(sse_data.to_vec());
let mut lines = tokio::io::BufReader::new(cursor).lines();
let sender = Some(tx);
let _ = parse_sse_lines(&mut lines, sender, &queued, &stats).await;
let channel_event = rx.try_recv().ok();
let queued_events = queued.lock().clone();
let total = (if channel_event.is_some() { 1 } else { 0 }) + queued_events.len();
assert_eq!(
total,
1,
"Event should be delivered exactly once, not duplicated. Channel: {}, Queued: {}",
channel_event.is_some(),
queued_events.len()
);
}
#[tokio::test]
async fn test_byte_stream_request_vs_notification_discrimination() {
let frames = br#"{"jsonrpc":"2.0","id":"srv-1","method":"sampling/createMessage","params":{}}{"jsonrpc":"2.0","method":"notifications/progress","params":{"progress":50}}{"jsonrpc":"2.0","id":"req_0","result":{"tools":[]}}"#;
let stream = futures::stream::iter(vec![Ok::<_, std::io::Error>(frames.to_vec())]);
let transport = HttpTransport::new("http://localhost:9999/mcp").unwrap();
let result = transport.test_handle_byte_stream(stream).await;
assert!(result.is_ok());
let events = transport.queued_events.lock();
assert_eq!(
events.len(),
2,
"Should have both request and notification events"
);
match &events[0] {
ServerEvent::Request(val) => {
assert_eq!(val["id"], "srv-1");
assert_eq!(val["method"], "sampling/createMessage");
}
other => panic!(
"Expected ServerEvent::Request as first event, got {:?}",
other
),
}
match &events[1] {
ServerEvent::Notification(val) => {
assert_eq!(val["method"], "notifications/progress");
}
other => panic!(
"Expected ServerEvent::Notification as second event, got {:?}",
other
),
}
}
}