use std::fmt;
use std::net::{IpAddr, SocketAddr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use futures::StreamExt;
use reqwest::header::{HeaderName, HeaderValue, ACCEPT, CONTENT_TYPE};
use reqwest::{Client, Url};
use serde_json::{json, Value};
use sha2::{Digest, Sha256};
use crate::{McpEndpoint, McpError, McpTool, McpToolResult, McpTransport};
const MAX_RESPONSE_BYTES: usize = 2 * 1024 * 1024;
#[derive(Clone, PartialEq, Eq)]
pub struct McpCredential {
pub header_name: String,
pub header_value: String,
}
impl fmt::Debug for McpCredential {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("McpCredential")
.field("header_name", &self.header_name)
.field("header_value", &"[REDACTED]")
.finish()
}
}
#[async_trait]
pub trait McpCredentialProvider: Send + Sync {
async fn resolve(&self, reference: &str) -> Result<McpCredential, McpError>;
}
pub struct StreamableHttpTransport<C> {
credentials: C,
request_id: AtomicU64,
}
impl<C> StreamableHttpTransport<C> {
pub fn new(credentials: C) -> Self {
Self {
credentials,
request_id: AtomicU64::new(1),
}
}
}
#[async_trait]
impl<C: McpCredentialProvider> McpTransport for StreamableHttpTransport<C> {
async fn list_tools(&self, endpoint: &McpEndpoint) -> Result<Vec<McpTool>, McpError> {
let result = self
.request(endpoint, "tools/list", json!({}), None)
.await?;
serde_json::from_value(
result
.get("tools")
.cloned()
.ok_or_else(|| McpError::Unavailable("tools/list omitted tools".into()))?,
)
.map_err(|_| McpError::Unavailable("tools/list returned an invalid catalog".into()))
}
async fn call_tool(
&self,
endpoint: &McpEndpoint,
context: &crate::McpCallContext,
name: &str,
arguments: Value,
) -> Result<McpToolResult, McpError> {
let idempotency_key = scoped_idempotency_key(endpoint, context);
let result = self
.request(
endpoint,
"tools/call",
json!({"name":name,"arguments":arguments}),
Some(&idempotency_key),
)
.await?;
normalize_tool_result(result)
}
}
fn normalize_tool_result(result: Value) -> Result<McpToolResult, McpError> {
let is_error = result
.get("isError")
.and_then(Value::as_bool)
.unwrap_or(false);
let content = if is_error {
result.get("content")
} else {
result
.get("structuredContent")
.or_else(|| result.get("content"))
}
.cloned()
.ok_or_else(|| McpError::Unavailable("tools/call omitted content".into()))?;
Ok(McpToolResult { content, is_error })
}
fn scoped_idempotency_key(endpoint: &McpEndpoint, context: &crate::McpCallContext) -> String {
let mut digest = Sha256::new();
for part in [
endpoint.id.as_str(),
context.tenant_id.as_str(),
context.subject_id.as_str(),
context.session_id.as_str(),
context.run_id.as_str(),
context.call_id.as_str(),
] {
digest.update((part.len() as u64).to_be_bytes());
digest.update(part.as_bytes());
}
digest.update(context.source_event_seq.to_be_bytes());
format!("af-mcp-{:x}", digest.finalize())
}
impl<C: McpCredentialProvider> StreamableHttpTransport<C> {
async fn request(
&self,
endpoint: &McpEndpoint,
method: &str,
params: Value,
idempotency_key: Option<&str>,
) -> Result<Value, McpError> {
let url = endpoint.validate()?;
let client = pinned_client(&url, endpoint.timeout_ms).await?;
let credential = match endpoint.credential_ref.as_deref() {
Some(reference) => Some(self.credentials.resolve(reference).await?),
None => None,
};
let initialize_id = self.next_id();
let initialize = send(
&client,
&url,
credential.as_ref(),
None,
None,
json!({
"jsonrpc":"2.0",
"id":initialize_id,
"method":"initialize",
"params":{
"protocolVersion":"2025-03-26",
"capabilities":{},
"clientInfo":{"name":"agent-factory","version":env!("CARGO_PKG_VERSION")}
}
}),
)
.await?;
rpc_result(&initialize.body, initialize_id)?;
let session = initialize.session_id.as_deref();
send(
&client,
&url,
credential.as_ref(),
session,
None,
json!({"jsonrpc":"2.0","method":"notifications/initialized"}),
)
.await?;
let id = self.next_id();
let response = send(
&client,
&url,
credential.as_ref(),
session,
idempotency_key,
json!({"jsonrpc":"2.0","id":id,"method":method,"params":params}),
)
.await?;
rpc_result(&response.body, id)
}
fn next_id(&self) -> u64 {
self.request_id.fetch_add(1, Ordering::Relaxed)
}
}
#[derive(Debug)]
struct HttpResponse {
body: Value,
session_id: Option<String>,
}
async fn send(
client: &Client,
url: &Url,
credential: Option<&McpCredential>,
session_id: Option<&str>,
idempotency_key: Option<&str>,
body: Value,
) -> Result<HttpResponse, McpError> {
let mut request = client
.post(url.clone())
.header(ACCEPT, "application/json, text/event-stream")
.header(CONTENT_TYPE, "application/json")
.json(&body);
if let Some(session_id) = session_id {
request = request.header("Mcp-Session-Id", session_id);
}
if let Some(idempotency_key) = idempotency_key {
request = request.header("Idempotency-Key", idempotency_key);
}
if let Some(credential) = credential {
let name = HeaderName::from_bytes(credential.header_name.as_bytes())
.map_err(|_| McpError::Rejected("credential header name is invalid".into()))?;
if name != reqwest::header::AUTHORIZATION && !name.as_str().starts_with("x-") {
return Err(McpError::Rejected(
"credential headers must be Authorization or X-*".into(),
));
}
let value = HeaderValue::from_str(&credential.header_value)
.map_err(|_| McpError::Rejected("credential header value is invalid".into()))?;
request = request.header(name, value);
}
let response = request
.send()
.await
.map_err(|error| McpError::Unavailable(format!("HTTP request failed: {error}")))?;
let status = response.status();
let session_id = response
.headers()
.get("Mcp-Session-Id")
.and_then(|value| value.to_str().ok())
.map(str::to_string);
if status == reqwest::StatusCode::ACCEPTED || status == reqwest::StatusCode::NO_CONTENT {
return Ok(HttpResponse {
body: Value::Null,
session_id,
});
}
if !status.is_success() {
return Err(McpError::Unavailable(format!(
"MCP endpoint returned HTTP {status}"
)));
}
let content_type = response
.headers()
.get(CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string();
if response
.content_length()
.is_some_and(|length| length > MAX_RESPONSE_BYTES as u64)
{
return Err(McpError::Rejected("MCP response exceeds 2 MiB".into()));
}
let mut bytes = Vec::new();
let mut chunks = response.bytes_stream();
while let Some(chunk) = chunks.next().await {
let chunk = chunk
.map_err(|error| McpError::Unavailable(format!("response read failed: {error}")))?;
append_limited(&mut bytes, &chunk)?;
}
let body = if content_type.starts_with("text/event-stream") {
parse_sse(&bytes)?
} else if content_type.starts_with("application/json") {
serde_json::from_slice(&bytes)
.map_err(|_| McpError::Unavailable("MCP returned invalid JSON".into()))?
} else {
return Err(McpError::Unavailable(
"MCP returned an unsupported content type".into(),
));
};
Ok(HttpResponse { body, session_id })
}
fn append_limited(target: &mut Vec<u8>, chunk: &[u8]) -> Result<(), McpError> {
if target.len().saturating_add(chunk.len()) > MAX_RESPONSE_BYTES {
return Err(McpError::Rejected("MCP response exceeds 2 MiB".into()));
}
target.extend_from_slice(chunk);
Ok(())
}
fn rpc_result(body: &Value, id: u64) -> Result<Value, McpError> {
if body.get("id").and_then(Value::as_u64) != Some(id) {
return Err(McpError::Unavailable("MCP response id mismatch".into()));
}
if body.get("error").is_some() {
return Err(McpError::Unavailable(
"MCP returned a JSON-RPC error".into(),
));
}
body.get("result")
.cloned()
.ok_or_else(|| McpError::Unavailable("MCP response omitted result".into()))
}
fn parse_sse(bytes: &[u8]) -> Result<Value, McpError> {
let text = std::str::from_utf8(bytes)
.map_err(|_| McpError::Unavailable("MCP returned invalid SSE text".into()))?;
text.lines()
.filter_map(|line| line.strip_prefix("data:"))
.map(str::trim)
.find(|line| !line.is_empty())
.ok_or_else(|| McpError::Unavailable("MCP SSE response omitted data".into()))
.and_then(|data| {
serde_json::from_str(data)
.map_err(|_| McpError::Unavailable("MCP returned invalid SSE JSON".into()))
})
}
async fn pinned_client(url: &Url, timeout_ms: u64) -> Result<Client, McpError> {
let host = url
.host_str()
.ok_or_else(|| McpError::Rejected("missing host".into()))?;
let port = url.port_or_known_default().unwrap_or(443);
let addresses = tokio::net::lookup_host((host, port))
.await
.map_err(|_| McpError::Unavailable("MCP DNS resolution failed".into()))?
.collect::<Vec<SocketAddr>>();
if addresses.is_empty() || addresses.iter().any(|address| !is_public(address.ip())) {
return Err(McpError::Rejected(
"MCP DNS resolved to a non-public address".into(),
));
}
Client::builder()
.redirect(reqwest::redirect::Policy::none())
.resolve(host, addresses[0])
.timeout(Duration::from_millis(timeout_ms))
.build()
.map_err(|_| McpError::Unavailable("MCP HTTP client initialization failed".into()))
}
fn is_public(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => {
let octets = ip.octets();
!ip.is_private()
&& !ip.is_loopback()
&& !ip.is_link_local()
&& !ip.is_broadcast()
&& !ip.is_documentation()
&& !ip.is_unspecified()
&& !ip.is_multicast()
&& !(octets[0] == 100 && (64..=127).contains(&octets[1]))
}
IpAddr::V6(ip) => {
if let Some(mapped) = ip.to_ipv4_mapped() {
return is_public(IpAddr::V4(mapped));
}
!ip.is_loopback()
&& !ip.is_unspecified()
&& !ip.is_multicast()
&& !(ip.segments()[0] & 0xfe00 == 0xfc00)
&& !(ip.segments()[0] & 0xffc0 == 0xfe80)
&& !(ip.segments()[0] == 0x2001 && ip.segments()[1] == 0x0db8)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
async fn mock_response(response: &'static str, delay: Duration) -> Url {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = vec![0; 4096];
let _ = socket.read(&mut request).await;
tokio::time::sleep(delay).await;
let _ = socket.write_all(response.as_bytes()).await;
});
Url::parse(&format!("http://{address}/mcp")).unwrap()
}
#[test]
fn parses_sse_and_rejects_private_networks() {
assert_eq!(
parse_sse(b"event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\n\n")
.unwrap()["id"],
1
);
assert!(!is_public("127.0.0.1".parse().unwrap()));
assert!(!is_public("10.0.0.1".parse().unwrap()));
assert!(is_public("1.1.1.1".parse().unwrap()));
let tool: McpTool = serde_json::from_value(json!({
"name":"echo",
"description":"",
"inputSchema":{"type":"object"}
}))
.unwrap();
assert_eq!(tool.input_schema["type"], "object");
assert_eq!(tool.output_schema, Value::Null);
assert_eq!(
normalize_tool_result(json!({
"content":[{"type":"text","text":"fallback"}],
"structuredContent":{"value":1}
}))
.unwrap()
.content,
json!({"value":1})
);
assert_eq!(
normalize_tool_result(json!({"content":[{"type":"text","text":"fallback"}]}))
.unwrap()
.content,
json!([{"type":"text","text":"fallback"}])
);
}
#[test]
fn response_limit_stops_before_appending_oversized_chunk() {
let mut body = vec![0; MAX_RESPONSE_BYTES - 1];
append_limited(&mut body, &[1]).unwrap();
assert_eq!(body.len(), MAX_RESPONSE_BYTES);
assert!(append_limited(&mut body, &[2]).is_err());
assert_eq!(body.len(), MAX_RESPONSE_BYTES);
}
#[test]
fn credential_debug_is_redacted() {
let credential = McpCredential {
header_name: "authorization".into(),
header_value: "Bearer top-secret".into(),
};
let debug = format!("{credential:?}");
assert!(debug.contains("authorization"));
assert!(!debug.contains("top-secret"));
}
#[test]
fn idempotency_key_is_stable_and_scoped_to_the_durable_call() {
let endpoint = McpEndpoint {
id: "search".into(),
url: "https://mcp.example.com".into(),
namespace: "reference".into(),
allowed_hosts: ["mcp.example.com".into()].into_iter().collect(),
allowed_tools: ["query".into()].into_iter().collect(),
credential_ref: None,
timeout_ms: 1_000,
failure_threshold: 1,
recovery_ms: 1_000,
};
let context = crate::McpCallContext {
tenant_id: "tenant-a".parse().unwrap(),
subject_id: "subject".parse().unwrap(),
session_id: "session".parse().unwrap(),
run_id: "run".parse().unwrap(),
call_id: "call".parse().unwrap(),
source_event_seq: 7,
request_id: "request".parse().unwrap(),
};
let key = scoped_idempotency_key(&endpoint, &context);
assert_eq!(key, scoped_idempotency_key(&endpoint, &context));
let mut other_tenant = context;
other_tenant.tenant_id = "tenant-b".parse().unwrap();
assert_ne!(key, scoped_idempotency_key(&endpoint, &other_tenant));
assert!(!key.contains("tenant-a"));
}
#[tokio::test]
async fn http_sender_classifies_success_rate_limit_server_error_and_timeout_without_secrets() {
let ok = mock_response("HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 36\r\n\r\n{\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}", Duration::ZERO).await;
let client = Client::builder()
.timeout(Duration::from_secs(1))
.build()
.unwrap();
assert_eq!(
send(&client, &ok, None, None, None, json!({}))
.await
.unwrap()
.body["id"],
1
);
for status in ["429 Too Many Requests", "503 Service Unavailable"] {
let response = Box::leak(
format!("HTTP/1.1 {status}\r\nContent-Length: 0\r\n\r\n").into_boxed_str(),
);
let url = mock_response(response, Duration::ZERO).await;
let credential = McpCredential {
header_name: "authorization".into(),
header_value: "Bearer top-secret".into(),
};
let error = send(&client, &url, Some(&credential), None, None, json!({}))
.await
.unwrap_err();
assert!(matches!(error, McpError::Unavailable(_)));
assert!(!error.to_string().contains("top-secret"));
}
let slow = mock_response(
"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n",
Duration::from_millis(100),
)
.await;
let impatient = Client::builder()
.timeout(Duration::from_millis(5))
.build()
.unwrap();
assert!(matches!(
send(&impatient, &slow, None, None, None, json!({})).await,
Err(McpError::Unavailable(_))
));
}
}