#[cfg(test)]
mod tests;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use serde::de::DeserializeOwned;
use serde_json::Value as JsonValue;
use tokio::sync::{Mutex, oneshot};
use tokio::task::JoinHandle;
use url::Url;
use super::config::{McpSseConfigError, McpSseLimits, McpSseServerConfig};
use super::endpoint::{EndpointError, resolve_endpoint};
use super::wire::{SseParser, SseWireError};
use crate::mcp::protocol::*;
const PROTOCOL_VERSION: &str = "2024-11-05";
const MAX_DIAGNOSTIC_BODY_BYTES: usize = 8 * 1024;
#[derive(Debug, thiserror::Error)]
pub enum McpSseError {
#[error("invalid MCP SSE configuration: {0}")]
Config(#[from] McpSseConfigError),
#[error("invalid MCP SSE endpoint: {0}")]
Endpoint(#[from] EndpointError),
#[error("failed to reach the MCP SSE server: {0}")]
Transport(String),
#[error("MCP SSE server answered the {method} request with HTTP {status}")]
HttpStatus {
method: &'static str,
status: reqwest::StatusCode,
},
#[error(
"MCP SSE server answered with a redirect, which is not followed because it would send \
credentials to an unvalidated origin"
)]
RedirectRefused,
#[error(
"MCP SSE server answered with content type '{content_type}', expected text/event-stream"
)]
UnexpectedContentType { content_type: String },
#[error("MCP SSE stream framing error: {0}")]
Wire(#[from] SseWireError),
#[error("MCP SSE endpoint event exceeded the {limit} byte limit")]
EndpointTooLarge { limit: usize },
#[error("MCP SSE server kept paginating tools/list past {limit} pages")]
TooManyToolPages { limit: usize },
#[error("MCP SSE server returned JSON-RPC error: {0}")]
JsonRpc(JsonRpcError),
#[error("failed to parse the MCP SSE response: {0}")]
ParseError(String),
#[error("timed out after {0:?} waiting for the MCP SSE server")]
Timeout(Duration),
#[error("the MCP SSE stream closed before the request completed")]
StreamClosed,
#[error(
"the MCP SSE server accepted the '{method}' request but never answered it; \
the call may have executed and must not be retried automatically"
)]
RequestIndeterminate { method: String },
#[error("the MCP SSE client is shut down")]
Shutdown,
}
type PendingReply = Result<JsonValue, McpSseError>;
#[derive(Default)]
struct Pending {
waiters: HashMap<u64, PendingWaiter>,
closed: bool,
}
struct PendingWaiter {
reply: oneshot::Sender<PendingReply>,
method: String,
}
pub struct McpSseClient {
http: reqwest::Client,
endpoint: Url,
headers: HeaderMap,
limits: McpSseLimits,
next_id: AtomicU64,
pending: Arc<Mutex<Pending>>,
reader: JoinHandle<()>,
server_info: Option<McpServerInfo>,
tools: Vec<McpToolDefinition>,
server_name: String,
stream_url: Url,
}
impl McpSseClient {
pub async fn connect(config: &McpSseServerConfig) -> Result<Self, McpSseError> {
let stream_url = config.validate()?;
let headers = build_headers(config)?;
let http = build_http_client(&config.limits)?;
let response = tokio::time::timeout(
config.limits.connect_timeout,
http.get(stream_url.clone())
.header(reqwest::header::ACCEPT, "text/event-stream")
.headers(headers.clone())
.send(),
)
.await
.map_err(|_| McpSseError::Timeout(config.limits.connect_timeout))?
.map_err(transport_error)?;
check_stream_response(&response)?;
let pending: Arc<Mutex<Pending>> = Arc::new(Mutex::new(Pending::default()));
let (endpoint_tx, endpoint_rx) = oneshot::channel();
let reader = tokio::spawn(read_stream(
response,
Arc::clone(&pending),
endpoint_tx,
config.limits.clone(),
));
let endpoint = match tokio::time::timeout(config.limits.connect_timeout, endpoint_rx).await
{
Ok(Ok(Ok(raw))) => match resolve_endpoint(&stream_url, &raw) {
Ok(endpoint) => endpoint,
Err(error) => {
reader.abort();
return Err(error.into());
}
},
Ok(Ok(Err(error))) => {
reader.abort();
return Err(error);
}
Ok(Err(_)) => {
reader.abort();
return Err(McpSseError::StreamClosed);
}
Err(_) => {
reader.abort();
return Err(McpSseError::Timeout(config.limits.connect_timeout));
}
};
let mut client = Self {
http,
endpoint,
headers,
limits: config.limits.clone(),
next_id: AtomicU64::new(1),
pending,
reader,
server_info: None,
tools: Vec::new(),
server_name: config.name.clone(),
stream_url,
};
client.initialize().await?;
client.discover_tools().await?;
Ok(client)
}
pub fn server_name(&self) -> &str {
&self.server_name
}
pub fn stream_url(&self) -> &Url {
&self.stream_url
}
pub fn server_info(&self) -> Option<&McpServerInfo> {
self.server_info.as_ref()
}
pub fn tools(&self) -> &[McpToolDefinition] {
&self.tools
}
pub async fn call_tool(
&self,
tool_name: &str,
arguments: Option<JsonValue>,
) -> Result<McpToolCallResult, McpSseError> {
let params = McpToolCallParams {
name: tool_name.to_string(),
arguments,
};
self.request("tools/call", Some(params), self.limits.call_tool_timeout)
.await
}
pub async fn shutdown(&self) {
self.reader.abort();
let mut pending = self.pending.lock().await;
pending.closed = true;
drain_pending(&mut pending);
}
async fn request<P: serde::Serialize, R: DeserializeOwned>(
&self,
method: &'static str,
params: Option<P>,
timeout: Duration,
) -> Result<R, McpSseError> {
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
let params = params
.map(|params| serde_json::to_value(params))
.transpose()
.map_err(|error| McpSseError::ParseError(error.to_string()))?
.filter(|params| !params.is_null());
let request = JsonRpcRequest::new(id, method, params);
let (reply_tx, reply_rx) = oneshot::channel();
{
let mut pending = self.pending.lock().await;
if pending.closed {
return Err(McpSseError::StreamClosed);
}
pending.waiters.insert(
id,
PendingWaiter {
reply: reply_tx,
method: method.to_string(),
},
);
}
if let Err(error) = self.post(&request).await {
self.forget(id).await;
return Err(error);
}
let result = match tokio::time::timeout(timeout, reply_rx).await {
Ok(Ok(result)) => result?,
Ok(Err(_)) => return Err(indeterminate(method)),
Err(_) => {
self.forget(id).await;
return Err(if method == "tools/call" {
indeterminate(method)
} else {
McpSseError::Timeout(timeout)
});
}
};
serde_json::from_value(result)
.map_err(|error| McpSseError::ParseError(format!("deserialize response: {error}")))
}
async fn forget(&self, id: u64) {
self.pending.lock().await.waiters.remove(&id);
}
async fn notify(&self, method: &str) -> Result<(), McpSseError> {
let notification = serde_json::json!({"jsonrpc": "2.0", "method": method});
self.post(¬ification).await
}
async fn post<T: serde::Serialize>(&self, message: &T) -> Result<(), McpSseError> {
let response = self
.http
.post(self.endpoint.clone())
.headers(self.headers.clone())
.json(message)
.send()
.await
.map_err(transport_error)?;
let status = response.status();
if status.is_redirection() {
return Err(McpSseError::RedirectRefused);
}
if !status.is_success() {
drain_bounded(response).await;
return Err(McpSseError::HttpStatus {
method: "POST",
status,
});
}
drain_bounded(response).await;
Ok(())
}
async fn initialize(&mut self) -> Result<(), McpSseError> {
let params = McpInitializeParams {
protocol_version: PROTOCOL_VERSION.to_string(),
capabilities: serde_json::json!({}),
client_info: McpClientInfo {
name: "mentra".to_string(),
version: env!("CARGO_PKG_VERSION").to_string(),
},
};
let result: McpInitializeResult = self
.request("initialize", Some(params), self.limits.initialize_timeout)
.await?;
self.server_info = Some(result.server_info);
self.notify("notifications/initialized").await
}
async fn discover_tools(&mut self) -> Result<(), McpSseError> {
let mut tools = Vec::new();
let mut cursor: Option<String> = None;
let mut pages = 0_usize;
loop {
let params = McpListToolsParams {
cursor: cursor.clone(),
};
let page: McpListToolsResult = self
.request("tools/list", Some(params), self.limits.list_tools_timeout)
.await?;
tools.extend(page.tools);
pages += 1;
if pages >= self.limits.max_tool_pages {
return Err(McpSseError::TooManyToolPages {
limit: self.limits.max_tool_pages,
});
}
match page.next_cursor {
Some(next) if !next.is_empty() => cursor = Some(next),
_ => break,
}
}
self.tools = tools;
Ok(())
}
}
impl Drop for McpSseClient {
fn drop(&mut self) {
self.reader.abort();
}
}
impl std::fmt::Debug for McpSseClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpSseClient")
.field("server_name", &self.server_name)
.field("stream_url", &self.stream_url.as_str())
.field("tools", &self.tools.len())
.finish_non_exhaustive()
}
}
async fn read_stream(
response: reqwest::Response,
pending: Arc<Mutex<Pending>>,
endpoint_tx: oneshot::Sender<Result<String, McpSseError>>,
limits: McpSseLimits,
) {
let mut parser = SseParser::new(limits.max_event_bytes);
let mut body = response.bytes_stream();
let mut endpoint_tx = Some(endpoint_tx);
loop {
let next = tokio::time::timeout(limits.stream_idle_timeout, body.next()).await;
let chunk = match next {
Ok(Some(Ok(chunk))) => chunk,
Ok(Some(Err(_))) | Ok(None) => break,
Err(_) => break,
};
let events = match parser.feed(&chunk) {
Ok(events) => events,
Err(error) => {
notify_endpoint(&mut endpoint_tx, Err(McpSseError::Wire(error)));
break;
}
};
for event in events {
match event.event.as_str() {
"endpoint" => {
if event.data.len() > limits.max_endpoint_bytes {
notify_endpoint(
&mut endpoint_tx,
Err(McpSseError::EndpointTooLarge {
limit: limits.max_endpoint_bytes,
}),
);
break;
}
notify_endpoint(&mut endpoint_tx, Ok(event.data));
}
"message" => deliver_message(&pending, &event.data).await,
_ => {}
}
}
}
notify_endpoint(&mut endpoint_tx, Err(McpSseError::StreamClosed));
let mut pending = pending.lock().await;
pending.closed = true;
drain_pending(&mut pending);
}
async fn deliver_message(pending: &Arc<Mutex<Pending>>, data: &str) {
let Ok(response) = serde_json::from_str::<JsonRpcResponse>(data) else {
return;
};
if response.result.is_none() && response.error.is_none() {
return;
}
let JsonRpcId::Number(id) = response.id else {
return;
};
let Some(waiter) = pending.lock().await.waiters.remove(&id) else {
return;
};
let reply = match response.error {
Some(error) => Err(McpSseError::JsonRpc(error)),
None => Ok(response.result.unwrap_or(JsonValue::Null)),
};
let _ = waiter.reply.send(reply);
}
fn notify_endpoint(
endpoint_tx: &mut Option<oneshot::Sender<Result<String, McpSseError>>>,
outcome: Result<String, McpSseError>,
) {
if let Some(tx) = endpoint_tx.take() {
let _ = tx.send(outcome);
}
}
fn drain_pending(pending: &mut Pending) {
for (_, waiter) in pending.waiters.drain() {
let _ = waiter.reply.send(Err(indeterminate(&waiter.method)));
}
}
fn indeterminate(method: &str) -> McpSseError {
if method == "tools/call" {
McpSseError::RequestIndeterminate {
method: method.to_string(),
}
} else {
McpSseError::StreamClosed
}
}
fn build_http_client(limits: &McpSseLimits) -> Result<reqwest::Client, McpSseError> {
reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(limits.connect_timeout)
.build()
.map_err(|error| McpSseError::Transport(error.to_string()))
}
fn build_headers(config: &McpSseServerConfig) -> Result<HeaderMap, McpSseError> {
let mut headers = HeaderMap::new();
for (name, value) in &config.headers {
let name = HeaderName::try_from(name.as_str()).map_err(|_| {
McpSseConfigError::InvalidHeaderName {
name: name.to_string(),
}
})?;
let mut value = HeaderValue::try_from(value.expose_secret()).map_err(|_| {
McpSseConfigError::InvalidHeaderValue {
name: name.to_string(),
}
})?;
value.set_sensitive(true);
headers.insert(name, value);
}
Ok(headers)
}
fn check_stream_response(response: &reqwest::Response) -> Result<(), McpSseError> {
let status = response.status();
if status.is_redirection() {
return Err(McpSseError::RedirectRefused);
}
if !status.is_success() {
return Err(McpSseError::HttpStatus {
method: "GET",
status,
});
}
let content_type = response
.headers()
.get(reqwest::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or_default();
if !content_type
.trim_start()
.to_ascii_lowercase()
.starts_with("text/event-stream")
{
return Err(McpSseError::UnexpectedContentType {
content_type: content_type.to_string(),
});
}
Ok(())
}
async fn drain_bounded(response: reqwest::Response) {
let mut body = response.bytes_stream();
let mut seen = 0_usize;
while let Some(Ok(chunk)) = body.next().await {
seen += chunk.len();
if seen >= MAX_DIAGNOSTIC_BODY_BYTES {
break;
}
}
}
fn transport_error(error: reqwest::Error) -> McpSseError {
let reason = if error.is_timeout() {
"timed out"
} else if error.is_connect() {
"could not connect"
} else if error.is_request() {
"the request could not be sent"
} else {
"the connection failed"
};
McpSseError::Transport(reason.to_string())
}