1use std::fmt;
2use std::net::{IpAddr, SocketAddr};
3use std::sync::atomic::{AtomicU64, Ordering};
4use std::time::Duration;
5
6use async_trait::async_trait;
7use futures::StreamExt;
8use reqwest::header::{HeaderName, HeaderValue, ACCEPT, CONTENT_TYPE};
9use reqwest::{Client, Url};
10use serde_json::{json, Value};
11use sha2::{Digest, Sha256};
12
13use crate::{McpEndpoint, McpError, McpTool, McpToolResult, McpTransport};
14
15const MAX_RESPONSE_BYTES: usize = 2 * 1024 * 1024;
16
17#[derive(Clone, PartialEq, Eq)]
19pub struct McpCredential {
20 pub header_name: String,
22 pub header_value: String,
24}
25
26impl fmt::Debug for McpCredential {
27 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
28 formatter
29 .debug_struct("McpCredential")
30 .field("header_name", &self.header_name)
31 .field("header_value", &"[REDACTED]")
32 .finish()
33 }
34}
35
36#[async_trait]
38pub trait McpCredentialProvider: Send + Sync {
39 async fn resolve(&self, reference: &str) -> Result<McpCredential, McpError>;
41}
42
43pub struct StreamableHttpTransport<C> {
45 credentials: C,
46 request_id: AtomicU64,
47}
48
49impl<C> StreamableHttpTransport<C> {
50 pub fn new(credentials: C) -> Self {
52 Self {
53 credentials,
54 request_id: AtomicU64::new(1),
55 }
56 }
57}
58
59#[async_trait]
60impl<C: McpCredentialProvider> McpTransport for StreamableHttpTransport<C> {
61 async fn list_tools(&self, endpoint: &McpEndpoint) -> Result<Vec<McpTool>, McpError> {
62 let result = self
63 .request(endpoint, "tools/list", json!({}), None)
64 .await?;
65 serde_json::from_value(
66 result
67 .get("tools")
68 .cloned()
69 .ok_or_else(|| McpError::Unavailable("tools/list omitted tools".into()))?,
70 )
71 .map_err(|_| McpError::Unavailable("tools/list returned an invalid catalog".into()))
72 }
73
74 async fn call_tool(
75 &self,
76 endpoint: &McpEndpoint,
77 context: &crate::McpCallContext,
78 name: &str,
79 arguments: Value,
80 ) -> Result<McpToolResult, McpError> {
81 let idempotency_key = scoped_idempotency_key(endpoint, context);
82 let result = self
83 .request(
84 endpoint,
85 "tools/call",
86 json!({"name":name,"arguments":arguments}),
87 Some(&idempotency_key),
88 )
89 .await?;
90 normalize_tool_result(result)
91 }
92}
93
94fn normalize_tool_result(result: Value) -> Result<McpToolResult, McpError> {
95 let is_error = result
96 .get("isError")
97 .and_then(Value::as_bool)
98 .unwrap_or(false);
99 let content = if is_error {
100 result.get("content")
101 } else {
102 result
103 .get("structuredContent")
104 .or_else(|| result.get("content"))
105 }
106 .cloned()
107 .ok_or_else(|| McpError::Unavailable("tools/call omitted content".into()))?;
108 Ok(McpToolResult { content, is_error })
109}
110
111fn scoped_idempotency_key(endpoint: &McpEndpoint, context: &crate::McpCallContext) -> String {
112 let mut digest = Sha256::new();
113 for part in [
114 endpoint.id.as_str(),
115 context.tenant_id.as_str(),
116 context.subject_id.as_str(),
117 context.session_id.as_str(),
118 context.run_id.as_str(),
119 context.call_id.as_str(),
120 ] {
121 digest.update((part.len() as u64).to_be_bytes());
122 digest.update(part.as_bytes());
123 }
124 digest.update(context.source_event_seq.to_be_bytes());
125 format!("af-mcp-{:x}", digest.finalize())
126}
127
128impl<C: McpCredentialProvider> StreamableHttpTransport<C> {
129 async fn request(
130 &self,
131 endpoint: &McpEndpoint,
132 method: &str,
133 params: Value,
134 idempotency_key: Option<&str>,
135 ) -> Result<Value, McpError> {
136 let url = endpoint.validate()?;
137 let client = pinned_client(&url, endpoint.timeout_ms).await?;
138 let credential = match endpoint.credential_ref.as_deref() {
139 Some(reference) => Some(self.credentials.resolve(reference).await?),
140 None => None,
141 };
142 let initialize_id = self.next_id();
143 let initialize = send(
144 &client,
145 &url,
146 credential.as_ref(),
147 None,
148 None,
149 json!({
150 "jsonrpc":"2.0",
151 "id":initialize_id,
152 "method":"initialize",
153 "params":{
154 "protocolVersion":"2025-03-26",
155 "capabilities":{},
156 "clientInfo":{"name":"agent-factory","version":env!("CARGO_PKG_VERSION")}
157 }
158 }),
159 )
160 .await?;
161 rpc_result(&initialize.body, initialize_id)?;
162 let session = initialize.session_id.as_deref();
163 send(
164 &client,
165 &url,
166 credential.as_ref(),
167 session,
168 None,
169 json!({"jsonrpc":"2.0","method":"notifications/initialized"}),
170 )
171 .await?;
172 let id = self.next_id();
173 let response = send(
174 &client,
175 &url,
176 credential.as_ref(),
177 session,
178 idempotency_key,
179 json!({"jsonrpc":"2.0","id":id,"method":method,"params":params}),
180 )
181 .await?;
182 rpc_result(&response.body, id)
183 }
184
185 fn next_id(&self) -> u64 {
186 self.request_id.fetch_add(1, Ordering::Relaxed)
187 }
188}
189
190#[derive(Debug)]
191struct HttpResponse {
192 body: Value,
193 session_id: Option<String>,
194}
195
196async fn send(
197 client: &Client,
198 url: &Url,
199 credential: Option<&McpCredential>,
200 session_id: Option<&str>,
201 idempotency_key: Option<&str>,
202 body: Value,
203) -> Result<HttpResponse, McpError> {
204 let mut request = client
205 .post(url.clone())
206 .header(ACCEPT, "application/json, text/event-stream")
207 .header(CONTENT_TYPE, "application/json")
208 .json(&body);
209 if let Some(session_id) = session_id {
210 request = request.header("Mcp-Session-Id", session_id);
211 }
212 if let Some(idempotency_key) = idempotency_key {
213 request = request.header("Idempotency-Key", idempotency_key);
214 }
215 if let Some(credential) = credential {
216 let name = HeaderName::from_bytes(credential.header_name.as_bytes())
217 .map_err(|_| McpError::Rejected("credential header name is invalid".into()))?;
218 if name != reqwest::header::AUTHORIZATION && !name.as_str().starts_with("x-") {
219 return Err(McpError::Rejected(
220 "credential headers must be Authorization or X-*".into(),
221 ));
222 }
223 let value = HeaderValue::from_str(&credential.header_value)
224 .map_err(|_| McpError::Rejected("credential header value is invalid".into()))?;
225 request = request.header(name, value);
226 }
227 let response = request
228 .send()
229 .await
230 .map_err(|error| McpError::Unavailable(format!("HTTP request failed: {error}")))?;
231 let status = response.status();
232 let session_id = response
233 .headers()
234 .get("Mcp-Session-Id")
235 .and_then(|value| value.to_str().ok())
236 .map(str::to_string);
237 if status == reqwest::StatusCode::ACCEPTED || status == reqwest::StatusCode::NO_CONTENT {
238 return Ok(HttpResponse {
239 body: Value::Null,
240 session_id,
241 });
242 }
243 if !status.is_success() {
244 return Err(McpError::Unavailable(format!(
245 "MCP endpoint returned HTTP {status}"
246 )));
247 }
248 let content_type = response
249 .headers()
250 .get(CONTENT_TYPE)
251 .and_then(|value| value.to_str().ok())
252 .unwrap_or_default()
253 .to_string();
254 if response
255 .content_length()
256 .is_some_and(|length| length > MAX_RESPONSE_BYTES as u64)
257 {
258 return Err(McpError::Rejected("MCP response exceeds 2 MiB".into()));
259 }
260 let mut bytes = Vec::new();
261 let mut chunks = response.bytes_stream();
262 while let Some(chunk) = chunks.next().await {
263 let chunk = chunk
264 .map_err(|error| McpError::Unavailable(format!("response read failed: {error}")))?;
265 append_limited(&mut bytes, &chunk)?;
266 }
267 let body = if content_type.starts_with("text/event-stream") {
268 parse_sse(&bytes)?
269 } else if content_type.starts_with("application/json") {
270 serde_json::from_slice(&bytes)
271 .map_err(|_| McpError::Unavailable("MCP returned invalid JSON".into()))?
272 } else {
273 return Err(McpError::Unavailable(
274 "MCP returned an unsupported content type".into(),
275 ));
276 };
277 Ok(HttpResponse { body, session_id })
278}
279
280fn append_limited(target: &mut Vec<u8>, chunk: &[u8]) -> Result<(), McpError> {
281 if target.len().saturating_add(chunk.len()) > MAX_RESPONSE_BYTES {
282 return Err(McpError::Rejected("MCP response exceeds 2 MiB".into()));
283 }
284 target.extend_from_slice(chunk);
285 Ok(())
286}
287
288fn rpc_result(body: &Value, id: u64) -> Result<Value, McpError> {
289 if body.get("id").and_then(Value::as_u64) != Some(id) {
290 return Err(McpError::Unavailable("MCP response id mismatch".into()));
291 }
292 if body.get("error").is_some() {
293 return Err(McpError::Unavailable(
294 "MCP returned a JSON-RPC error".into(),
295 ));
296 }
297 body.get("result")
298 .cloned()
299 .ok_or_else(|| McpError::Unavailable("MCP response omitted result".into()))
300}
301
302fn parse_sse(bytes: &[u8]) -> Result<Value, McpError> {
303 let text = std::str::from_utf8(bytes)
304 .map_err(|_| McpError::Unavailable("MCP returned invalid SSE text".into()))?;
305 text.lines()
306 .filter_map(|line| line.strip_prefix("data:"))
307 .map(str::trim)
308 .find(|line| !line.is_empty())
309 .ok_or_else(|| McpError::Unavailable("MCP SSE response omitted data".into()))
310 .and_then(|data| {
311 serde_json::from_str(data)
312 .map_err(|_| McpError::Unavailable("MCP returned invalid SSE JSON".into()))
313 })
314}
315
316async fn pinned_client(url: &Url, timeout_ms: u64) -> Result<Client, McpError> {
317 let host = url
318 .host_str()
319 .ok_or_else(|| McpError::Rejected("missing host".into()))?;
320 let port = url.port_or_known_default().unwrap_or(443);
321 let addresses = tokio::net::lookup_host((host, port))
322 .await
323 .map_err(|_| McpError::Unavailable("MCP DNS resolution failed".into()))?
324 .collect::<Vec<SocketAddr>>();
325 if addresses.is_empty() || addresses.iter().any(|address| !is_public(address.ip())) {
326 return Err(McpError::Rejected(
327 "MCP DNS resolved to a non-public address".into(),
328 ));
329 }
330 Client::builder()
331 .redirect(reqwest::redirect::Policy::none())
332 .resolve(host, addresses[0])
333 .timeout(Duration::from_millis(timeout_ms))
334 .build()
335 .map_err(|_| McpError::Unavailable("MCP HTTP client initialization failed".into()))
336}
337
338fn is_public(ip: IpAddr) -> bool {
339 match ip {
340 IpAddr::V4(ip) => {
341 let octets = ip.octets();
342 !ip.is_private()
343 && !ip.is_loopback()
344 && !ip.is_link_local()
345 && !ip.is_broadcast()
346 && !ip.is_documentation()
347 && !ip.is_unspecified()
348 && !ip.is_multicast()
349 && !(octets[0] == 100 && (64..=127).contains(&octets[1]))
350 }
351 IpAddr::V6(ip) => {
352 if let Some(mapped) = ip.to_ipv4_mapped() {
353 return is_public(IpAddr::V4(mapped));
354 }
355 !ip.is_loopback()
356 && !ip.is_unspecified()
357 && !ip.is_multicast()
358 && !(ip.segments()[0] & 0xfe00 == 0xfc00)
359 && !(ip.segments()[0] & 0xffc0 == 0xfe80)
360 && !(ip.segments()[0] == 0x2001 && ip.segments()[1] == 0x0db8)
361 }
362 }
363}
364
365#[cfg(test)]
366mod tests {
367 use super::*;
368 use tokio::io::{AsyncReadExt, AsyncWriteExt};
369
370 async fn mock_response(response: &'static str, delay: Duration) -> Url {
371 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
372 let address = listener.local_addr().unwrap();
373 tokio::spawn(async move {
374 let (mut socket, _) = listener.accept().await.unwrap();
375 let mut request = vec![0; 4096];
376 let _ = socket.read(&mut request).await;
377 tokio::time::sleep(delay).await;
378 let _ = socket.write_all(response.as_bytes()).await;
379 });
380 Url::parse(&format!("http://{address}/mcp")).unwrap()
381 }
382
383 #[test]
384 fn parses_sse_and_rejects_private_networks() {
385 assert_eq!(
386 parse_sse(b"event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\n\n")
387 .unwrap()["id"],
388 1
389 );
390 assert!(!is_public("127.0.0.1".parse().unwrap()));
391 assert!(!is_public("10.0.0.1".parse().unwrap()));
392 assert!(is_public("1.1.1.1".parse().unwrap()));
393 let tool: McpTool = serde_json::from_value(json!({
394 "name":"echo",
395 "description":"",
396 "inputSchema":{"type":"object"}
397 }))
398 .unwrap();
399 assert_eq!(tool.input_schema["type"], "object");
400 assert_eq!(tool.output_schema, Value::Null);
401 assert_eq!(
402 normalize_tool_result(json!({
403 "content":[{"type":"text","text":"fallback"}],
404 "structuredContent":{"value":1}
405 }))
406 .unwrap()
407 .content,
408 json!({"value":1})
409 );
410 assert_eq!(
411 normalize_tool_result(json!({"content":[{"type":"text","text":"fallback"}]}))
412 .unwrap()
413 .content,
414 json!([{"type":"text","text":"fallback"}])
415 );
416 }
417
418 #[test]
419 fn response_limit_stops_before_appending_oversized_chunk() {
420 let mut body = vec![0; MAX_RESPONSE_BYTES - 1];
421 append_limited(&mut body, &[1]).unwrap();
422 assert_eq!(body.len(), MAX_RESPONSE_BYTES);
423 assert!(append_limited(&mut body, &[2]).is_err());
424 assert_eq!(body.len(), MAX_RESPONSE_BYTES);
425 }
426
427 #[test]
428 fn credential_debug_is_redacted() {
429 let credential = McpCredential {
430 header_name: "authorization".into(),
431 header_value: "Bearer top-secret".into(),
432 };
433 let debug = format!("{credential:?}");
434 assert!(debug.contains("authorization"));
435 assert!(!debug.contains("top-secret"));
436 }
437
438 #[test]
439 fn idempotency_key_is_stable_and_scoped_to_the_durable_call() {
440 let endpoint = McpEndpoint {
441 id: "search".into(),
442 url: "https://mcp.example.com".into(),
443 namespace: "reference".into(),
444 allowed_hosts: ["mcp.example.com".into()].into_iter().collect(),
445 allowed_tools: ["query".into()].into_iter().collect(),
446 credential_ref: None,
447 timeout_ms: 1_000,
448 failure_threshold: 1,
449 recovery_ms: 1_000,
450 };
451 let context = crate::McpCallContext {
452 tenant_id: "tenant-a".parse().unwrap(),
453 subject_id: "subject".parse().unwrap(),
454 session_id: "session".parse().unwrap(),
455 run_id: "run".parse().unwrap(),
456 call_id: "call".parse().unwrap(),
457 source_event_seq: 7,
458 request_id: "request".parse().unwrap(),
459 };
460 let key = scoped_idempotency_key(&endpoint, &context);
461 assert_eq!(key, scoped_idempotency_key(&endpoint, &context));
462 let mut other_tenant = context;
463 other_tenant.tenant_id = "tenant-b".parse().unwrap();
464 assert_ne!(key, scoped_idempotency_key(&endpoint, &other_tenant));
465 assert!(!key.contains("tenant-a"));
466 }
467
468 #[tokio::test]
469 async fn http_sender_classifies_success_rate_limit_server_error_and_timeout_without_secrets() {
470 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;
471 let client = Client::builder()
472 .timeout(Duration::from_secs(1))
473 .build()
474 .unwrap();
475 assert_eq!(
476 send(&client, &ok, None, None, None, json!({}))
477 .await
478 .unwrap()
479 .body["id"],
480 1
481 );
482
483 for status in ["429 Too Many Requests", "503 Service Unavailable"] {
484 let response = Box::leak(
485 format!("HTTP/1.1 {status}\r\nContent-Length: 0\r\n\r\n").into_boxed_str(),
486 );
487 let url = mock_response(response, Duration::ZERO).await;
488 let credential = McpCredential {
489 header_name: "authorization".into(),
490 header_value: "Bearer top-secret".into(),
491 };
492 let error = send(&client, &url, Some(&credential), None, None, json!({}))
493 .await
494 .unwrap_err();
495 assert!(matches!(error, McpError::Unavailable(_)));
496 assert!(!error.to_string().contains("top-secret"));
497 }
498
499 let slow = mock_response(
500 "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n",
501 Duration::from_millis(100),
502 )
503 .await;
504 let impatient = Client::builder()
505 .timeout(Duration::from_millis(5))
506 .build()
507 .unwrap();
508 assert!(matches!(
509 send(&impatient, &slow, None, None, None, json!({})).await,
510 Err(McpError::Unavailable(_))
511 ));
512 }
513}