Skip to main content

tower_mcp/client/
http.rs

1//! HTTP client transport for remote MCP servers.
2//!
3//! Provides [`HttpClientTransport`] which connects to an MCP server using
4//! the Streamable HTTP transport protocol (MCP spec 2025-11-25). Manages
5//! session lifecycle, SSE stream for server notifications, and HTTP POST
6//! for client requests.
7//!
8//! # Example
9//!
10//! ```rust,no_run
11//! use tower_mcp::client::{McpClient, HttpClientTransport};
12//!
13//! # async fn example() -> Result<(), tower_mcp::BoxError> {
14//! let transport = HttpClientTransport::new("http://localhost:3000");
15//! let client = McpClient::connect(transport).await?;
16//!
17//! let info = client.initialize("my-client", "1.0.0").await?;
18//! println!("Connected to: {}", info.server_info.name);
19//! # Ok(())
20//! # }
21//! ```
22//!
23//! # Authentication
24//!
25//! ```rust,no_run
26//! use tower_mcp::client::{McpClient, HttpClientTransport};
27//!
28//! # async fn example() -> Result<(), tower_mcp::BoxError> {
29//! // Bearer token
30//! let transport = HttpClientTransport::new("http://localhost:3000")
31//!     .bearer_token("sk-your-token-here");
32//!
33//! // Custom API key header
34//! let transport = HttpClientTransport::new("http://localhost:3000")
35//!     .api_key_header("X-API-Key", "your-key");
36//!
37//! // Basic auth
38//! let transport = HttpClientTransport::new("http://localhost:3000")
39//!     .basic_auth("user", "password");
40//! # Ok(())
41//! # }
42//! ```
43
44use std::collections::HashMap;
45use std::sync::Arc;
46use std::sync::atomic::{AtomicBool, Ordering};
47use std::time::Duration;
48
49use async_trait::async_trait;
50use base64::Engine;
51use tokio::sync::{Notify, RwLock, mpsc};
52use tokio::task::JoinHandle;
53
54use super::transport::ClientTransport;
55use crate::error::{Error, Result};
56use crate::protocol::{RequestId, notifications};
57
58const MCP_METHOD_HEADER: &str = "mcp-method";
59const MCP_NAME_HEADER: &str = "mcp-name";
60const MCP_PARAM_HEADER_PREFIX: &str = "mcp-param-";
61const BASE64_SENTINEL_PREFIX: &str = "=?base64?";
62const BASE64_SENTINEL_SUFFIX: &str = "?=";
63
64#[derive(Debug, Clone)]
65struct CustomHeaderMapping {
66    suffix: String,
67    property_path: Vec<String>,
68}
69
70#[cfg(feature = "oauth-client")]
71#[derive(Clone)]
72struct ScopeEscalationRuntime {
73    handler: Arc<dyn OAuthScopeEscalationHandler>,
74    state: Arc<tokio::sync::Mutex<ScopeEscalationState>>,
75    max_attempts: usize,
76}
77
78#[cfg(feature = "oauth-client")]
79struct ScopeEscalationState {
80    scopes: Vec<String>,
81    revision: usize,
82}
83
84#[cfg(feature = "oauth-client")]
85impl ScopeEscalationRuntime {
86    fn new<P>(handler: Arc<P>, config: OAuthScopeEscalationConfig) -> Self
87    where
88        P: OAuthScopeEscalationHandler,
89    {
90        Self {
91            handler,
92            state: Arc::new(tokio::sync::Mutex::new(ScopeEscalationState {
93                scopes: config.initial_scopes().to_vec(),
94                revision: 0,
95            })),
96            max_attempts: config.maximum_attempts(),
97        }
98    }
99
100    async fn respond_to_challenge(
101        &self,
102        challenge: OAuthScopeChallenge,
103        resource: &str,
104        operation: &str,
105        attempt: usize,
106        observed_revision: usize,
107    ) -> std::result::Result<ScopeEscalationDecision, OAuthClientError> {
108        // Serialize the complete reauthorization flow. A second operation
109        // challenged by the same scope can reuse the token produced by the
110        // first rather than opening a duplicate browser/headless flow.
111        let mut state = self.state.lock().await;
112        let previous_scopes = state.scopes.clone();
113        let mut requested_scopes = previous_scopes.clone();
114        for scope in &challenge.required_scopes {
115            if !requested_scopes.contains(scope) {
116                requested_scopes.push(scope.clone());
117            }
118        }
119
120        if requested_scopes == previous_scopes && state.revision > observed_revision {
121            return Ok(ScopeEscalationDecision {
122                revision: state.revision,
123            });
124        }
125        // The scope may have been requested previously without being granted,
126        // or the token may otherwise be stale. Reauthorize the same union
127        // again, but only within the caller's hard attempt limit.
128
129        self.handler
130            .reauthorize(OAuthScopeEscalationRequest {
131                resource: resource.to_string(),
132                operation: operation.to_string(),
133                challenge,
134                previous_scopes,
135                requested_scopes: requested_scopes.clone(),
136                attempt,
137            })
138            .await?;
139
140        state.scopes = requested_scopes;
141        state.revision += 1;
142        Ok(ScopeEscalationDecision {
143            revision: state.revision,
144        })
145    }
146}
147
148#[cfg(feature = "oauth-client")]
149struct ScopeEscalationDecision {
150    revision: usize,
151}
152
153#[cfg(feature = "oauth-client")]
154use super::oauth::{
155    OAuthClientError, OAuthScopeChallenge, OAuthScopeEscalationConfig, OAuthScopeEscalationHandler,
156    OAuthScopeEscalationRequest, TokenProvider,
157};
158
159/// Configuration for [`HttpClientTransport`].
160///
161/// # Example
162///
163/// ```rust,no_run
164/// use tower_mcp::client::{HttpClientTransport, HttpClientConfig};
165/// use std::time::Duration;
166///
167/// let config = HttpClientConfig {
168///     request_timeout: Duration::from_secs(60),
169///     ..Default::default()
170/// };
171/// let transport = HttpClientTransport::with_config("http://localhost:3000", config);
172/// ```
173#[derive(Debug, Clone)]
174pub struct HttpClientConfig {
175    /// Custom headers to include on every request (e.g., auth tokens).
176    pub headers: HashMap<String, String>,
177    /// Whether to automatically open the standalone SSE notification stream
178    /// after initialization. Associated POST response streams used for
179    /// sampling, elicitation, and roots requests work independently.
180    /// Default: `true`.
181    pub auto_sse: bool,
182    /// Capacity of the internal message channel.
183    /// Default: 256.
184    pub channel_capacity: usize,
185    /// Timeout for HTTP requests.
186    /// Default: 30 seconds.
187    pub request_timeout: Duration,
188    /// Timeout for notification POSTs (frames without an `id`), capped at
189    /// `request_timeout`.
190    ///
191    /// Notifications are awaited inline so `notifications/initialized` is
192    /// ordered ahead of the first request, which blocks the client's message
193    /// loop for the duration. This bounds that block independently of
194    /// `request_timeout` so a server that stalls a notification's `202` cannot
195    /// freeze the client for the full request timeout.
196    /// Default: 5 seconds.
197    pub notification_timeout: Duration,
198    /// Whether to attempt SSE reconnection on disconnect.
199    /// Default: `true`.
200    pub sse_reconnect: bool,
201    /// Delay before SSE reconnection attempts.
202    /// Default: 1 second.
203    pub sse_reconnect_delay: Duration,
204    /// Maximum SSE reconnection attempts before giving up.
205    /// Default: 5.
206    pub max_sse_reconnect_attempts: u32,
207    /// Whether to support automatic session recovery on expiry.
208    /// When enabled, HTTP 404 responses (with a session ID attached) and
209    /// JSON-RPC -32005 errors trigger re-initialization.
210    /// Default: `true`.
211    pub session_recovery: bool,
212    /// Maximum size in bytes buffered for a single SSE event.
213    ///
214    /// A server that streams an event without ever terminating it would
215    /// otherwise grow the parse buffer without bound. When a single
216    /// event's buffered size exceeds this cap, the stream is terminated
217    /// with [`Error::SseEventTooLarge`].
218    /// Default: 16 MiB (matching rmcp).
219    pub max_sse_event_size: usize,
220}
221
222/// Default maximum buffered size for a single SSE event (16 MiB, matching
223/// rmcp). See [`HttpClientConfig::max_sse_event_size`].
224pub const DEFAULT_MAX_SSE_EVENT_SIZE: usize = 16 * 1024 * 1024;
225
226impl Default for HttpClientConfig {
227    fn default() -> Self {
228        Self {
229            headers: HashMap::new(),
230            auto_sse: true,
231            channel_capacity: 256,
232            request_timeout: Duration::from_secs(30),
233            notification_timeout: Duration::from_secs(5),
234            sse_reconnect: true,
235            sse_reconnect_delay: Duration::from_secs(1),
236            max_sse_reconnect_attempts: 5,
237            session_recovery: true,
238            max_sse_event_size: DEFAULT_MAX_SSE_EVENT_SIZE,
239        }
240    }
241}
242
243impl HttpClientConfig {
244    /// Set a Bearer token for authentication.
245    pub fn bearer_token(mut self, token: impl Into<String>) -> Self {
246        self.headers.insert(
247            "Authorization".to_string(),
248            format!("Bearer {}", token.into()),
249        );
250        self
251    }
252
253    /// Set an API key using a custom header name.
254    pub fn api_key_header(mut self, name: impl Into<String>, key: impl Into<String>) -> Self {
255        self.headers.insert(name.into(), key.into());
256        self
257    }
258
259    /// Set Basic authentication credentials.
260    pub fn basic_auth(mut self, username: impl AsRef<str>, password: impl AsRef<str>) -> Self {
261        use base64::Engine;
262        let encoded = base64::engine::general_purpose::STANDARD.encode(format!(
263            "{}:{}",
264            username.as_ref(),
265            password.as_ref()
266        ));
267        self.headers
268            .insert("Authorization".to_string(), format!("Basic {}", encoded));
269        self
270    }
271
272    /// Add a custom header.
273    pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
274        self.headers.insert(name.into(), value.into());
275        self
276    }
277}
278
279/// Client transport for MCP servers over Streamable HTTP.
280///
281/// Connects to a remote MCP server using the Streamable HTTP transport
282/// protocol. Manages session lifecycle (`mcp-session-id`), opens an SSE
283/// stream for server-initiated messages, and sends client requests via
284/// HTTP POST.
285///
286/// # How it works
287///
288/// The transport bridges HTTP's request/response model with the
289/// `ClientTransport` trait's `send()`/`recv()` message-passing model:
290///
291/// - **`send()`** POSTs JSON-RPC messages to the server and queues the
292///   response body into an internal channel for `recv()` to return.
293/// - **`recv()`** reads from that channel, which also receives SSE events
294///   from a background task.
295///
296/// After the `initialize` handshake establishes a session, an SSE stream
297/// is automatically opened to receive server notifications and
298/// server-initiated requests.
299///
300/// # Example
301///
302/// ```rust,no_run
303/// use tower_mcp::client::{McpClient, HttpClientTransport};
304///
305/// # async fn example() -> Result<(), tower_mcp::BoxError> {
306/// let transport = HttpClientTransport::new("http://localhost:3000");
307/// let client = McpClient::connect(transport).await?;
308///
309/// let info = client.initialize("my-client", "1.0.0").await?;
310/// let tools = client.list_tools().await?;
311/// client.shutdown().await?;
312/// # Ok(())
313/// # }
314/// ```
315pub struct HttpClientTransport {
316    /// The base URL of the MCP server endpoint.
317    url: String,
318    /// reqwest HTTP client (reused across requests).
319    client: reqwest::Client,
320    /// Session ID received from the server after `initialize`.
321    session_id: Option<String>,
322    /// Negotiated protocol version.
323    protocol_version: Option<String>,
324    /// Validated custom-header mappings learned from the latest tools/list.
325    tool_header_mappings: HashMap<String, Vec<CustomHeaderMapping>>,
326    /// Channel receiver for incoming messages (POST responses + SSE events).
327    incoming_rx: mpsc::Receiver<String>,
328    /// Channel sender used by `send()` to queue POST response bodies
329    /// and cloned for the SSE background task.
330    incoming_tx: mpsc::Sender<String>,
331    /// Handle to the SSE background task, if running.
332    sse_task: Option<JoinHandle<()>>,
333    /// In-flight POST response streams, keyed by their JSON-RPC request ID.
334    ///
335    /// Final `subscriptions/listen` requests stay here until cancelled or the
336    /// server closes them. Other completed tasks are pruned on subsequent
337    /// sends.
338    request_tasks: HashMap<RequestId, JoinHandle<()>>,
339    /// The last SSE event ID received, for stream resumption.
340    last_event_id: Arc<RwLock<Option<String>>>,
341    /// Server-requested retry delay from SSE `retry:` field.
342    sse_retry_delay: Arc<RwLock<Option<Duration>>>,
343    /// Signal to tell the SSE loop to close its current stream and reconnect.
344    sse_reconnect_signal: Arc<Notify>,
345    /// Whether the transport is still connected.
346    connected: Arc<AtomicBool>,
347    /// Configuration options.
348    config: HttpClientConfig,
349    /// Dynamic token provider for OAuth or other token-based auth.
350    #[cfg(feature = "oauth-client")]
351    token_provider: Option<Arc<dyn TokenProvider>>,
352    /// Runtime policy for insufficient-scope challenges.
353    #[cfg(feature = "oauth-client")]
354    scope_escalation: Option<ScopeEscalationRuntime>,
355}
356
357impl HttpClientTransport {
358    /// Create a new HTTP client transport targeting the given URL.
359    ///
360    /// Uses default configuration. The URL should be the MCP server's
361    /// Streamable HTTP endpoint (e.g., `http://localhost:3000`).
362    ///
363    /// # Example
364    ///
365    /// ```rust,no_run
366    /// use tower_mcp::client::HttpClientTransport;
367    ///
368    /// let transport = HttpClientTransport::new("http://localhost:3000");
369    /// ```
370    pub fn new(url: impl Into<String>) -> Self {
371        Self::with_config(url, HttpClientConfig::default())
372    }
373
374    /// Create with custom configuration.
375    ///
376    /// # Example
377    ///
378    /// ```rust,no_run
379    /// use tower_mcp::client::{HttpClientTransport, HttpClientConfig};
380    /// use std::time::Duration;
381    ///
382    /// let config = HttpClientConfig {
383    ///     request_timeout: Duration::from_secs(60),
384    ///     sse_reconnect: false,
385    ///     ..Default::default()
386    /// };
387    /// let transport = HttpClientTransport::with_config("http://localhost:3000", config);
388    /// ```
389    pub fn with_config(url: impl Into<String>, config: HttpClientConfig) -> Self {
390        let (tx, rx) = mpsc::channel(config.channel_capacity);
391        Self {
392            url: url.into(),
393            client: reqwest::Client::new(),
394            session_id: None,
395            protocol_version: None,
396            tool_header_mappings: HashMap::new(),
397            incoming_rx: rx,
398            incoming_tx: tx,
399            sse_task: None,
400            request_tasks: HashMap::new(),
401            last_event_id: Arc::new(RwLock::new(None)),
402            sse_retry_delay: Arc::new(RwLock::new(None)),
403            sse_reconnect_signal: Arc::new(Notify::new()),
404            connected: Arc::new(AtomicBool::new(true)),
405            config,
406            #[cfg(feature = "oauth-client")]
407            token_provider: None,
408            #[cfg(feature = "oauth-client")]
409            scope_escalation: None,
410        }
411    }
412
413    /// Create with an existing `reqwest::Client`.
414    ///
415    /// Use this when you need custom TLS configuration, proxy settings,
416    /// or connection pooling.
417    ///
418    /// # Example
419    ///
420    /// ```rust,no_run
421    /// use tower_mcp::client::HttpClientTransport;
422    ///
423    /// let client = reqwest::Client::builder()
424    ///     .danger_accept_invalid_certs(true) // for development
425    ///     .build()
426    ///     .unwrap();
427    /// let transport = HttpClientTransport::with_client("https://mcp.example.com", client);
428    /// ```
429    pub fn with_client(url: impl Into<String>, client: reqwest::Client) -> Self {
430        let config = HttpClientConfig::default();
431        let (tx, rx) = mpsc::channel(config.channel_capacity);
432        Self {
433            url: url.into(),
434            client,
435            session_id: None,
436            protocol_version: None,
437            tool_header_mappings: HashMap::new(),
438            incoming_rx: rx,
439            incoming_tx: tx,
440            sse_task: None,
441            request_tasks: HashMap::new(),
442            last_event_id: Arc::new(RwLock::new(None)),
443            sse_retry_delay: Arc::new(RwLock::new(None)),
444            sse_reconnect_signal: Arc::new(Notify::new()),
445            connected: Arc::new(AtomicBool::new(true)),
446            config,
447            #[cfg(feature = "oauth-client")]
448            token_provider: None,
449            #[cfg(feature = "oauth-client")]
450            scope_escalation: None,
451        }
452    }
453
454    /// Set a Bearer token for `Authorization: Bearer <token>` authentication.
455    ///
456    /// The token is included on every HTTP request (POST and SSE GET).
457    ///
458    /// # Example
459    ///
460    /// ```rust,no_run
461    /// use tower_mcp::client::HttpClientTransport;
462    ///
463    /// let transport = HttpClientTransport::new("http://localhost:3000")
464    ///     .bearer_token("sk-my-secret-token");
465    /// ```
466    pub fn bearer_token(mut self, token: impl Into<String>) -> Self {
467        self.config.headers.insert(
468            "Authorization".to_string(),
469            format!("Bearer {}", token.into()),
470        );
471        self
472    }
473
474    /// Set an API key for authentication.
475    ///
476    /// Sends as `Authorization: Bearer <key>`. Use
477    /// [`api_key_header`](Self::api_key_header) for a custom header name.
478    ///
479    /// # Example
480    ///
481    /// ```rust,no_run
482    /// use tower_mcp::client::HttpClientTransport;
483    ///
484    /// let transport = HttpClientTransport::new("http://localhost:3000")
485    ///     .api_key("sk-my-api-key");
486    /// ```
487    pub fn api_key(self, key: impl Into<String>) -> Self {
488        self.bearer_token(key)
489    }
490
491    /// Set an API key using a custom header name.
492    ///
493    /// Sends the key as the raw header value (no `Bearer` prefix).
494    ///
495    /// # Example
496    ///
497    /// ```rust,no_run
498    /// use tower_mcp::client::HttpClientTransport;
499    ///
500    /// let transport = HttpClientTransport::new("http://localhost:3000")
501    ///     .api_key_header("X-API-Key", "sk-my-api-key");
502    /// ```
503    pub fn api_key_header(mut self, name: impl Into<String>, key: impl Into<String>) -> Self {
504        self.config.headers.insert(name.into(), key.into());
505        self
506    }
507
508    /// Set Basic authentication credentials.
509    ///
510    /// Encodes `username:password` as Base64 and sends as
511    /// `Authorization: Basic <encoded>`.
512    ///
513    /// # Example
514    ///
515    /// ```rust,no_run
516    /// use tower_mcp::client::HttpClientTransport;
517    ///
518    /// let transport = HttpClientTransport::new("http://localhost:3000")
519    ///     .basic_auth("admin", "secret");
520    /// ```
521    pub fn basic_auth(mut self, username: impl AsRef<str>, password: impl AsRef<str>) -> Self {
522        use base64::Engine;
523        let encoded = base64::engine::general_purpose::STANDARD.encode(format!(
524            "{}:{}",
525            username.as_ref(),
526            password.as_ref()
527        ));
528        self.config
529            .headers
530            .insert("Authorization".to_string(), format!("Basic {}", encoded));
531        self
532    }
533
534    /// Add a custom header to every request.
535    ///
536    /// Can be called multiple times to add multiple headers.
537    ///
538    /// # Example
539    ///
540    /// ```rust,no_run
541    /// use tower_mcp::client::HttpClientTransport;
542    ///
543    /// let transport = HttpClientTransport::new("http://localhost:3000")
544    ///     .header("X-Custom-Header", "my-value")
545    ///     .header("X-Request-Source", "my-app");
546    /// ```
547    pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
548        self.config.headers.insert(name.into(), value.into());
549        self
550    }
551
552    /// Disable automatic session recovery.
553    ///
554    /// By default, if the server returns a session expired error (HTTP 404
555    /// with session ID or JSON-RPC -32005), the client will automatically
556    /// re-initialize and retry the failed operation. Call this to disable
557    /// that behavior and surface the error to the caller instead.
558    pub fn disable_session_recovery(mut self) -> Self {
559        self.config.session_recovery = false;
560        self
561    }
562
563    /// Set a dynamic token provider for authentication.
564    ///
565    /// The provider's [`TokenProvider::get_token()`] is called before each
566    /// HTTP request, and the returned token is sent as `Authorization: Bearer <token>`.
567    /// This overrides any static `Authorization` header set via [`bearer_token()`](Self::bearer_token)
568    /// or [`basic_auth()`](Self::basic_auth).
569    ///
570    /// Use [`OAuthClientCredentials`](super::OAuthClientCredentials) for
571    /// OAuth 2.0 Client Credentials grants, or implement [`TokenProvider`]
572    /// for custom token acquisition logic.
573    ///
574    /// # Example
575    ///
576    /// ```rust,no_run
577    /// use tower_mcp::client::{HttpClientTransport, OAuthClientCredentials};
578    ///
579    /// # fn example() -> Result<(), tower_mcp::BoxError> {
580    /// let provider = OAuthClientCredentials::builder()
581    ///     .client_id("my-client")
582    ///     .client_secret("my-secret")
583    ///     .token_endpoint("https://auth.example.com/token")
584    ///     .resource("http://localhost:3000")
585    ///     .build()?;
586    ///
587    /// let transport = HttpClientTransport::new("http://localhost:3000")
588    ///     .with_token_provider(provider);
589    /// # Ok(())
590    /// # }
591    /// ```
592    #[cfg(feature = "oauth-client")]
593    pub fn with_token_provider(mut self, provider: impl TokenProvider) -> Self {
594        self.token_provider = Some(Arc::new(provider));
595        self.scope_escalation = None;
596        self
597    }
598
599    /// Set a token provider with bounded runtime scope escalation.
600    ///
601    /// When an MCP operation receives an HTTP 403 Bearer challenge with
602    /// `error="insufficient_scope"`, the transport unions the challenged
603    /// scopes with the scopes already tracked by `config`, invokes
604    /// [`OAuthScopeEscalationHandler::reauthorize`], asks the provider for a
605    /// fresh token, and retries the same operation. Reauthorization is
606    /// serialized across concurrent requests, and each operation is bounded
607    /// by [`OAuthScopeEscalationConfig::maximum_attempts`].
608    ///
609    /// The provider and handler are the same value so the handler can update
610    /// the token returned by [`TokenProvider::get_token`].
611    #[cfg(feature = "oauth-client")]
612    pub fn with_scope_aware_token_provider<P>(
613        mut self,
614        provider: P,
615        config: OAuthScopeEscalationConfig,
616    ) -> Self
617    where
618        P: TokenProvider + OAuthScopeEscalationHandler,
619    {
620        let provider = Arc::new(provider);
621        self.token_provider = Some(provider.clone());
622        self.scope_escalation = Some(ScopeEscalationRuntime::new(provider, config));
623        self
624    }
625
626    fn outgoing_custom_headers(&self, parsed: &serde_json::Value) -> Vec<(String, String)> {
627        if parsed.get("method").and_then(serde_json::Value::as_str) != Some("tools/call") {
628            return Vec::new();
629        }
630        let Some(params) = parsed.get("params") else {
631            return Vec::new();
632        };
633        let Some(name) = params.get("name").and_then(serde_json::Value::as_str) else {
634            return Vec::new();
635        };
636        let Some(mappings) = self.tool_header_mappings.get(name) else {
637            return Vec::new();
638        };
639        let arguments = params.get("arguments").unwrap_or(&serde_json::Value::Null);
640
641        mappings
642            .iter()
643            .filter_map(|mapping| {
644                let value = value_at_property_path(arguments, &mapping.property_path)?;
645                if value.is_null() {
646                    return None;
647                }
648                let rendered = json_value_to_header_string(value)?;
649                Some((
650                    format!("{MCP_PARAM_HEADER_PREFIX}{}", mapping.suffix),
651                    encode_header_value(&rendered),
652                ))
653            })
654            .collect()
655    }
656
657    fn normalize_incoming_message(&mut self, message: String) -> String {
658        if self.protocol_version.as_deref() != Some(crate::protocol::PROTOCOL_VERSION_2026_07_28) {
659            return message;
660        }
661        let Ok(mut parsed) = serde_json::from_str::<serde_json::Value>(&message) else {
662            return message;
663        };
664        let Some(tools) = parsed
665            .get_mut("result")
666            .and_then(|result| result.get_mut("tools"))
667            .and_then(serde_json::Value::as_array_mut)
668        else {
669            return message;
670        };
671
672        self.tool_header_mappings.clear();
673        tools.retain(|tool| {
674            let Some(name) = tool.get("name").and_then(serde_json::Value::as_str) else {
675                return false;
676            };
677            let Some(schema) = tool.get("inputSchema") else {
678                return false;
679            };
680            match custom_header_mappings(schema) {
681                Ok(mappings) => {
682                    self.tool_header_mappings.insert(name.to_string(), mappings);
683                    true
684                }
685                Err(error) => {
686                    tracing::warn!(tool = %name, %error, "Excluding tool with invalid x-mcp-header annotations");
687                    false
688                }
689            }
690        });
691
692        parsed.to_string()
693    }
694
695    /// Start the SSE background stream after session is established.
696    fn start_sse_stream(&mut self) {
697        let url = self.url.clone();
698        let client = self.client.clone();
699        let session_id = self.session_id.clone().unwrap();
700        let protocol_version = self.protocol_version.clone();
701        let tx = self.incoming_tx.clone();
702        let last_event_id = self.last_event_id.clone();
703        let sse_retry_delay = self.sse_retry_delay.clone();
704        let reconnect_signal = self.sse_reconnect_signal.clone();
705        let connected = self.connected.clone();
706        let config = self.config.clone();
707        #[cfg(feature = "oauth-client")]
708        let token_provider = self.token_provider.clone();
709
710        self.sse_task = Some(tokio::spawn(async move {
711            sse_stream_loop(SseLoopParams {
712                url,
713                client,
714                session_id,
715                protocol_version,
716                tx,
717                last_event_id,
718                sse_retry_delay,
719                reconnect_signal,
720                connected,
721                config,
722                #[cfg(feature = "oauth-client")]
723                token_provider,
724            })
725            .await;
726        }));
727    }
728}
729
730fn custom_header_mappings(
731    schema: &serde_json::Value,
732) -> std::result::Result<Vec<CustomHeaderMapping>, String> {
733    fn annotation_count(value: &serde_json::Value) -> usize {
734        match value {
735            serde_json::Value::Object(object) => {
736                usize::from(object.contains_key("x-mcp-header"))
737                    + object.values().map(annotation_count).sum::<usize>()
738            }
739            serde_json::Value::Array(values) => values.iter().map(annotation_count).sum::<usize>(),
740            _ => 0,
741        }
742    }
743
744    fn is_tchar(byte: u8) -> bool {
745        byte.is_ascii_alphanumeric()
746            || matches!(
747                byte,
748                b'!' | b'#'
749                    | b'$'
750                    | b'%'
751                    | b'&'
752                    | b'\''
753                    | b'*'
754                    | b'+'
755                    | b'-'
756                    | b'.'
757                    | b'^'
758                    | b'_'
759                    | b'`'
760                    | b'|'
761                    | b'~'
762            )
763    }
764
765    fn primitive_header_type(schema: &serde_json::Value) -> bool {
766        match schema.get("type") {
767            Some(serde_json::Value::String(kind)) => {
768                matches!(kind.as_str(), "string" | "number" | "integer" | "boolean")
769            }
770            Some(serde_json::Value::Array(kinds)) => {
771                let mut primitive = false;
772                for kind in kinds {
773                    match kind.as_str() {
774                        Some("string" | "number" | "integer" | "boolean") if !primitive => {
775                            primitive = true;
776                        }
777                        Some("null") => {}
778                        _ => return false,
779                    }
780                }
781                primitive
782            }
783            _ => false,
784        }
785    }
786
787    fn walk(
788        schema: &serde_json::Value,
789        path: &mut Vec<String>,
790        seen: &mut std::collections::HashSet<String>,
791        mappings: &mut Vec<CustomHeaderMapping>,
792    ) -> std::result::Result<(), String> {
793        let Some(properties) = schema
794            .get("properties")
795            .and_then(serde_json::Value::as_object)
796        else {
797            return Ok(());
798        };
799        for (property_name, property_schema) in properties {
800            path.push(property_name.clone());
801            if let Some(annotation) = property_schema.get("x-mcp-header") {
802                let suffix = annotation
803                    .as_str()
804                    .ok_or_else(|| format!("annotation at {} is not a string", path.join(".")))?;
805                if suffix.is_empty() || !suffix.bytes().all(is_tchar) {
806                    return Err(format!(
807                        "invalid header suffix {suffix:?} at {}",
808                        path.join(".")
809                    ));
810                }
811                if !primitive_header_type(property_schema) {
812                    return Err(format!(
813                        "annotation at {} is not on a primitive property",
814                        path.join(".")
815                    ));
816                }
817                if !seen.insert(suffix.to_ascii_lowercase()) {
818                    return Err(format!("duplicate header suffix {suffix:?}"));
819                }
820                mappings.push(CustomHeaderMapping {
821                    suffix: suffix.to_string(),
822                    property_path: path.clone(),
823                });
824            }
825            walk(property_schema, path, seen, mappings)?;
826            path.pop();
827        }
828        Ok(())
829    }
830
831    let mut mappings = Vec::new();
832    walk(
833        schema,
834        &mut Vec::new(),
835        &mut std::collections::HashSet::new(),
836        &mut mappings,
837    )?;
838    if annotation_count(schema) != mappings.len() {
839        return Err(
840            "x-mcp-header annotation is not statically reachable through properties".to_string(),
841        );
842    }
843    Ok(mappings)
844}
845
846fn value_at_property_path<'a>(
847    root: &'a serde_json::Value,
848    path: &[String],
849) -> Option<&'a serde_json::Value> {
850    path.iter().try_fold(root, |value, key| value.get(key))
851}
852
853fn json_value_to_header_string(value: &serde_json::Value) -> Option<String> {
854    match value {
855        serde_json::Value::String(value) => Some(value.clone()),
856        serde_json::Value::Number(value) => Some(value.to_string()),
857        serde_json::Value::Bool(value) => Some(value.to_string()),
858        serde_json::Value::Null | serde_json::Value::Array(_) | serde_json::Value::Object(_) => {
859            None
860        }
861    }
862}
863
864fn encode_header_value(value: &str) -> String {
865    let unsafe_for_header =
866        value.trim() != value || value.bytes().any(|byte| !(0x20..=0x7e).contains(&byte));
867    if unsafe_for_header {
868        format!(
869            "{BASE64_SENTINEL_PREFIX}{}{BASE64_SENTINEL_SUFFIX}",
870            base64::engine::general_purpose::STANDARD.encode(value)
871        )
872    } else {
873        value.to_string()
874    }
875}
876
877#[cfg(feature = "oauth-client")]
878fn bearer_headers(token: &str) -> std::result::Result<reqwest::header::HeaderMap, String> {
879    let value = reqwest::header::HeaderValue::from_str(&format!("Bearer {token}"))
880        .map_err(|_| "token provider returned an invalid bearer token".to_string())?;
881    let mut headers = reqwest::header::HeaderMap::new();
882    headers.insert(reqwest::header::AUTHORIZATION, value);
883    // RequestBuilder::header appends, which can leave a stale Authorization
884    // value in front of the fresh token. Supplying a HeaderMap replaces the
885    // existing value instead.
886    Ok(headers)
887}
888
889fn is_jsonrpc_error_response(value: &serde_json::Value) -> bool {
890    value.get("error").is_some_and(serde_json::Value::is_object)
891        && value.pointer("/error/code").is_some()
892        && value.pointer("/error/message").is_some()
893}
894
895struct HttpRequestSendError {
896    message: String,
897    connection_failed: bool,
898}
899
900impl HttpRequestSendError {
901    fn request(error: reqwest::Error) -> Self {
902        Self {
903            message: format!("HTTP request failed: {error}"),
904            connection_failed: true,
905        }
906    }
907
908    #[cfg(feature = "oauth-client")]
909    fn oauth(error: OAuthClientError) -> Self {
910        Self {
911            message: error.to_string(),
912            connection_failed: false,
913        }
914    }
915}
916
917async fn send_http_request(
918    request: reqwest::RequestBuilder,
919    resource: &str,
920    operation: &str,
921    #[cfg(feature = "oauth-client")] token_provider: Option<Arc<dyn TokenProvider>>,
922    #[cfg(feature = "oauth-client")] scope_escalation: Option<ScopeEscalationRuntime>,
923    #[cfg(feature = "oauth-client")] initial_scope_revision: usize,
924) -> std::result::Result<reqwest::Response, HttpRequestSendError> {
925    #[cfg(feature = "oauth-client")]
926    let mut request = request;
927    #[cfg(not(feature = "oauth-client"))]
928    let _ = (resource, operation);
929
930    #[cfg(feature = "oauth-client")]
931    let mut observed_revision = initial_scope_revision;
932    #[cfg(feature = "oauth-client")]
933    let mut attempts = 0;
934
935    loop {
936        #[cfg(feature = "oauth-client")]
937        let retry_request = request.try_clone();
938
939        let response = request
940            .send()
941            .await
942            .map_err(HttpRequestSendError::request)?;
943
944        #[cfg(feature = "oauth-client")]
945        {
946            let challenge = if response.status() == reqwest::StatusCode::FORBIDDEN {
947                scope_challenge(response.headers())
948            } else {
949                None
950            };
951            let Some(challenge) = challenge else {
952                return Ok(response);
953            };
954            let (Some(runtime), Some(provider), Some(mut retry_request)) = (
955                scope_escalation.as_ref(),
956                token_provider.as_ref(),
957                retry_request,
958            ) else {
959                return Ok(response);
960            };
961            if attempts >= runtime.max_attempts {
962                return Ok(response);
963            }
964
965            attempts += 1;
966            let decision = runtime
967                .respond_to_challenge(challenge, resource, operation, attempts, observed_revision)
968                .await
969                .map_err(HttpRequestSendError::oauth)?;
970            observed_revision = decision.revision;
971
972            let token = provider
973                .get_token()
974                .await
975                .map_err(HttpRequestSendError::oauth)?;
976            let headers = bearer_headers(&token).map_err(|message| {
977                HttpRequestSendError::oauth(OAuthClientError::ScopeEscalation(message))
978            })?;
979            retry_request = retry_request.headers(headers);
980            request = retry_request;
981        }
982
983        #[cfg(not(feature = "oauth-client"))]
984        return Ok(response);
985    }
986}
987
988#[cfg(feature = "oauth-client")]
989fn scope_challenge(headers: &reqwest::header::HeaderMap) -> Option<OAuthScopeChallenge> {
990    headers
991        .get_all(reqwest::header::WWW_AUTHENTICATE)
992        .iter()
993        .filter_map(|value| value.to_str().ok())
994        .find_map(OAuthScopeChallenge::from_www_authenticate)
995}
996
997fn http_status_error(status: reqwest::StatusCode, headers: &reqwest::header::HeaderMap) -> String {
998    #[cfg(feature = "oauth-client")]
999    if let Some(challenge) = scope_challenge(headers) {
1000        let mut message = format!(
1001            "server returned HTTP {status}: insufficient_scope requires {}",
1002            challenge.required_scopes.join(" ")
1003        );
1004        if let Some(resource_metadata) = challenge.resource_metadata {
1005            message.push_str(&format!(" (resource metadata: {resource_metadata})"));
1006        }
1007        return message;
1008    }
1009
1010    #[cfg(not(feature = "oauth-client"))]
1011    let _ = headers;
1012    format!("server returned HTTP {status}")
1013}
1014
1015fn operation_label(parsed: Option<&serde_json::Value>) -> String {
1016    let Some(method) = parsed
1017        .and_then(|value| value.get("method"))
1018        .and_then(serde_json::Value::as_str)
1019    else {
1020        return "unknown".to_string();
1021    };
1022    let target = match method {
1023        "tools/call" | "prompts/get" => parsed
1024            .and_then(|value| value.pointer("/params/name"))
1025            .and_then(serde_json::Value::as_str),
1026        "resources/read" => parsed
1027            .and_then(|value| value.pointer("/params/uri"))
1028            .and_then(serde_json::Value::as_str),
1029        "tasks/get" | "tasks/update" | "tasks/cancel" => parsed
1030            .and_then(|value| value.pointer("/params/taskId"))
1031            .and_then(serde_json::Value::as_str),
1032        _ => None,
1033    };
1034    match target {
1035        Some(target) => format!("{method}:{target}"),
1036        None => method.to_string(),
1037    }
1038}
1039
1040#[async_trait]
1041impl ClientTransport for HttpClientTransport {
1042    async fn send(&mut self, message: &str) -> Result<()> {
1043        if !self.connected.load(Ordering::Acquire) {
1044            return Err(Error::Transport("Transport closed".to_string()));
1045        }
1046
1047        // Notifications (frames without an `id`) are awaited inline (below) to
1048        // keep `notifications/initialized` ordered before the first request
1049        // (#967). That inline await blocks the whole message loop, so it must
1050        // be bounded independently: a server that stalls the notification's
1051        // 202 (observed: a multi-instance server holding the POST for the full
1052        // request timeout) would otherwise freeze the client with no output.
1053        let parsed_message = serde_json::from_str::<serde_json::Value>(message).ok();
1054        let is_notification = parsed_message
1055            .as_ref()
1056            .map(|v| v.get("id").is_none())
1057            .unwrap_or(false);
1058        let method = parsed_message
1059            .as_ref()
1060            .and_then(|value| value.get("method"))
1061            .and_then(serde_json::Value::as_str);
1062        let operation = operation_label(parsed_message.as_ref());
1063        let outbound_version = parsed_message
1064            .as_ref()
1065            .and_then(|value| {
1066                value.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion")
1067            })
1068            .and_then(serde_json::Value::as_str)
1069            .map(str::to_string);
1070        let is_modern_request =
1071            outbound_version.as_deref() == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28);
1072        if is_modern_request {
1073            self.protocol_version = outbound_version.clone();
1074            // Final requests are sessionless even when a transitional peer
1075            // incorrectly returns a legacy session header from discovery.
1076            self.session_id = None;
1077        }
1078        let timeout = if is_notification {
1079            self.config
1080                .notification_timeout
1081                .min(self.config.request_timeout)
1082        } else {
1083            self.config.request_timeout
1084        };
1085
1086        // Build request with headers
1087        let mut request = self
1088            .client
1089            .post(&self.url)
1090            .header("Content-Type", "application/json")
1091            .header("Accept", "application/json, text/event-stream");
1092        // A subscription is intentionally long-lived; its handle owns the
1093        // timeout/cancellation decision. Every ordinary request remains
1094        // bounded by the configured request timeout.
1095        if method != Some("subscriptions/listen") {
1096            request = request.timeout(timeout);
1097        }
1098
1099        if !is_modern_request && let Some(ref session_id) = self.session_id {
1100            request = request.header("mcp-session-id", session_id);
1101        }
1102
1103        if let Some(version) = outbound_version.as_ref().or(self.protocol_version.as_ref()) {
1104            request = request.header("mcp-protocol-version", version);
1105        }
1106
1107        if let Some(method) = method {
1108            request = request.header(MCP_METHOD_HEADER, method);
1109            let name = match method {
1110                "tools/call" | "prompts/get" => parsed_message
1111                    .as_ref()
1112                    .and_then(|value| value.pointer("/params/name"))
1113                    .and_then(serde_json::Value::as_str),
1114                "resources/read" => parsed_message
1115                    .as_ref()
1116                    .and_then(|value| value.pointer("/params/uri"))
1117                    .and_then(serde_json::Value::as_str),
1118                "tasks/get" | "tasks/update" | "tasks/cancel" => parsed_message
1119                    .as_ref()
1120                    .and_then(|value| value.pointer("/params/taskId"))
1121                    .and_then(serde_json::Value::as_str),
1122                _ => None,
1123            };
1124            if let Some(name) = name {
1125                request = request.header(MCP_NAME_HEADER, name);
1126            }
1127        }
1128
1129        if let Some(parsed) = parsed_message.as_ref() {
1130            for (name, value) in self.outgoing_custom_headers(parsed) {
1131                request = request.header(name, value);
1132            }
1133        }
1134
1135        for (key, value) in &self.config.headers {
1136            request = request.header(key.as_str(), value.as_str());
1137        }
1138
1139        #[cfg(feature = "oauth-client")]
1140        let initial_scope_revision = match &self.scope_escalation {
1141            Some(runtime) => runtime.state.lock().await.revision,
1142            None => 0,
1143        };
1144
1145        // Dynamic token provider overrides static Authorization header
1146        #[cfg(feature = "oauth-client")]
1147        if let Some(ref provider) = self.token_provider {
1148            let token = provider
1149                .get_token()
1150                .await
1151                .map_err(|e| Error::Transport(format!("Token provider error: {}", e)))?;
1152            request = request.headers(bearer_headers(&token).map_err(Error::Transport)?);
1153        }
1154
1155        let request = request.body(message.to_string());
1156
1157        // Notifications are awaited inline (bounded above) even after session
1158        // establishment, so `notifications/initialized` reaches the server
1159        // ahead of the first request rather than racing it on a pooled
1160        // connection, which strict servers rejected (#967).
1161        //
1162        // After session is established, send requests in a background task
1163        // so the message loop can continue processing incoming SSE messages.
1164        // This prevents a deadlock when the server blocks on a
1165        // bidirectional request (sampling/elicitation) that requires the
1166        // client to handle a request on the originating POST response stream.
1167        if !is_notification && (self.session_id.is_some() || is_modern_request) {
1168            let tx = self.incoming_tx.clone();
1169            // The caller is parked on this request id in the message loop's
1170            // correlation map. Every failure branch below delivers a frame
1171            // carrying it, so a background POST that dies (network error,
1172            // timeout, HTTP error, empty body, response stream closed early)
1173            // wakes the caller with an error instead of hanging it forever.
1174            let req_id = parsed_message
1175                .as_ref()
1176                .and_then(|value| value.get("id"))
1177                .cloned();
1178            let request_id = req_id
1179                .clone()
1180                .and_then(|value| serde_json::from_value(value).ok());
1181            let is_subscription = method == Some("subscriptions/listen");
1182            let connected = self.connected.clone();
1183            let last_event_id = self.last_event_id.clone();
1184            let sse_retry_delay = self.sse_retry_delay.clone();
1185            let sse_reconnect_signal = self.sse_reconnect_signal.clone();
1186            let max_sse_event_size = self.config.max_sse_event_size;
1187            let request_resource = self.url.clone();
1188            #[cfg(feature = "oauth-client")]
1189            let token_provider = self.token_provider.clone();
1190            #[cfg(feature = "oauth-client")]
1191            let scope_escalation = self.scope_escalation.clone();
1192            self.request_tasks.retain(|_, task| !task.is_finished());
1193            let task = tokio::spawn(async move {
1194                let response_result = send_http_request(
1195                    request,
1196                    &request_resource,
1197                    &operation,
1198                    #[cfg(feature = "oauth-client")]
1199                    token_provider,
1200                    #[cfg(feature = "oauth-client")]
1201                    scope_escalation,
1202                    #[cfg(feature = "oauth-client")]
1203                    initial_scope_revision,
1204                )
1205                .await;
1206                let response = match response_result {
1207                    Ok(r) => r,
1208                    Err(e) => {
1209                        let connection_failed = e.connection_failed;
1210                        tracing::error!(error = %e.message, "Background HTTP request failed");
1211                        if let Some(id) = &req_id {
1212                            let _ = tx.send(transport_error_frame(id, &e.message)).await;
1213                        }
1214                        if connection_failed {
1215                            connected.store(false, Ordering::Release);
1216                        }
1217                        return;
1218                    }
1219                };
1220
1221                let status = response.status();
1222
1223                // 202 Accepted = notification acknowledged, no body
1224                if status == reqwest::StatusCode::ACCEPTED {
1225                    return;
1226                }
1227
1228                if !status.is_success() {
1229                    let status_error = http_status_error(status, response.headers());
1230                    let body = response.text().await.unwrap_or_default();
1231
1232                    // Forward a JSON-RPC error body so the message loop can
1233                    // detect -32005 (SessionNotFound) and trigger session
1234                    // recovery. If the server did not echo our request id (a
1235                    // null/absent id that is not the session-level -32005
1236                    // signal), inject it so the awaiting caller is woken by
1237                    // this error instead of hanging.
1238                    if !body.is_empty()
1239                        && let Ok(mut v) = serde_json::from_str::<serde_json::Value>(&body)
1240                        && is_jsonrpc_error_response(&v)
1241                    {
1242                        let is_session_signal =
1243                            v.pointer("/error/code").and_then(|c| c.as_i64()) == Some(-32005);
1244                        if !is_session_signal
1245                            && v.get("id").is_none_or(|id| id.is_null())
1246                            && let Some(id) = &req_id
1247                        {
1248                            v["id"] = id.clone();
1249                        }
1250                        let _ = tx.send(v.to_string()).await;
1251                        return;
1252                    }
1253
1254                    tracing::error!(status = %status, body = %body, "HTTP error from server");
1255                    if let Some(id) = &req_id {
1256                        let _ = tx.send(transport_error_frame(id, &status_error)).await;
1257                    }
1258                    connected.store(false, Ordering::Release);
1259                    return;
1260                }
1261
1262                // Check if response is SSE-formatted
1263                let is_sse = response
1264                    .headers()
1265                    .get("content-type")
1266                    .and_then(|v| v.to_str().ok())
1267                    .is_some_and(|ct| ct.contains("text/event-stream"));
1268
1269                if is_sse {
1270                    // Stream SSE response to extract id/retry fields
1271                    let mut stream = response.bytes_stream();
1272                    let mut parser = SseParser::with_limit(max_sse_event_size);
1273                    let mut had_retry = false;
1274                    let mut had_data = false;
1275                    let mut subscription_acknowledged = false;
1276
1277                    use futures::StreamExt;
1278                    while let Some(result) = stream.next().await {
1279                        match result {
1280                            Ok(bytes) => {
1281                                let text = String::from_utf8_lossy(&bytes);
1282                                let events = match parser.feed(&text) {
1283                                    Ok(events) => events,
1284                                    Err(e) => {
1285                                        // A single event exceeded the cap;
1286                                        // terminate the stream instead of
1287                                        // buffering without bound. The
1288                                        // response for this request is lost,
1289                                        // so the transport is unusable.
1290                                        tracing::error!(error = %e, "POST SSE stream terminated");
1291                                        connected.store(false, Ordering::Release);
1292                                        return;
1293                                    }
1294                                };
1295                                for event in events {
1296                                    if let Some(ref id) = event.id {
1297                                        *last_event_id.write().await = Some(id.clone());
1298                                    }
1299                                    if let Some(retry_ms) = event.retry {
1300                                        *sse_retry_delay.write().await =
1301                                            Some(Duration::from_millis(retry_ms));
1302                                        had_retry = true;
1303                                    }
1304                                    if !event.data.is_empty() {
1305                                        had_data = true;
1306                                        let value =
1307                                            serde_json::from_str::<serde_json::Value>(&event.data);
1308                                        let value = match value {
1309                                            Ok(value) => value,
1310                                            Err(error) if is_subscription => {
1311                                                if let Some(id) = &req_id {
1312                                                    let _ = tx
1313                                                        .send(transport_error_frame(
1314                                                            id,
1315                                                            &format!(
1316                                                                "subscription stream returned invalid JSON: {error}"
1317                                                            ),
1318                                                        ))
1319                                                        .await;
1320                                                }
1321                                                return;
1322                                            }
1323                                            Err(_) => {
1324                                                let _ = tx.send(event.data).await;
1325                                                continue;
1326                                            }
1327                                        };
1328                                        let is_terminal =
1329                                            value.get("id").zip(req_id.as_ref()).is_some_and(
1330                                                |(actual, expected)| {
1331                                                    json_request_ids_match(actual, expected)
1332                                                },
1333                                            ) && (value.get("result").is_some()
1334                                                || value.get("error").is_some());
1335
1336                                        if is_subscription {
1337                                            let violation = if is_terminal {
1338                                                if value.get("error").is_some() {
1339                                                    None
1340                                                } else if !subscription_acknowledged {
1341                                                    Some(
1342                                                        "subscriptions/listen completed before acknowledgment",
1343                                                    )
1344                                                } else if !value
1345                                                    .pointer(
1346                                                        "/result/_meta/io.modelcontextprotocol~1subscriptionId",
1347                                                    )
1348                                                    .zip(req_id.as_ref())
1349                                                    .is_some_and(|(actual, expected)| {
1350                                                        json_request_ids_match(actual, expected)
1351                                                    })
1352                                                {
1353                                                    Some(
1354                                                        "subscriptions/listen result carried a missing or mismatched subscription ID",
1355                                                    )
1356                                                } else {
1357                                                    None
1358                                                }
1359                                            } else if value.get("method").is_some()
1360                                                && value.get("id").is_none()
1361                                            {
1362                                                let correlated = value
1363                                                    .pointer(
1364                                                        "/params/_meta/io.modelcontextprotocol~1subscriptionId",
1365                                                    )
1366                                                    .zip(req_id.as_ref())
1367                                                    .is_some_and(|(actual, expected)| {
1368                                                        json_request_ids_match(actual, expected)
1369                                                    });
1370                                                let is_acknowledgment = value
1371                                                    .get("method")
1372                                                    .and_then(serde_json::Value::as_str)
1373                                                    == Some(
1374                                                        notifications::SUBSCRIPTIONS_ACKNOWLEDGED,
1375                                                    );
1376                                                if !correlated {
1377                                                    Some(
1378                                                        "subscription notification carried a missing or mismatched subscription ID",
1379                                                    )
1380                                                } else if !subscription_acknowledged
1381                                                    && !is_acknowledgment
1382                                                {
1383                                                    Some(
1384                                                        "subscription notification arrived before acknowledgment",
1385                                                    )
1386                                                } else if subscription_acknowledged
1387                                                    && is_acknowledgment
1388                                                {
1389                                                    Some(
1390                                                        "subscription stream sent a duplicate acknowledgment",
1391                                                    )
1392                                                } else {
1393                                                    if is_acknowledgment {
1394                                                        subscription_acknowledged = true;
1395                                                    }
1396                                                    None
1397                                                }
1398                                            } else {
1399                                                Some(
1400                                                    "subscription stream returned an unrelated JSON-RPC message",
1401                                                )
1402                                            };
1403                                            if let Some(message) = violation {
1404                                                if let Some(id) = &req_id {
1405                                                    let _ = tx
1406                                                        .send(transport_error_frame(id, message))
1407                                                        .await;
1408                                                }
1409                                                return;
1410                                            }
1411                                        }
1412                                        let _ = tx.send(event.data).await;
1413                                        if is_terminal {
1414                                            // The request is complete. Close the response
1415                                            // body ourselves even if a non-conforming server
1416                                            // leaves the SSE stream open after its final reply.
1417                                            return;
1418                                        }
1419                                    }
1420                                }
1421                            }
1422                            Err(e) => {
1423                                tracing::warn!(error = %e, "POST SSE stream error");
1424                                break;
1425                            }
1426                        }
1427                    }
1428
1429                    // If the POST SSE stream closed with a retry hint but no data,
1430                    // the server expects us to reconnect the GET notification stream.
1431                    // Signal the SSE loop to close its current stream and reconnect
1432                    // with the updated last_event_id and sse_retry_delay.
1433                    if !is_modern_request && had_retry && !had_data {
1434                        sse_reconnect_signal.notify_one();
1435                    } else {
1436                        // The response stream closed without ever delivering a
1437                        // terminal response. Acknowledgments and ordinary
1438                        // notifications do not complete a request, so wake the
1439                        // caller rather than leave it hanging.
1440                        if let Some(id) = &req_id {
1441                            let reason = if had_data {
1442                                "server closed the response stream before the final reply"
1443                            } else {
1444                                "server closed the response stream without a reply"
1445                            };
1446                            let _ = tx.send(transport_error_frame(id, reason)).await;
1447                        }
1448                    }
1449                } else {
1450                    // Non-SSE response: read body and queue for recv().
1451                    match response.text().await {
1452                        Ok(body) if !body.is_empty() => {
1453                            let msgs = extract_json_messages(&body);
1454                            if msgs.is_empty() {
1455                                // A non-empty body that yields no JSON-RPC
1456                                // frames leaves the request uncorrelated.
1457                                if let Some(id) = &req_id {
1458                                    let _ = tx
1459                                        .send(transport_error_frame(
1460                                            id,
1461                                            "server returned an unparseable response body",
1462                                        ))
1463                                        .await;
1464                                }
1465                            } else {
1466                                for msg in msgs {
1467                                    let _ = tx.send(msg).await;
1468                                }
1469                            }
1470                        }
1471                        Ok(_) => {
1472                            // 2xx with an empty body: no frame to correlate the
1473                            // request, so wake the caller rather than hang.
1474                            if let Some(id) = &req_id {
1475                                let _ = tx
1476                                    .send(transport_error_frame(
1477                                        id,
1478                                        "server returned an empty response body",
1479                                    ))
1480                                    .await;
1481                            }
1482                        }
1483                        Err(e) => {
1484                            tracing::error!(error = %e, "Failed to read response body");
1485                            if let Some(id) = &req_id {
1486                                let _ = tx
1487                                    .send(transport_error_frame(
1488                                        id,
1489                                        &format!("failed to read response body: {e}"),
1490                                    ))
1491                                    .await;
1492                            }
1493                            connected.store(false, Ordering::Release);
1494                        }
1495                    }
1496                }
1497            });
1498            if let Some(request_id) = request_id {
1499                self.request_tasks.insert(request_id, task);
1500            }
1501            return Ok(());
1502        }
1503
1504        // Pre-session (initialize) and notifications: handle synchronously.
1505        // For initialize this extracts session headers and starts the SSE
1506        // stream; for notifications the expected response is a bare 202.
1507        let response = send_http_request(
1508            request,
1509            &self.url,
1510            &operation,
1511            #[cfg(feature = "oauth-client")]
1512            self.token_provider.clone(),
1513            #[cfg(feature = "oauth-client")]
1514            self.scope_escalation.clone(),
1515            #[cfg(feature = "oauth-client")]
1516            initial_scope_revision,
1517        )
1518        .await
1519        .map_err(|e| Error::Transport(e.message))?;
1520
1521        let status = response.status();
1522
1523        // Extract session headers before consuming the body
1524        let new_session_id = response
1525            .headers()
1526            .get("mcp-session-id")
1527            .and_then(|v| v.to_str().ok())
1528            .map(|s| s.to_string());
1529        let new_protocol_version = response
1530            .headers()
1531            .get("mcp-protocol-version")
1532            .and_then(|v| v.to_str().ok())
1533            .map(|s| s.to_string());
1534
1535        // 202 Accepted = notification acknowledged, no body
1536        if status == reqwest::StatusCode::ACCEPTED {
1537            // Still update session state if headers present
1538            if !is_modern_request && let Some(sid) = new_session_id {
1539                self.session_id = Some(sid);
1540            }
1541            if let Some(pv) = new_protocol_version {
1542                self.protocol_version = Some(pv);
1543            }
1544            return Ok(());
1545        }
1546
1547        if !status.is_success() {
1548            #[cfg(feature = "oauth-client")]
1549            let status_error = http_status_error(status, response.headers());
1550            let body = response.text().await.unwrap_or_default();
1551            if is_modern_request
1552                && let Ok(mut error) = serde_json::from_str::<serde_json::Value>(&body)
1553                && is_jsonrpc_error_response(&error)
1554            {
1555                if error.get("id").is_none_or(serde_json::Value::is_null)
1556                    && let Some(id) = parsed_message.as_ref().and_then(|value| value.get("id"))
1557                {
1558                    error["id"] = id.clone();
1559                }
1560                self.incoming_tx
1561                    .send(error.to_string())
1562                    .await
1563                    .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1564                return Ok(());
1565            }
1566            // 404 only signals an expired session once a session exists.
1567            // Before that (the initial `initialize`), a 404 means the URL is
1568            // wrong, and reporting it as "Session expired" sends users
1569            // hunting session bugs instead of checking the endpoint path.
1570            if status == reqwest::StatusCode::NOT_FOUND
1571                && self.config.session_recovery
1572                && self.session_id.is_some()
1573            {
1574                return Err(Error::SessionExpired);
1575            }
1576            if status == reqwest::StatusCode::NOT_FOUND && self.session_id.is_none() {
1577                return Err(Error::Transport(format!(
1578                    "HTTP 404 from {}: MCP endpoint not found (check the endpoint path; \
1579                     some servers serve MCP at the root, others at /mcp)",
1580                    self.url
1581                )));
1582            }
1583            #[cfg(feature = "oauth-client")]
1584            if status == reqwest::StatusCode::FORBIDDEN
1585                && status_error.contains("insufficient_scope")
1586            {
1587                return Err(Error::Transport(if body.is_empty() {
1588                    status_error
1589                } else {
1590                    format!("{status_error}: {body}")
1591                }));
1592            }
1593            return Err(Error::Transport(format!(
1594                "HTTP {status} from server: {body}"
1595            )));
1596        }
1597
1598        // Update session state
1599        if !is_modern_request && let Some(sid) = new_session_id {
1600            let is_new_session = self.session_id.is_none();
1601            self.session_id = Some(sid);
1602
1603            if is_new_session && self.config.auto_sse {
1604                self.start_sse_stream();
1605            }
1606        }
1607        if let Some(pv) = new_protocol_version {
1608            self.protocol_version = Some(pv);
1609        }
1610
1611        // Read response body and queue for recv()
1612        let body = response
1613            .text()
1614            .await
1615            .map_err(|e| Error::Transport(format!("Failed to read response: {}", e)))?;
1616
1617        for msg in extract_json_messages(&body) {
1618            self.incoming_tx
1619                .send(msg)
1620                .await
1621                .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1622        }
1623
1624        Ok(())
1625    }
1626
1627    async fn recv(&mut self) -> Result<Option<String>> {
1628        match self.incoming_rx.recv().await {
1629            // All response paths converge here, including background final
1630            // POSTs. Normalize tools/list results on the transport-owning
1631            // task so validated x-mcp-header mappings are available to the
1632            // next tools/call without sharing mutable state across tasks.
1633            Some(msg) => Ok(Some(self.normalize_incoming_message(msg))),
1634            None => {
1635                self.connected.store(false, Ordering::Release);
1636                Ok(None)
1637            }
1638        }
1639    }
1640
1641    fn is_connected(&self) -> bool {
1642        self.connected.load(Ordering::Acquire)
1643    }
1644
1645    async fn close(&mut self) -> Result<()> {
1646        self.connected.store(false, Ordering::Release);
1647
1648        for (_, task) in self.request_tasks.drain() {
1649            task.abort();
1650        }
1651
1652        // Abort SSE task
1653        if let Some(task) = self.sse_task.take() {
1654            task.abort();
1655        }
1656
1657        // Send DELETE to terminate the session (best effort)
1658        if let Some(ref session_id) = self.session_id {
1659            let mut request = self
1660                .client
1661                .delete(&self.url)
1662                .header("mcp-session-id", session_id)
1663                .timeout(Duration::from_secs(5));
1664
1665            for (key, value) in &self.config.headers {
1666                request = request.header(key.as_str(), value.as_str());
1667            }
1668
1669            // Dynamic token provider overrides static Authorization header
1670            #[cfg(feature = "oauth-client")]
1671            if let Some(ref provider) = self.token_provider
1672                && let Ok(token) = provider.get_token().await
1673                && let Ok(headers) = bearer_headers(&token)
1674            {
1675                request = request.headers(headers);
1676            }
1677
1678            let _ = request.send().await;
1679        }
1680
1681        self.session_id = None;
1682        Ok(())
1683    }
1684
1685    async fn reset_session(&mut self) {
1686        tracing::info!("Resetting session for re-initialization");
1687
1688        for (_, task) in self.request_tasks.drain() {
1689            task.abort();
1690        }
1691
1692        // Abort SSE task
1693        if let Some(task) = self.sse_task.take() {
1694            task.abort();
1695        }
1696
1697        // Clear session state but keep the transport alive
1698        self.session_id = None;
1699        self.protocol_version = None;
1700        *self.last_event_id.write().await = None;
1701        *self.sse_retry_delay.write().await = None;
1702
1703        // Drain any stale messages from the channel
1704        while self.incoming_rx.try_recv().is_ok() {}
1705    }
1706
1707    fn supports_session_recovery(&self) -> bool {
1708        self.config.session_recovery
1709    }
1710
1711    async fn cancel_request(&mut self, request_id: &RequestId) -> Result<()> {
1712        if let Some(task) = self.request_tasks.remove(request_id) {
1713            // Dropping reqwest's response byte stream closes this request's
1714            // HTTP response body, which is the final protocol's cancellation
1715            // signal. Other concurrent POST streams remain alive.
1716            task.abort();
1717            let _ = task.await;
1718        }
1719        Ok(())
1720    }
1721}
1722
1723// =============================================================================
1724// SSE Stream Background Loop
1725// =============================================================================
1726
1727/// Parameters for the SSE background loop.
1728struct SseLoopParams {
1729    url: String,
1730    client: reqwest::Client,
1731    session_id: String,
1732    protocol_version: Option<String>,
1733    tx: mpsc::Sender<String>,
1734    last_event_id: Arc<RwLock<Option<String>>>,
1735    sse_retry_delay: Arc<RwLock<Option<Duration>>>,
1736    reconnect_signal: Arc<Notify>,
1737    connected: Arc<AtomicBool>,
1738    config: HttpClientConfig,
1739    #[cfg(feature = "oauth-client")]
1740    token_provider: Option<Arc<dyn TokenProvider>>,
1741}
1742
1743/// Background loop that maintains the SSE stream connection.
1744///
1745/// Opens a GET request with `Accept: text/event-stream` and parses
1746/// incoming SSE events. Events are pushed into the mpsc channel for
1747/// `recv()` to return. Supports reconnection with `Last-Event-ID`.
1748async fn sse_stream_loop(params: SseLoopParams) {
1749    let SseLoopParams {
1750        url,
1751        client,
1752        session_id,
1753        protocol_version,
1754        tx,
1755        last_event_id,
1756        sse_retry_delay,
1757        reconnect_signal,
1758        connected,
1759        config,
1760        #[cfg(feature = "oauth-client")]
1761        token_provider,
1762    } = params;
1763    let mut reconnect_attempts = 0u32;
1764
1765    loop {
1766        if !connected.load(Ordering::Acquire) {
1767            break;
1768        }
1769
1770        let mut request = client
1771            .get(&url)
1772            .header("Accept", "text/event-stream")
1773            .header("mcp-session-id", &session_id);
1774
1775        if let Some(ref version) = protocol_version {
1776            request = request.header("mcp-protocol-version", version);
1777        }
1778
1779        for (key, value) in &config.headers {
1780            request = request.header(key.as_str(), value.as_str());
1781        }
1782
1783        // Dynamic token provider overrides static Authorization header
1784        #[cfg(feature = "oauth-client")]
1785        if let Some(ref provider) = token_provider {
1786            match provider.get_token().await {
1787                Ok(token) => match bearer_headers(&token) {
1788                    Ok(headers) => request = request.headers(headers),
1789                    Err(error) => {
1790                        tracing::warn!(%error, "Token provider failed for SSE connection");
1791                        break;
1792                    }
1793                },
1794                Err(e) => {
1795                    tracing::warn!(error = %e, "Token provider failed for SSE connection");
1796                    break;
1797                }
1798            }
1799        }
1800
1801        // Send Last-Event-ID for stream resumption
1802        if let Some(ref lei) = *last_event_id.read().await {
1803            request = request.header("Last-Event-ID", lei.clone());
1804        }
1805
1806        let response = match request.send().await {
1807            Ok(r) if r.status().is_success() => {
1808                reconnect_attempts = 0;
1809                r
1810            }
1811            Ok(r) => {
1812                tracing::warn!(status = %r.status(), "SSE connection rejected");
1813                break;
1814            }
1815            Err(e) => {
1816                tracing::warn!(error = %e, "SSE connection failed");
1817                if !config.sse_reconnect || reconnect_attempts >= config.max_sse_reconnect_attempts
1818                {
1819                    break;
1820                }
1821                reconnect_attempts += 1;
1822                let delay = sse_retry_delay
1823                    .read()
1824                    .await
1825                    .unwrap_or(config.sse_reconnect_delay);
1826                tokio::time::sleep(delay).await;
1827                continue;
1828            }
1829        };
1830
1831        // Parse SSE stream, also listening for reconnect signals from POST handlers
1832        let mut stream = response.bytes_stream();
1833        let mut parser = SseParser::with_limit(config.max_sse_event_size);
1834
1835        use futures::StreamExt;
1836        loop {
1837            tokio::select! {
1838                chunk = stream.next() => {
1839                    match chunk {
1840                        Some(Ok(bytes)) => {
1841                            let text = String::from_utf8_lossy(&bytes);
1842                            let events = match parser.feed(&text) {
1843                                Ok(events) => events,
1844                                Err(e) => {
1845                                    // A single event exceeded the cap;
1846                                    // terminate the stream (no reconnect,
1847                                    // the server would just repeat it)
1848                                    // instead of buffering without bound.
1849                                    tracing::error!(error = %e, "SSE stream terminated");
1850                                    connected.store(false, Ordering::Release);
1851                                    return;
1852                                }
1853                            };
1854                            for event in events {
1855                                if let Some(ref id) = event.id {
1856                                    *last_event_id.write().await = Some(id.clone());
1857                                }
1858                                if let Some(retry_ms) = event.retry {
1859                                    *sse_retry_delay.write().await = Some(Duration::from_millis(retry_ms));
1860                                }
1861                                if !event.data.is_empty() && tx.send(event.data).await.is_err() {
1862                                    return; // Channel closed, transport dropped
1863                                }
1864                            }
1865                        }
1866                        Some(Err(e)) => {
1867                            tracing::warn!(error = %e, "SSE stream error");
1868                            break;
1869                        }
1870                        None => {
1871                            tracing::debug!("SSE stream ended");
1872                            break;
1873                        }
1874                    }
1875                }
1876                _ = reconnect_signal.notified() => {
1877                    tracing::debug!("SSE reconnect signal received, closing current stream");
1878                    break;
1879                }
1880            }
1881        }
1882
1883        // Attempt reconnection
1884        if !config.sse_reconnect
1885            || !connected.load(Ordering::Acquire)
1886            || reconnect_attempts >= config.max_sse_reconnect_attempts
1887        {
1888            break;
1889        }
1890        reconnect_attempts += 1;
1891        let delay = sse_retry_delay
1892            .read()
1893            .await
1894            .unwrap_or(config.sse_reconnect_delay);
1895        tracing::info!(
1896            attempt = reconnect_attempts,
1897            max = config.max_sse_reconnect_attempts,
1898            delay_ms = delay.as_millis() as u64,
1899            "Reconnecting SSE stream"
1900        );
1901        tokio::time::sleep(delay).await;
1902    }
1903}
1904
1905// =============================================================================
1906// SSE Parser
1907// =============================================================================
1908
1909/// Extract JSON messages from a response body.
1910///
1911/// If the body is SSE-formatted (`event: message\ndata: ...\n\n`), extracts the
1912/// `data:` content from each event. Otherwise returns the body as-is.
1913/// Build a JSON-RPC error frame carrying `id`.
1914///
1915/// A post-session request POST is spawned in the background (so the message
1916/// loop can keep servicing the SSE stream), which means its failures happen
1917/// out of band from the caller. The caller is parked on the request id in the
1918/// message loop's correlation map; if the background POST fails before a real
1919/// response reaches the incoming channel, we must still deliver a frame with
1920/// this id or the caller hangs until the process exits. `-32000` is the
1921/// generic server-error code; the message names the transport-level cause.
1922fn transport_error_frame(id: &serde_json::Value, message: &str) -> String {
1923    serde_json::json!({
1924        "jsonrpc": "2.0",
1925        "id": id,
1926        "error": { "code": -32000, "message": message },
1927    })
1928    .to_string()
1929}
1930
1931fn json_request_ids_match(left: &serde_json::Value, right: &serde_json::Value) -> bool {
1932    left == right
1933        || match (left, right) {
1934            (serde_json::Value::Number(number), serde_json::Value::String(value))
1935            | (serde_json::Value::String(value), serde_json::Value::Number(number)) => number
1936                .as_i64()
1937                .is_some_and(|number| value.parse::<i64>() == Ok(number)),
1938            _ => false,
1939        }
1940}
1941
1942fn extract_json_messages(body: &str) -> Vec<String> {
1943    let trimmed = body.trim();
1944    if trimmed.is_empty() {
1945        return Vec::new();
1946    }
1947
1948    // Heuristic: SSE bodies start with "event:" or "data:" or "id:" or ":"
1949    let looks_like_sse = trimmed.starts_with("event:")
1950        || trimmed.starts_with("data:")
1951        || trimmed.starts_with("id:")
1952        || trimmed.starts_with(':');
1953
1954    if looks_like_sse {
1955        // The body is already fully in memory here, so no event-size cap
1956        // applies: an unlimited parser never returns an error.
1957        let mut parser = SseParser::new();
1958        let events = parser.feed(body).unwrap_or_default();
1959        events.into_iter().map(|e| e.data).collect()
1960    } else {
1961        vec![trimmed.to_string()]
1962    }
1963}
1964
1965/// A parsed SSE event.
1966#[derive(Debug)]
1967struct SseEvent {
1968    /// Event ID (from `id:` line), if present. String per SSE spec.
1969    id: Option<String>,
1970    /// Event data (from `data:` lines, joined with newlines).
1971    data: String,
1972    /// Server-requested retry delay in milliseconds (from `retry:` line).
1973    retry: Option<u64>,
1974}
1975
1976/// Incremental SSE parser.
1977///
1978/// Handles partial chunks from the byte stream, buffering incomplete
1979/// lines across `feed()` calls. When constructed with
1980/// [`with_limit`](Self::with_limit), a single event whose buffered size
1981/// exceeds the limit terminates parsing with
1982/// [`Error::SseEventTooLarge`] instead of growing without bound
1983/// (rmcp #970 analog).
1984struct SseParser {
1985    /// Partial line buffer (when a chunk ends mid-line).
1986    buffer: String,
1987    /// Current event being parsed.
1988    current_id: Option<String>,
1989    current_data: Vec<String>,
1990    current_retry: Option<u64>,
1991    /// Total bytes across `current_data` lines, tracked incrementally.
1992    data_len: usize,
1993    /// Maximum buffered size for a single event.
1994    max_event_size: usize,
1995}
1996
1997impl SseParser {
1998    /// Create a parser with no event-size limit.
1999    fn new() -> Self {
2000        Self::with_limit(usize::MAX)
2001    }
2002
2003    /// Create a parser that rejects events buffering more than
2004    /// `max_event_size` bytes.
2005    fn with_limit(max_event_size: usize) -> Self {
2006        Self {
2007            buffer: String::new(),
2008            current_id: None,
2009            current_data: Vec::new(),
2010            current_retry: None,
2011            data_len: 0,
2012            max_event_size,
2013        }
2014    }
2015
2016    /// Feed a chunk of text and return any complete events.
2017    ///
2018    /// Returns [`Error::SseEventTooLarge`] when the bytes buffered for a
2019    /// single in-progress event exceed the configured limit. The parser
2020    /// should not be fed further after an error.
2021    fn feed(&mut self, text: &str) -> Result<Vec<SseEvent>> {
2022        self.buffer.push_str(text);
2023        let mut events = Vec::new();
2024
2025        // Process complete lines
2026        while let Some(newline_pos) = self.buffer.find('\n') {
2027            let line = self.buffer[..newline_pos]
2028                .trim_end_matches('\r')
2029                .to_string();
2030            self.buffer = self.buffer[newline_pos + 1..].to_string();
2031
2032            if line.is_empty() {
2033                // Empty line = end of event
2034                if !self.current_data.is_empty() || self.current_retry.is_some() {
2035                    events.push(SseEvent {
2036                        id: self.current_id.take(),
2037                        data: self.current_data.join("\n"),
2038                        retry: self.current_retry.take(),
2039                    });
2040                    self.current_data.clear();
2041                    self.data_len = 0;
2042                }
2043                self.current_id = None;
2044                self.current_retry = None;
2045            } else if let Some(value) = line.strip_prefix("id:") {
2046                let trimmed = value.trim();
2047                if !trimmed.is_empty() {
2048                    self.current_id = Some(trimmed.to_string());
2049                }
2050            } else if let Some(value) = line.strip_prefix("data:") {
2051                let data = value.trim().to_string();
2052                self.data_len += data.len();
2053                self.current_data.push(data);
2054            } else if let Some(value) = line.strip_prefix("retry:") {
2055                self.current_retry = value.trim().parse().ok();
2056            }
2057            // Lines starting with ':' are comments (keep-alive) -- ignored
2058            // Lines starting with 'event:' are event types -- ignored (we only care about data)
2059        }
2060
2061        // Everything still buffered belongs to a single unfinished event
2062        // (or an unfinished line of one). Cap it so a server that never
2063        // terminates an event can't grow the buffers without bound.
2064        let buffered = self.buffer.len() + self.data_len;
2065        if buffered > self.max_event_size {
2066            return Err(Error::SseEventTooLarge {
2067                size: buffered,
2068                limit: self.max_event_size,
2069            });
2070        }
2071
2072        Ok(events)
2073    }
2074}
2075
2076#[cfg(test)]
2077mod tests {
2078    use super::*;
2079
2080    // =========================================================================
2081    // SseParser tests
2082    // =========================================================================
2083
2084    #[test]
2085    fn test_parse_complete_event() {
2086        let mut parser = SseParser::new();
2087        let events = parser
2088            .feed("id: 1\nevent: message\ndata: {\"hello\":\"world\"}\n\n")
2089            .unwrap();
2090
2091        assert_eq!(events.len(), 1);
2092        assert_eq!(events[0].id, Some("1".to_string()));
2093        assert_eq!(events[0].data, "{\"hello\":\"world\"}");
2094    }
2095
2096    #[test]
2097    fn test_parse_multiple_events() {
2098        let mut parser = SseParser::new();
2099        let events = parser
2100            .feed("id: 1\ndata: first\n\nid: 2\ndata: second\n\nid: 3\ndata: third\n\n")
2101            .unwrap();
2102
2103        assert_eq!(events.len(), 3);
2104        assert_eq!(events[0].data, "first");
2105        assert_eq!(events[1].data, "second");
2106        assert_eq!(events[2].data, "third");
2107        assert_eq!(events[0].id, Some("1".to_string()));
2108        assert_eq!(events[1].id, Some("2".to_string()));
2109        assert_eq!(events[2].id, Some("3".to_string()));
2110    }
2111
2112    #[test]
2113    fn test_parse_partial_chunks() {
2114        let mut parser = SseParser::new();
2115
2116        // First chunk: partial event
2117        let events = parser.feed("id: 1\nda").unwrap();
2118        assert!(events.is_empty());
2119
2120        // Second chunk: completes the event
2121        let events = parser.feed("ta: hello\n\n").unwrap();
2122        assert_eq!(events.len(), 1);
2123        assert_eq!(events[0].id, Some("1".to_string()));
2124        assert_eq!(events[0].data, "hello");
2125    }
2126
2127    #[test]
2128    fn test_parse_multiline_data() {
2129        let mut parser = SseParser::new();
2130        let events = parser
2131            .feed("id: 1\ndata: line1\ndata: line2\ndata: line3\n\n")
2132            .unwrap();
2133
2134        assert_eq!(events.len(), 1);
2135        assert_eq!(events[0].data, "line1\nline2\nline3");
2136    }
2137
2138    #[test]
2139    fn test_parse_comment_lines() {
2140        let mut parser = SseParser::new();
2141        let events = parser.feed(": keep-alive\nid: 1\ndata: hello\n\n").unwrap();
2142
2143        assert_eq!(events.len(), 1);
2144        assert_eq!(events[0].data, "hello");
2145    }
2146
2147    #[test]
2148    fn test_parse_event_without_id() {
2149        let mut parser = SseParser::new();
2150        let events = parser.feed("data: no-id-event\n\n").unwrap();
2151
2152        assert_eq!(events.len(), 1);
2153        assert_eq!(events[0].id, None);
2154        assert_eq!(events[0].data, "no-id-event");
2155    }
2156
2157    #[test]
2158    fn test_empty_data_no_event() {
2159        let mut parser = SseParser::new();
2160        let events = parser.feed("id: 1\n\n").unwrap();
2161
2162        // No data lines = no event produced
2163        assert!(events.is_empty());
2164    }
2165
2166    #[test]
2167    fn test_parse_crlf_line_endings() {
2168        let mut parser = SseParser::new();
2169        let events = parser.feed("id: 1\r\ndata: crlf\r\n\r\n").unwrap();
2170
2171        assert_eq!(events.len(), 1);
2172        assert_eq!(events[0].data, "crlf");
2173    }
2174
2175    #[test]
2176    fn test_parse_json_data() {
2177        let mut parser = SseParser::new();
2178        let json = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{"token":"t1","progress":50}}"#;
2179        let input = format!("id: 42\nevent: message\ndata: {}\n\n", json);
2180        let events = parser.feed(&input).unwrap();
2181
2182        assert_eq!(events.len(), 1);
2183        assert_eq!(events[0].id, Some("42".to_string()));
2184
2185        // Verify it's valid JSON
2186        let parsed: serde_json::Value = serde_json::from_str(&events[0].data).unwrap();
2187        assert_eq!(parsed["method"], "notifications/progress");
2188    }
2189
2190    #[test]
2191    fn test_event_exceeding_limit_is_rejected() {
2192        let mut parser = SseParser::with_limit(64);
2193
2194        // An unterminated data line larger than the limit trips the cap.
2195        let big = "data: ".to_string() + &"x".repeat(128);
2196        let err = parser.feed(&big).unwrap_err();
2197        match err {
2198            Error::SseEventTooLarge { size, limit } => {
2199                assert!(size > 64, "size {} should exceed limit", size);
2200                assert_eq!(limit, 64);
2201            }
2202            other => panic!("expected SseEventTooLarge, got {:?}", other),
2203        }
2204    }
2205
2206    #[test]
2207    fn test_accumulated_data_lines_count_toward_limit() {
2208        let mut parser = SseParser::with_limit(64);
2209
2210        // Many complete data lines belonging to one unterminated event.
2211        let mut result = Ok(Vec::new());
2212        for _ in 0..10 {
2213            result = parser.feed("data: 0123456789\n");
2214            if result.is_err() {
2215                break;
2216            }
2217        }
2218        assert!(matches!(result, Err(Error::SseEventTooLarge { .. })));
2219    }
2220
2221    #[test]
2222    fn test_events_within_limit_pass() {
2223        let mut parser = SseParser::with_limit(64);
2224        let events = parser.feed("data: hello\n\ndata: world\n\n").unwrap();
2225        assert_eq!(events.len(), 2);
2226    }
2227
2228    // =========================================================================
2229    // Config tests
2230    // =========================================================================
2231
2232    #[test]
2233    fn test_default_config() {
2234        let config = HttpClientConfig::default();
2235        assert!(config.auto_sse);
2236        assert_eq!(config.channel_capacity, 256);
2237        assert_eq!(config.request_timeout, Duration::from_secs(30));
2238        assert!(config.sse_reconnect);
2239        assert_eq!(config.sse_reconnect_delay, Duration::from_secs(1));
2240        assert_eq!(config.max_sse_reconnect_attempts, 5);
2241        assert!(config.headers.is_empty());
2242    }
2243
2244    // =========================================================================
2245    // Transport constructor tests
2246    // =========================================================================
2247
2248    #[test]
2249    fn test_new_transport() {
2250        let transport = HttpClientTransport::new("http://localhost:3000");
2251        assert_eq!(transport.url, "http://localhost:3000");
2252        assert!(transport.session_id.is_none());
2253        assert!(transport.protocol_version.is_none());
2254        assert!(transport.is_connected());
2255    }
2256
2257    #[test]
2258    fn test_with_config() {
2259        let config = HttpClientConfig {
2260            request_timeout: Duration::from_secs(60),
2261            sse_reconnect: false,
2262            ..Default::default()
2263        };
2264        let transport = HttpClientTransport::with_config("http://example.com", config);
2265        assert_eq!(transport.url, "http://example.com");
2266        assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2267        assert!(!transport.config.sse_reconnect);
2268    }
2269
2270    #[test]
2271    fn test_with_client() {
2272        let client = reqwest::Client::new();
2273        let transport = HttpClientTransport::with_client("http://example.com", client);
2274        assert_eq!(transport.url, "http://example.com");
2275        assert!(transport.is_connected());
2276    }
2277
2278    // =========================================================================
2279    // Auth builder tests
2280    // =========================================================================
2281
2282    #[test]
2283    fn test_bearer_token() {
2284        let transport =
2285            HttpClientTransport::new("http://localhost:3000").bearer_token("sk-test-token");
2286        assert_eq!(
2287            transport.config.headers.get("Authorization").unwrap(),
2288            "Bearer sk-test-token"
2289        );
2290    }
2291
2292    #[test]
2293    fn test_api_key() {
2294        let transport = HttpClientTransport::new("http://localhost:3000").api_key("sk-api-key-123");
2295        assert_eq!(
2296            transport.config.headers.get("Authorization").unwrap(),
2297            "Bearer sk-api-key-123"
2298        );
2299    }
2300
2301    #[test]
2302    fn test_api_key_header() {
2303        let transport =
2304            HttpClientTransport::new("http://localhost:3000").api_key_header("X-API-Key", "my-key");
2305        assert_eq!(transport.config.headers.get("X-API-Key").unwrap(), "my-key");
2306        assert!(!transport.config.headers.contains_key("Authorization"));
2307    }
2308
2309    #[test]
2310    fn test_basic_auth() {
2311        let transport =
2312            HttpClientTransport::new("http://localhost:3000").basic_auth("admin", "secret");
2313        let header = transport.config.headers.get("Authorization").unwrap();
2314        assert!(header.starts_with("Basic "));
2315        use base64::Engine;
2316        let decoded = base64::engine::general_purpose::STANDARD
2317            .decode(header.strip_prefix("Basic ").unwrap())
2318            .unwrap();
2319        assert_eq!(String::from_utf8(decoded).unwrap(), "admin:secret");
2320    }
2321
2322    #[test]
2323    fn test_custom_header() {
2324        let transport = HttpClientTransport::new("http://localhost:3000")
2325            .header("X-Custom", "value1")
2326            .header("X-Another", "value2");
2327        assert_eq!(transport.config.headers.get("X-Custom").unwrap(), "value1");
2328        assert_eq!(transport.config.headers.get("X-Another").unwrap(), "value2");
2329    }
2330
2331    #[test]
2332    fn test_chaining_with_config() {
2333        let config = HttpClientConfig {
2334            request_timeout: Duration::from_secs(60),
2335            ..Default::default()
2336        };
2337        let transport =
2338            HttpClientTransport::with_config("http://localhost:3000", config).bearer_token("tk");
2339        assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2340        assert_eq!(
2341            transport.config.headers.get("Authorization").unwrap(),
2342            "Bearer tk"
2343        );
2344    }
2345
2346    #[test]
2347    fn test_last_auth_wins() {
2348        let transport = HttpClientTransport::new("http://localhost:3000")
2349            .bearer_token("token1")
2350            .basic_auth("user", "pass");
2351        let header = transport.config.headers.get("Authorization").unwrap();
2352        assert!(header.starts_with("Basic "));
2353    }
2354
2355    #[test]
2356    fn test_config_bearer_token() {
2357        let config = HttpClientConfig::default().bearer_token("tk-123");
2358        assert_eq!(
2359            config.headers.get("Authorization").unwrap(),
2360            "Bearer tk-123"
2361        );
2362    }
2363
2364    #[test]
2365    fn test_config_header() {
2366        let config = HttpClientConfig::default().header("X-Foo", "bar");
2367        assert_eq!(config.headers.get("X-Foo").unwrap(), "bar");
2368    }
2369
2370    #[test]
2371    fn test_config_api_key_header() {
2372        let config = HttpClientConfig::default().api_key_header("X-Key", "secret");
2373        assert_eq!(config.headers.get("X-Key").unwrap(), "secret");
2374    }
2375
2376    #[test]
2377    fn test_config_basic_auth() {
2378        let config = HttpClientConfig::default().basic_auth("user", "pw");
2379        let header = config.headers.get("Authorization").unwrap();
2380        assert!(header.starts_with("Basic "));
2381    }
2382
2383    #[test]
2384    fn sep_2243_encodes_only_unsafe_values() {
2385        assert_eq!(encode_header_value("us west 1"), "us west 1");
2386        assert_eq!(encode_header_value(""), "");
2387        assert_eq!(encode_header_value(" padded "), "=?base64?IHBhZGRlZCA=?=");
2388        assert_eq!(
2389            encode_header_value("Hello, 世界"),
2390            "=?base64?SGVsbG8sIOS4lueVjA==?="
2391        );
2392    }
2393
2394    #[test]
2395    fn oauth_error_body_is_not_misclassified_as_jsonrpc() {
2396        assert!(!is_jsonrpc_error_response(&serde_json::json!({
2397            "error": "insufficient_scope",
2398            "error_description": "Token has insufficient scope"
2399        })));
2400        assert!(is_jsonrpc_error_response(&serde_json::json!({
2401            "jsonrpc": "2.0",
2402            "id": 1,
2403            "error": {
2404                "code": -32022,
2405                "message": "Unsupported protocol version"
2406            }
2407        })));
2408    }
2409
2410    #[test]
2411    fn sep_2243_validates_custom_header_annotations() {
2412        let mappings = custom_header_mappings(&serde_json::json!({
2413            "type": "object",
2414            "properties": {
2415                "region": {"type": "string", "x-mcp-header": "Region"},
2416                "priority": {"type": "integer", "x-mcp-header": "Priority"},
2417                "ratio": {"type": "number", "x-mcp-header": "Ratio"}
2418            }
2419        }))
2420        .unwrap();
2421        assert_eq!(mappings.len(), 3);
2422
2423        for invalid in [
2424            serde_json::json!({
2425                "type": "object",
2426                "properties": {"value": {"type": "object", "x-mcp-header": "Value"}}
2427            }),
2428            serde_json::json!({
2429                "type": "object",
2430                "properties": {
2431                    "a": {"type": "string", "x-mcp-header": "Region"},
2432                    "b": {"type": "string", "x-mcp-header": "region"}
2433                }
2434            }),
2435            serde_json::json!({
2436                "type": "object",
2437                "properties": {"value": {"type": "string", "x-mcp-header": "Bad Header"}}
2438            }),
2439        ] {
2440            assert!(custom_header_mappings(&invalid).is_err());
2441        }
2442    }
2443
2444    #[test]
2445    fn sep_2243_filters_invalid_tools_and_caches_valid_mappings() {
2446        let mut transport = HttpClientTransport::new("http://localhost:3000");
2447        transport.protocol_version = Some(crate::protocol::PROTOCOL_VERSION_2026_07_28.to_string());
2448        let normalized = transport.normalize_incoming_message(
2449            serde_json::json!({
2450                "jsonrpc": "2.0",
2451                "id": 1,
2452                "result": {
2453                    "tools": [
2454                        {
2455                            "name": "valid",
2456                            "inputSchema": {
2457                                "type": "object",
2458                                "properties": {
2459                                    "region": {"type": "string", "x-mcp-header": "Region"}
2460                                }
2461                            }
2462                        },
2463                        {
2464                            "name": "invalid",
2465                            "inputSchema": {
2466                                "type": "object",
2467                                "properties": {
2468                                    "value": {"type": "array", "x-mcp-header": "Value"}
2469                                }
2470                            }
2471                        }
2472                    ]
2473                }
2474            })
2475            .to_string(),
2476        );
2477        let parsed: serde_json::Value = serde_json::from_str(&normalized).unwrap();
2478        assert_eq!(parsed["result"]["tools"].as_array().unwrap().len(), 1);
2479        assert!(transport.tool_header_mappings.contains_key("valid"));
2480        assert!(!transport.tool_header_mappings.contains_key("invalid"));
2481    }
2482}