Skip to main content

af_mcp_client/
transport.rs

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/// Resolved credential injected as one HTTP header.
18#[derive(Clone, PartialEq, Eq)]
19pub struct McpCredential {
20    /// Header carrying the credential (for example `authorization`).
21    pub header_name: String,
22    /// Header value; never logged.
23    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/// Host-owned resolution of credential references to secrets.
37#[async_trait]
38pub trait McpCredentialProvider: Send + Sync {
39    /// Resolve one reference; unknown references fail closed.
40    async fn resolve(&self, reference: &str) -> Result<McpCredential, McpError>;
41}
42
43/// Streamable HTTP transport with DNS pinning and no redirects.
44pub struct StreamableHttpTransport<C> {
45    credentials: C,
46    request_id: AtomicU64,
47}
48
49impl<C> StreamableHttpTransport<C> {
50    /// Transport resolving credentials through `credentials`.
51    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}