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
917#[cfg(not(feature = "oauth-client"))]
918async fn send_http_request(
919    request: reqwest::RequestBuilder,
920    resource: &str,
921    operation: &str,
922) -> std::result::Result<reqwest::Response, HttpRequestSendError> {
923    let _ = (resource, operation);
924    request.send().await.map_err(HttpRequestSendError::request)
925}
926
927#[cfg(feature = "oauth-client")]
928async fn send_http_request(
929    mut request: reqwest::RequestBuilder,
930    resource: &str,
931    operation: &str,
932    token_provider: Option<Arc<dyn TokenProvider>>,
933    scope_escalation: Option<ScopeEscalationRuntime>,
934    initial_scope_revision: usize,
935) -> std::result::Result<reqwest::Response, HttpRequestSendError> {
936    let mut observed_revision = initial_scope_revision;
937    let mut attempts = 0;
938
939    loop {
940        let retry_request = request.try_clone();
941
942        let response = request
943            .send()
944            .await
945            .map_err(HttpRequestSendError::request)?;
946
947        let challenge = if response.status() == reqwest::StatusCode::FORBIDDEN {
948            scope_challenge(response.headers())
949        } else {
950            None
951        };
952        let Some(challenge) = challenge else {
953            return Ok(response);
954        };
955        let (Some(runtime), Some(provider), Some(mut retry_request)) = (
956            scope_escalation.as_ref(),
957            token_provider.as_ref(),
958            retry_request,
959        ) else {
960            return Ok(response);
961        };
962        if attempts >= runtime.max_attempts {
963            return Ok(response);
964        }
965
966        attempts += 1;
967        let decision = runtime
968            .respond_to_challenge(challenge, resource, operation, attempts, observed_revision)
969            .await
970            .map_err(HttpRequestSendError::oauth)?;
971        observed_revision = decision.revision;
972
973        let token = provider
974            .get_token()
975            .await
976            .map_err(HttpRequestSendError::oauth)?;
977        let headers = bearer_headers(&token).map_err(|message| {
978            HttpRequestSendError::oauth(OAuthClientError::ScopeEscalation(message))
979        })?;
980        retry_request = retry_request.headers(headers);
981        request = retry_request;
982    }
983}
984
985#[cfg(feature = "oauth-client")]
986fn scope_challenge(headers: &reqwest::header::HeaderMap) -> Option<OAuthScopeChallenge> {
987    headers
988        .get_all(reqwest::header::WWW_AUTHENTICATE)
989        .iter()
990        .filter_map(|value| value.to_str().ok())
991        .find_map(OAuthScopeChallenge::from_www_authenticate)
992}
993
994fn http_status_error(status: reqwest::StatusCode, headers: &reqwest::header::HeaderMap) -> String {
995    #[cfg(feature = "oauth-client")]
996    if let Some(challenge) = scope_challenge(headers) {
997        let mut message = format!(
998            "server returned HTTP {status}: insufficient_scope requires {}",
999            challenge.required_scopes.join(" ")
1000        );
1001        if let Some(resource_metadata) = challenge.resource_metadata {
1002            message.push_str(&format!(" (resource metadata: {resource_metadata})"));
1003        }
1004        return message;
1005    }
1006
1007    #[cfg(not(feature = "oauth-client"))]
1008    let _ = headers;
1009    format!("server returned HTTP {status}")
1010}
1011
1012fn operation_label(parsed: Option<&serde_json::Value>) -> String {
1013    let Some(method) = parsed
1014        .and_then(|value| value.get("method"))
1015        .and_then(serde_json::Value::as_str)
1016    else {
1017        return "unknown".to_string();
1018    };
1019    let target = match method {
1020        "tools/call" | "prompts/get" => parsed
1021            .and_then(|value| value.pointer("/params/name"))
1022            .and_then(serde_json::Value::as_str),
1023        "resources/read" => parsed
1024            .and_then(|value| value.pointer("/params/uri"))
1025            .and_then(serde_json::Value::as_str),
1026        "tasks/get" | "tasks/update" | "tasks/cancel" => parsed
1027            .and_then(|value| value.pointer("/params/taskId"))
1028            .and_then(serde_json::Value::as_str),
1029        _ => None,
1030    };
1031    match target {
1032        Some(target) => format!("{method}:{target}"),
1033        None => method.to_string(),
1034    }
1035}
1036
1037#[async_trait]
1038impl ClientTransport for HttpClientTransport {
1039    async fn send(&mut self, message: &str) -> Result<()> {
1040        if !self.connected.load(Ordering::Acquire) {
1041            return Err(Error::Transport("Transport closed".to_string()));
1042        }
1043
1044        // Notifications (frames without an `id`) are awaited inline (below) to
1045        // keep `notifications/initialized` ordered before the first request
1046        // (#967). That inline await blocks the whole message loop, so it must
1047        // be bounded independently: a server that stalls the notification's
1048        // 202 (observed: a multi-instance server holding the POST for the full
1049        // request timeout) would otherwise freeze the client with no output.
1050        let parsed_message = serde_json::from_str::<serde_json::Value>(message).ok();
1051        let is_notification = parsed_message
1052            .as_ref()
1053            .map(|v| v.get("id").is_none())
1054            .unwrap_or(false);
1055        let method = parsed_message
1056            .as_ref()
1057            .and_then(|value| value.get("method"))
1058            .and_then(serde_json::Value::as_str);
1059        let operation = operation_label(parsed_message.as_ref());
1060        let outbound_version = parsed_message
1061            .as_ref()
1062            .and_then(|value| {
1063                value.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion")
1064            })
1065            .and_then(serde_json::Value::as_str)
1066            .map(str::to_string);
1067        let is_modern_request =
1068            outbound_version.as_deref() == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28);
1069        if is_modern_request {
1070            self.protocol_version = outbound_version.clone();
1071            // Final requests are sessionless even when a transitional peer
1072            // incorrectly returns a legacy session header from discovery.
1073            self.session_id = None;
1074        }
1075        let timeout = if is_notification {
1076            self.config
1077                .notification_timeout
1078                .min(self.config.request_timeout)
1079        } else {
1080            self.config.request_timeout
1081        };
1082
1083        // Build request with headers
1084        let mut request = self
1085            .client
1086            .post(&self.url)
1087            .header("Content-Type", "application/json")
1088            .header("Accept", "application/json, text/event-stream");
1089        // A subscription is intentionally long-lived; its handle owns the
1090        // timeout/cancellation decision. Every ordinary request remains
1091        // bounded by the configured request timeout.
1092        if method != Some("subscriptions/listen") {
1093            request = request.timeout(timeout);
1094        }
1095
1096        if !is_modern_request && let Some(ref session_id) = self.session_id {
1097            request = request.header("mcp-session-id", session_id);
1098        }
1099
1100        if let Some(version) = outbound_version.as_ref().or(self.protocol_version.as_ref()) {
1101            request = request.header("mcp-protocol-version", version);
1102        }
1103
1104        if let Some(method) = method {
1105            request = request.header(MCP_METHOD_HEADER, method);
1106            let name = match method {
1107                "tools/call" | "prompts/get" => parsed_message
1108                    .as_ref()
1109                    .and_then(|value| value.pointer("/params/name"))
1110                    .and_then(serde_json::Value::as_str),
1111                "resources/read" => parsed_message
1112                    .as_ref()
1113                    .and_then(|value| value.pointer("/params/uri"))
1114                    .and_then(serde_json::Value::as_str),
1115                "tasks/get" | "tasks/update" | "tasks/cancel" => parsed_message
1116                    .as_ref()
1117                    .and_then(|value| value.pointer("/params/taskId"))
1118                    .and_then(serde_json::Value::as_str),
1119                _ => None,
1120            };
1121            if let Some(name) = name {
1122                request = request.header(MCP_NAME_HEADER, name);
1123            }
1124        }
1125
1126        if let Some(parsed) = parsed_message.as_ref() {
1127            for (name, value) in self.outgoing_custom_headers(parsed) {
1128                request = request.header(name, value);
1129            }
1130        }
1131
1132        for (key, value) in &self.config.headers {
1133            request = request.header(key.as_str(), value.as_str());
1134        }
1135
1136        #[cfg(feature = "oauth-client")]
1137        let initial_scope_revision = match &self.scope_escalation {
1138            Some(runtime) => runtime.state.lock().await.revision,
1139            None => 0,
1140        };
1141
1142        // Dynamic token provider overrides static Authorization header
1143        #[cfg(feature = "oauth-client")]
1144        if let Some(ref provider) = self.token_provider {
1145            let token = provider
1146                .get_token()
1147                .await
1148                .map_err(|e| Error::Transport(format!("Token provider error: {}", e)))?;
1149            request = request.headers(bearer_headers(&token).map_err(Error::Transport)?);
1150        }
1151
1152        let request = request.body(message.to_string());
1153
1154        // Notifications are awaited inline (bounded above) even after session
1155        // establishment, so `notifications/initialized` reaches the server
1156        // ahead of the first request rather than racing it on a pooled
1157        // connection, which strict servers rejected (#967).
1158        //
1159        // After session is established, send requests in a background task
1160        // so the message loop can continue processing incoming SSE messages.
1161        // This prevents a deadlock when the server blocks on a
1162        // bidirectional request (sampling/elicitation) that requires the
1163        // client to handle a request on the originating POST response stream.
1164        if !is_notification && (self.session_id.is_some() || is_modern_request) {
1165            let tx = self.incoming_tx.clone();
1166            // The caller is parked on this request id in the message loop's
1167            // correlation map. Every failure branch below delivers a frame
1168            // carrying it, so a background POST that dies (network error,
1169            // timeout, HTTP error, empty body, response stream closed early)
1170            // wakes the caller with an error instead of hanging it forever.
1171            let req_id = parsed_message
1172                .as_ref()
1173                .and_then(|value| value.get("id"))
1174                .cloned();
1175            let request_id = req_id
1176                .clone()
1177                .and_then(|value| serde_json::from_value(value).ok());
1178            let is_subscription = method == Some("subscriptions/listen");
1179            let connected = self.connected.clone();
1180            let last_event_id = self.last_event_id.clone();
1181            let sse_retry_delay = self.sse_retry_delay.clone();
1182            let sse_reconnect_signal = self.sse_reconnect_signal.clone();
1183            let max_sse_event_size = self.config.max_sse_event_size;
1184            let request_resource = self.url.clone();
1185            #[cfg(feature = "oauth-client")]
1186            let token_provider = self.token_provider.clone();
1187            #[cfg(feature = "oauth-client")]
1188            let scope_escalation = self.scope_escalation.clone();
1189            self.request_tasks.retain(|_, task| !task.is_finished());
1190            let task = tokio::spawn(async move {
1191                let response_result = send_http_request(
1192                    request,
1193                    &request_resource,
1194                    &operation,
1195                    #[cfg(feature = "oauth-client")]
1196                    token_provider,
1197                    #[cfg(feature = "oauth-client")]
1198                    scope_escalation,
1199                    #[cfg(feature = "oauth-client")]
1200                    initial_scope_revision,
1201                )
1202                .await;
1203                let response = match response_result {
1204                    Ok(r) => r,
1205                    Err(e) => {
1206                        let connection_failed = e.connection_failed;
1207                        tracing::error!(error = %e.message, "Background HTTP request failed");
1208                        if let Some(id) = &req_id {
1209                            let _ = tx.send(transport_error_frame(id, &e.message)).await;
1210                        }
1211                        if connection_failed {
1212                            connected.store(false, Ordering::Release);
1213                        }
1214                        return;
1215                    }
1216                };
1217
1218                let status = response.status();
1219
1220                // 202 Accepted = notification acknowledged, no body
1221                if status == reqwest::StatusCode::ACCEPTED {
1222                    return;
1223                }
1224
1225                if !status.is_success() {
1226                    let status_error = http_status_error(status, response.headers());
1227                    let body = response.text().await.unwrap_or_default();
1228
1229                    // Forward a JSON-RPC error body so the message loop can
1230                    // detect -32005 (SessionNotFound) and trigger session
1231                    // recovery. If the server did not echo our request id (a
1232                    // null/absent id that is not the session-level -32005
1233                    // signal), inject it so the awaiting caller is woken by
1234                    // this error instead of hanging.
1235                    if !body.is_empty()
1236                        && let Ok(mut v) = serde_json::from_str::<serde_json::Value>(&body)
1237                        && is_jsonrpc_error_response(&v)
1238                    {
1239                        let is_session_signal =
1240                            v.pointer("/error/code").and_then(|c| c.as_i64()) == Some(-32005);
1241                        if !is_session_signal
1242                            && v.get("id").is_none_or(|id| id.is_null())
1243                            && let Some(id) = &req_id
1244                        {
1245                            v["id"] = id.clone();
1246                        }
1247                        let _ = tx.send(v.to_string()).await;
1248                        return;
1249                    }
1250
1251                    tracing::error!(status = %status, body = %body, "HTTP error from server");
1252                    if let Some(id) = &req_id {
1253                        let _ = tx.send(transport_error_frame(id, &status_error)).await;
1254                    }
1255                    connected.store(false, Ordering::Release);
1256                    return;
1257                }
1258
1259                // Check if response is SSE-formatted
1260                let is_sse = response
1261                    .headers()
1262                    .get("content-type")
1263                    .and_then(|v| v.to_str().ok())
1264                    .is_some_and(|ct| ct.contains("text/event-stream"));
1265
1266                if is_sse {
1267                    // Stream SSE response to extract id/retry fields
1268                    let mut stream = response.bytes_stream();
1269                    let mut parser = SseParser::with_limit(max_sse_event_size);
1270                    let mut had_retry = false;
1271                    let mut had_data = false;
1272                    let mut subscription_acknowledged = false;
1273
1274                    use futures::StreamExt;
1275                    while let Some(result) = stream.next().await {
1276                        match result {
1277                            Ok(bytes) => {
1278                                let text = String::from_utf8_lossy(&bytes);
1279                                let events = match parser.feed(&text) {
1280                                    Ok(events) => events,
1281                                    Err(e) => {
1282                                        // A single event exceeded the cap;
1283                                        // terminate the stream instead of
1284                                        // buffering without bound. The
1285                                        // response for this request is lost,
1286                                        // so the transport is unusable.
1287                                        tracing::error!(error = %e, "POST SSE stream terminated");
1288                                        connected.store(false, Ordering::Release);
1289                                        return;
1290                                    }
1291                                };
1292                                for event in events {
1293                                    if let Some(ref id) = event.id {
1294                                        *last_event_id.write().await = Some(id.clone());
1295                                    }
1296                                    if let Some(retry_ms) = event.retry {
1297                                        *sse_retry_delay.write().await =
1298                                            Some(Duration::from_millis(retry_ms));
1299                                        had_retry = true;
1300                                    }
1301                                    if !event.data.is_empty() {
1302                                        had_data = true;
1303                                        let value =
1304                                            serde_json::from_str::<serde_json::Value>(&event.data);
1305                                        let value = match value {
1306                                            Ok(value) => value,
1307                                            Err(error) if is_subscription => {
1308                                                if let Some(id) = &req_id {
1309                                                    let _ = tx
1310                                                        .send(transport_error_frame(
1311                                                            id,
1312                                                            &format!(
1313                                                                "subscription stream returned invalid JSON: {error}"
1314                                                            ),
1315                                                        ))
1316                                                        .await;
1317                                                }
1318                                                return;
1319                                            }
1320                                            Err(_) => {
1321                                                let _ = tx.send(event.data).await;
1322                                                continue;
1323                                            }
1324                                        };
1325                                        let is_terminal =
1326                                            value.get("id").zip(req_id.as_ref()).is_some_and(
1327                                                |(actual, expected)| {
1328                                                    json_request_ids_match(actual, expected)
1329                                                },
1330                                            ) && (value.get("result").is_some()
1331                                                || value.get("error").is_some());
1332
1333                                        if is_subscription {
1334                                            let violation = if is_terminal {
1335                                                if value.get("error").is_some() {
1336                                                    None
1337                                                } else if !subscription_acknowledged {
1338                                                    Some(
1339                                                        "subscriptions/listen completed before acknowledgment",
1340                                                    )
1341                                                } else if !value
1342                                                    .pointer(
1343                                                        "/result/_meta/io.modelcontextprotocol~1subscriptionId",
1344                                                    )
1345                                                    .zip(req_id.as_ref())
1346                                                    .is_some_and(|(actual, expected)| {
1347                                                        json_request_ids_match(actual, expected)
1348                                                    })
1349                                                {
1350                                                    Some(
1351                                                        "subscriptions/listen result carried a missing or mismatched subscription ID",
1352                                                    )
1353                                                } else {
1354                                                    None
1355                                                }
1356                                            } else if value.get("method").is_some()
1357                                                && value.get("id").is_none()
1358                                            {
1359                                                let correlated = value
1360                                                    .pointer(
1361                                                        "/params/_meta/io.modelcontextprotocol~1subscriptionId",
1362                                                    )
1363                                                    .zip(req_id.as_ref())
1364                                                    .is_some_and(|(actual, expected)| {
1365                                                        json_request_ids_match(actual, expected)
1366                                                    });
1367                                                let is_acknowledgment = value
1368                                                    .get("method")
1369                                                    .and_then(serde_json::Value::as_str)
1370                                                    == Some(
1371                                                        notifications::SUBSCRIPTIONS_ACKNOWLEDGED,
1372                                                    );
1373                                                if !correlated {
1374                                                    Some(
1375                                                        "subscription notification carried a missing or mismatched subscription ID",
1376                                                    )
1377                                                } else if !subscription_acknowledged
1378                                                    && !is_acknowledgment
1379                                                {
1380                                                    Some(
1381                                                        "subscription notification arrived before acknowledgment",
1382                                                    )
1383                                                } else if subscription_acknowledged
1384                                                    && is_acknowledgment
1385                                                {
1386                                                    Some(
1387                                                        "subscription stream sent a duplicate acknowledgment",
1388                                                    )
1389                                                } else {
1390                                                    if is_acknowledgment {
1391                                                        subscription_acknowledged = true;
1392                                                    }
1393                                                    None
1394                                                }
1395                                            } else {
1396                                                Some(
1397                                                    "subscription stream returned an unrelated JSON-RPC message",
1398                                                )
1399                                            };
1400                                            if let Some(message) = violation {
1401                                                if let Some(id) = &req_id {
1402                                                    let _ = tx
1403                                                        .send(transport_error_frame(id, message))
1404                                                        .await;
1405                                                }
1406                                                return;
1407                                            }
1408                                        }
1409                                        let _ = tx.send(event.data).await;
1410                                        if is_terminal {
1411                                            // The request is complete. Close the response
1412                                            // body ourselves even if a non-conforming server
1413                                            // leaves the SSE stream open after its final reply.
1414                                            return;
1415                                        }
1416                                    }
1417                                }
1418                            }
1419                            Err(e) => {
1420                                tracing::warn!(error = %e, "POST SSE stream error");
1421                                break;
1422                            }
1423                        }
1424                    }
1425
1426                    // If the POST SSE stream closed with a retry hint but no data,
1427                    // the server expects us to reconnect the GET notification stream.
1428                    // Signal the SSE loop to close its current stream and reconnect
1429                    // with the updated last_event_id and sse_retry_delay.
1430                    if !is_modern_request && had_retry && !had_data {
1431                        sse_reconnect_signal.notify_one();
1432                    } else {
1433                        // The response stream closed without ever delivering a
1434                        // terminal response. Acknowledgments and ordinary
1435                        // notifications do not complete a request, so wake the
1436                        // caller rather than leave it hanging.
1437                        if let Some(id) = &req_id {
1438                            let reason = if had_data {
1439                                "server closed the response stream before the final reply"
1440                            } else {
1441                                "server closed the response stream without a reply"
1442                            };
1443                            let _ = tx.send(transport_error_frame(id, reason)).await;
1444                        }
1445                    }
1446                } else {
1447                    // Non-SSE response: read body and queue for recv().
1448                    match response.text().await {
1449                        Ok(body) if !body.is_empty() => {
1450                            let msgs = extract_json_messages(&body);
1451                            if msgs.is_empty() {
1452                                // A non-empty body that yields no JSON-RPC
1453                                // frames leaves the request uncorrelated.
1454                                if let Some(id) = &req_id {
1455                                    let _ = tx
1456                                        .send(transport_error_frame(
1457                                            id,
1458                                            "server returned an unparseable response body",
1459                                        ))
1460                                        .await;
1461                                }
1462                            } else {
1463                                for msg in msgs {
1464                                    let _ = tx.send(msg).await;
1465                                }
1466                            }
1467                        }
1468                        Ok(_) => {
1469                            // 2xx with an empty body: no frame to correlate the
1470                            // request, so wake the caller rather than hang.
1471                            if let Some(id) = &req_id {
1472                                let _ = tx
1473                                    .send(transport_error_frame(
1474                                        id,
1475                                        "server returned an empty response body",
1476                                    ))
1477                                    .await;
1478                            }
1479                        }
1480                        Err(e) => {
1481                            tracing::error!(error = %e, "Failed to read response body");
1482                            if let Some(id) = &req_id {
1483                                let _ = tx
1484                                    .send(transport_error_frame(
1485                                        id,
1486                                        &format!("failed to read response body: {e}"),
1487                                    ))
1488                                    .await;
1489                            }
1490                            connected.store(false, Ordering::Release);
1491                        }
1492                    }
1493                }
1494            });
1495            if let Some(request_id) = request_id {
1496                self.request_tasks.insert(request_id, task);
1497            }
1498            return Ok(());
1499        }
1500
1501        // Pre-session (initialize) and notifications: handle synchronously.
1502        // For initialize this extracts session headers and starts the SSE
1503        // stream; for notifications the expected response is a bare 202.
1504        let response = send_http_request(
1505            request,
1506            &self.url,
1507            &operation,
1508            #[cfg(feature = "oauth-client")]
1509            self.token_provider.clone(),
1510            #[cfg(feature = "oauth-client")]
1511            self.scope_escalation.clone(),
1512            #[cfg(feature = "oauth-client")]
1513            initial_scope_revision,
1514        )
1515        .await
1516        .map_err(|e| Error::Transport(e.message))?;
1517
1518        let status = response.status();
1519
1520        // Extract session headers before consuming the body
1521        let new_session_id = response
1522            .headers()
1523            .get("mcp-session-id")
1524            .and_then(|v| v.to_str().ok())
1525            .map(|s| s.to_string());
1526        let new_protocol_version = response
1527            .headers()
1528            .get("mcp-protocol-version")
1529            .and_then(|v| v.to_str().ok())
1530            .map(|s| s.to_string());
1531
1532        // 202 Accepted = notification acknowledged, no body
1533        if status == reqwest::StatusCode::ACCEPTED {
1534            // Still update session state if headers present
1535            if !is_modern_request && let Some(sid) = new_session_id {
1536                self.session_id = Some(sid);
1537            }
1538            if let Some(pv) = new_protocol_version {
1539                self.protocol_version = Some(pv);
1540            }
1541            return Ok(());
1542        }
1543
1544        if !status.is_success() {
1545            #[cfg(feature = "oauth-client")]
1546            let status_error = http_status_error(status, response.headers());
1547            let body = response.text().await.unwrap_or_default();
1548            if is_modern_request
1549                && let Ok(mut error) = serde_json::from_str::<serde_json::Value>(&body)
1550                && is_jsonrpc_error_response(&error)
1551            {
1552                if error.get("id").is_none_or(serde_json::Value::is_null)
1553                    && let Some(id) = parsed_message.as_ref().and_then(|value| value.get("id"))
1554                {
1555                    error["id"] = id.clone();
1556                }
1557                self.incoming_tx
1558                    .send(error.to_string())
1559                    .await
1560                    .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1561                return Ok(());
1562            }
1563            // 404 only signals an expired session once a session exists.
1564            // Before that (the initial `initialize`), a 404 means the URL is
1565            // wrong, and reporting it as "Session expired" sends users
1566            // hunting session bugs instead of checking the endpoint path.
1567            if status == reqwest::StatusCode::NOT_FOUND
1568                && self.config.session_recovery
1569                && self.session_id.is_some()
1570            {
1571                return Err(Error::SessionExpired);
1572            }
1573            if status == reqwest::StatusCode::NOT_FOUND && self.session_id.is_none() {
1574                return Err(Error::Transport(format!(
1575                    "HTTP 404 from {}: MCP endpoint not found (check the endpoint path; \
1576                     some servers serve MCP at the root, others at /mcp)",
1577                    self.url
1578                )));
1579            }
1580            #[cfg(feature = "oauth-client")]
1581            if status == reqwest::StatusCode::FORBIDDEN
1582                && status_error.contains("insufficient_scope")
1583            {
1584                return Err(Error::Transport(if body.is_empty() {
1585                    status_error
1586                } else {
1587                    format!("{status_error}: {body}")
1588                }));
1589            }
1590            return Err(Error::Transport(format!(
1591                "HTTP {status} from server: {body}"
1592            )));
1593        }
1594
1595        // Update session state
1596        if !is_modern_request && let Some(sid) = new_session_id {
1597            let is_new_session = self.session_id.is_none();
1598            self.session_id = Some(sid);
1599
1600            if is_new_session && self.config.auto_sse {
1601                self.start_sse_stream();
1602            }
1603        }
1604        if let Some(pv) = new_protocol_version {
1605            self.protocol_version = Some(pv);
1606        }
1607
1608        // Read response body and queue for recv()
1609        let body = response
1610            .text()
1611            .await
1612            .map_err(|e| Error::Transport(format!("Failed to read response: {}", e)))?;
1613
1614        for msg in extract_json_messages(&body) {
1615            self.incoming_tx
1616                .send(msg)
1617                .await
1618                .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1619        }
1620
1621        Ok(())
1622    }
1623
1624    async fn recv(&mut self) -> Result<Option<String>> {
1625        match self.incoming_rx.recv().await {
1626            // All response paths converge here, including background final
1627            // POSTs. Normalize tools/list results on the transport-owning
1628            // task so validated x-mcp-header mappings are available to the
1629            // next tools/call without sharing mutable state across tasks.
1630            Some(msg) => Ok(Some(self.normalize_incoming_message(msg))),
1631            None => {
1632                self.connected.store(false, Ordering::Release);
1633                Ok(None)
1634            }
1635        }
1636    }
1637
1638    fn is_connected(&self) -> bool {
1639        self.connected.load(Ordering::Acquire)
1640    }
1641
1642    async fn close(&mut self) -> Result<()> {
1643        self.connected.store(false, Ordering::Release);
1644
1645        for (_, task) in self.request_tasks.drain() {
1646            task.abort();
1647        }
1648
1649        // Abort SSE task
1650        if let Some(task) = self.sse_task.take() {
1651            task.abort();
1652        }
1653
1654        // Send DELETE to terminate the session (best effort)
1655        if let Some(ref session_id) = self.session_id {
1656            let mut request = self
1657                .client
1658                .delete(&self.url)
1659                .header("mcp-session-id", session_id)
1660                .timeout(Duration::from_secs(5));
1661
1662            for (key, value) in &self.config.headers {
1663                request = request.header(key.as_str(), value.as_str());
1664            }
1665
1666            // Dynamic token provider overrides static Authorization header
1667            #[cfg(feature = "oauth-client")]
1668            if let Some(ref provider) = self.token_provider
1669                && let Ok(token) = provider.get_token().await
1670                && let Ok(headers) = bearer_headers(&token)
1671            {
1672                request = request.headers(headers);
1673            }
1674
1675            let _ = request.send().await;
1676        }
1677
1678        self.session_id = None;
1679        Ok(())
1680    }
1681
1682    async fn reset_session(&mut self) {
1683        tracing::info!("Resetting session for re-initialization");
1684
1685        for (_, task) in self.request_tasks.drain() {
1686            task.abort();
1687        }
1688
1689        // Abort SSE task
1690        if let Some(task) = self.sse_task.take() {
1691            task.abort();
1692        }
1693
1694        // Clear session state but keep the transport alive
1695        self.session_id = None;
1696        self.protocol_version = None;
1697        *self.last_event_id.write().await = None;
1698        *self.sse_retry_delay.write().await = None;
1699
1700        // Drain any stale messages from the channel
1701        while self.incoming_rx.try_recv().is_ok() {}
1702    }
1703
1704    fn supports_session_recovery(&self) -> bool {
1705        self.config.session_recovery
1706    }
1707
1708    async fn cancel_request(&mut self, request_id: &RequestId) -> Result<()> {
1709        if let Some(task) = self.request_tasks.remove(request_id) {
1710            // Dropping reqwest's response byte stream closes this request's
1711            // HTTP response body, which is the final protocol's cancellation
1712            // signal. Other concurrent POST streams remain alive.
1713            task.abort();
1714            let _ = task.await;
1715        }
1716        Ok(())
1717    }
1718}
1719
1720// =============================================================================
1721// SSE Stream Background Loop
1722// =============================================================================
1723
1724/// Parameters for the SSE background loop.
1725struct SseLoopParams {
1726    url: String,
1727    client: reqwest::Client,
1728    session_id: String,
1729    protocol_version: Option<String>,
1730    tx: mpsc::Sender<String>,
1731    last_event_id: Arc<RwLock<Option<String>>>,
1732    sse_retry_delay: Arc<RwLock<Option<Duration>>>,
1733    reconnect_signal: Arc<Notify>,
1734    connected: Arc<AtomicBool>,
1735    config: HttpClientConfig,
1736    #[cfg(feature = "oauth-client")]
1737    token_provider: Option<Arc<dyn TokenProvider>>,
1738}
1739
1740/// Background loop that maintains the SSE stream connection.
1741///
1742/// Opens a GET request with `Accept: text/event-stream` and parses
1743/// incoming SSE events. Events are pushed into the mpsc channel for
1744/// `recv()` to return. Supports reconnection with `Last-Event-ID`.
1745async fn sse_stream_loop(params: SseLoopParams) {
1746    let SseLoopParams {
1747        url,
1748        client,
1749        session_id,
1750        protocol_version,
1751        tx,
1752        last_event_id,
1753        sse_retry_delay,
1754        reconnect_signal,
1755        connected,
1756        config,
1757        #[cfg(feature = "oauth-client")]
1758        token_provider,
1759    } = params;
1760    let mut reconnect_attempts = 0u32;
1761
1762    loop {
1763        if !connected.load(Ordering::Acquire) {
1764            break;
1765        }
1766
1767        let mut request = client
1768            .get(&url)
1769            .header("Accept", "text/event-stream")
1770            .header("mcp-session-id", &session_id);
1771
1772        if let Some(ref version) = protocol_version {
1773            request = request.header("mcp-protocol-version", version);
1774        }
1775
1776        for (key, value) in &config.headers {
1777            request = request.header(key.as_str(), value.as_str());
1778        }
1779
1780        // Dynamic token provider overrides static Authorization header
1781        #[cfg(feature = "oauth-client")]
1782        if let Some(ref provider) = token_provider {
1783            match provider.get_token().await {
1784                Ok(token) => match bearer_headers(&token) {
1785                    Ok(headers) => request = request.headers(headers),
1786                    Err(error) => {
1787                        tracing::warn!(%error, "Token provider failed for SSE connection");
1788                        break;
1789                    }
1790                },
1791                Err(e) => {
1792                    tracing::warn!(error = %e, "Token provider failed for SSE connection");
1793                    break;
1794                }
1795            }
1796        }
1797
1798        // Send Last-Event-ID for stream resumption
1799        if let Some(ref lei) = *last_event_id.read().await {
1800            request = request.header("Last-Event-ID", lei.clone());
1801        }
1802
1803        let response = match request.send().await {
1804            Ok(r) if r.status().is_success() => {
1805                reconnect_attempts = 0;
1806                r
1807            }
1808            Ok(r) => {
1809                tracing::warn!(status = %r.status(), "SSE connection rejected");
1810                break;
1811            }
1812            Err(e) => {
1813                tracing::warn!(error = %e, "SSE connection failed");
1814                if !config.sse_reconnect || reconnect_attempts >= config.max_sse_reconnect_attempts
1815                {
1816                    break;
1817                }
1818                reconnect_attempts += 1;
1819                let delay = sse_retry_delay
1820                    .read()
1821                    .await
1822                    .unwrap_or(config.sse_reconnect_delay);
1823                tokio::time::sleep(delay).await;
1824                continue;
1825            }
1826        };
1827
1828        // Parse SSE stream, also listening for reconnect signals from POST handlers
1829        let mut stream = response.bytes_stream();
1830        let mut parser = SseParser::with_limit(config.max_sse_event_size);
1831
1832        use futures::StreamExt;
1833        loop {
1834            tokio::select! {
1835                chunk = stream.next() => {
1836                    match chunk {
1837                        Some(Ok(bytes)) => {
1838                            let text = String::from_utf8_lossy(&bytes);
1839                            let events = match parser.feed(&text) {
1840                                Ok(events) => events,
1841                                Err(e) => {
1842                                    // A single event exceeded the cap;
1843                                    // terminate the stream (no reconnect,
1844                                    // the server would just repeat it)
1845                                    // instead of buffering without bound.
1846                                    tracing::error!(error = %e, "SSE stream terminated");
1847                                    connected.store(false, Ordering::Release);
1848                                    return;
1849                                }
1850                            };
1851                            for event in events {
1852                                if let Some(ref id) = event.id {
1853                                    *last_event_id.write().await = Some(id.clone());
1854                                }
1855                                if let Some(retry_ms) = event.retry {
1856                                    *sse_retry_delay.write().await = Some(Duration::from_millis(retry_ms));
1857                                }
1858                                if !event.data.is_empty() && tx.send(event.data).await.is_err() {
1859                                    return; // Channel closed, transport dropped
1860                                }
1861                            }
1862                        }
1863                        Some(Err(e)) => {
1864                            tracing::warn!(error = %e, "SSE stream error");
1865                            break;
1866                        }
1867                        None => {
1868                            tracing::debug!("SSE stream ended");
1869                            break;
1870                        }
1871                    }
1872                }
1873                _ = reconnect_signal.notified() => {
1874                    tracing::debug!("SSE reconnect signal received, closing current stream");
1875                    break;
1876                }
1877            }
1878        }
1879
1880        // Attempt reconnection
1881        if !config.sse_reconnect
1882            || !connected.load(Ordering::Acquire)
1883            || reconnect_attempts >= config.max_sse_reconnect_attempts
1884        {
1885            break;
1886        }
1887        reconnect_attempts += 1;
1888        let delay = sse_retry_delay
1889            .read()
1890            .await
1891            .unwrap_or(config.sse_reconnect_delay);
1892        tracing::info!(
1893            attempt = reconnect_attempts,
1894            max = config.max_sse_reconnect_attempts,
1895            delay_ms = delay.as_millis() as u64,
1896            "Reconnecting SSE stream"
1897        );
1898        tokio::time::sleep(delay).await;
1899    }
1900}
1901
1902// =============================================================================
1903// SSE Parser
1904// =============================================================================
1905
1906/// Extract JSON messages from a response body.
1907///
1908/// If the body is SSE-formatted (`event: message\ndata: ...\n\n`), extracts the
1909/// `data:` content from each event. Otherwise returns the body as-is.
1910/// Build a JSON-RPC error frame carrying `id`.
1911///
1912/// A post-session request POST is spawned in the background (so the message
1913/// loop can keep servicing the SSE stream), which means its failures happen
1914/// out of band from the caller. The caller is parked on the request id in the
1915/// message loop's correlation map; if the background POST fails before a real
1916/// response reaches the incoming channel, we must still deliver a frame with
1917/// this id or the caller hangs until the process exits. `-32000` is the
1918/// generic server-error code; the message names the transport-level cause.
1919fn transport_error_frame(id: &serde_json::Value, message: &str) -> String {
1920    serde_json::json!({
1921        "jsonrpc": "2.0",
1922        "id": id,
1923        "error": { "code": -32000, "message": message },
1924    })
1925    .to_string()
1926}
1927
1928fn json_request_ids_match(left: &serde_json::Value, right: &serde_json::Value) -> bool {
1929    left == right
1930        || match (left, right) {
1931            (serde_json::Value::Number(number), serde_json::Value::String(value))
1932            | (serde_json::Value::String(value), serde_json::Value::Number(number)) => number
1933                .as_i64()
1934                .is_some_and(|number| value.parse::<i64>() == Ok(number)),
1935            _ => false,
1936        }
1937}
1938
1939fn extract_json_messages(body: &str) -> Vec<String> {
1940    let trimmed = body.trim();
1941    if trimmed.is_empty() {
1942        return Vec::new();
1943    }
1944
1945    // Heuristic: SSE bodies start with "event:" or "data:" or "id:" or ":"
1946    let looks_like_sse = trimmed.starts_with("event:")
1947        || trimmed.starts_with("data:")
1948        || trimmed.starts_with("id:")
1949        || trimmed.starts_with(':');
1950
1951    if looks_like_sse {
1952        // The body is already fully in memory here, so no event-size cap
1953        // applies: an unlimited parser never returns an error.
1954        let mut parser = SseParser::new();
1955        let events = parser.feed(body).unwrap_or_default();
1956        events.into_iter().map(|e| e.data).collect()
1957    } else {
1958        vec![trimmed.to_string()]
1959    }
1960}
1961
1962/// A parsed SSE event.
1963#[derive(Debug)]
1964struct SseEvent {
1965    /// Event ID (from `id:` line), if present. String per SSE spec.
1966    id: Option<String>,
1967    /// Event data (from `data:` lines, joined with newlines).
1968    data: String,
1969    /// Server-requested retry delay in milliseconds (from `retry:` line).
1970    retry: Option<u64>,
1971}
1972
1973/// Incremental SSE parser.
1974///
1975/// Handles partial chunks from the byte stream, buffering incomplete
1976/// lines across `feed()` calls. When constructed with
1977/// [`with_limit`](Self::with_limit), a single event whose buffered size
1978/// exceeds the limit terminates parsing with
1979/// [`Error::SseEventTooLarge`] instead of growing without bound
1980/// (rmcp #970 analog).
1981struct SseParser {
1982    /// Partial line buffer (when a chunk ends mid-line).
1983    buffer: String,
1984    /// Current event being parsed.
1985    current_id: Option<String>,
1986    current_data: Vec<String>,
1987    current_retry: Option<u64>,
1988    /// Total bytes across `current_data` lines, tracked incrementally.
1989    data_len: usize,
1990    /// Maximum buffered size for a single event.
1991    max_event_size: usize,
1992}
1993
1994impl SseParser {
1995    /// Create a parser with no event-size limit.
1996    fn new() -> Self {
1997        Self::with_limit(usize::MAX)
1998    }
1999
2000    /// Create a parser that rejects events buffering more than
2001    /// `max_event_size` bytes.
2002    fn with_limit(max_event_size: usize) -> Self {
2003        Self {
2004            buffer: String::new(),
2005            current_id: None,
2006            current_data: Vec::new(),
2007            current_retry: None,
2008            data_len: 0,
2009            max_event_size,
2010        }
2011    }
2012
2013    /// Feed a chunk of text and return any complete events.
2014    ///
2015    /// Returns [`Error::SseEventTooLarge`] when the bytes buffered for a
2016    /// single in-progress event exceed the configured limit. The parser
2017    /// should not be fed further after an error.
2018    fn feed(&mut self, text: &str) -> Result<Vec<SseEvent>> {
2019        self.buffer.push_str(text);
2020        let mut events = Vec::new();
2021
2022        // Process complete lines
2023        while let Some(newline_pos) = self.buffer.find('\n') {
2024            let line = self.buffer[..newline_pos]
2025                .trim_end_matches('\r')
2026                .to_string();
2027            self.buffer = self.buffer[newline_pos + 1..].to_string();
2028
2029            if line.is_empty() {
2030                // Empty line = end of event
2031                if !self.current_data.is_empty() || self.current_retry.is_some() {
2032                    events.push(SseEvent {
2033                        id: self.current_id.take(),
2034                        data: self.current_data.join("\n"),
2035                        retry: self.current_retry.take(),
2036                    });
2037                    self.current_data.clear();
2038                    self.data_len = 0;
2039                }
2040                self.current_id = None;
2041                self.current_retry = None;
2042            } else if let Some(value) = line.strip_prefix("id:") {
2043                let trimmed = value.trim();
2044                if !trimmed.is_empty() {
2045                    self.current_id = Some(trimmed.to_string());
2046                }
2047            } else if let Some(value) = line.strip_prefix("data:") {
2048                let data = value.trim().to_string();
2049                self.data_len += data.len();
2050                self.current_data.push(data);
2051            } else if let Some(value) = line.strip_prefix("retry:") {
2052                self.current_retry = value.trim().parse().ok();
2053            }
2054            // Lines starting with ':' are comments (keep-alive) -- ignored
2055            // Lines starting with 'event:' are event types -- ignored (we only care about data)
2056        }
2057
2058        // Everything still buffered belongs to a single unfinished event
2059        // (or an unfinished line of one). Cap it so a server that never
2060        // terminates an event can't grow the buffers without bound.
2061        let buffered = self.buffer.len() + self.data_len;
2062        if buffered > self.max_event_size {
2063            return Err(Error::SseEventTooLarge {
2064                size: buffered,
2065                limit: self.max_event_size,
2066            });
2067        }
2068
2069        Ok(events)
2070    }
2071}
2072
2073#[cfg(test)]
2074mod tests {
2075    use super::*;
2076
2077    // =========================================================================
2078    // SseParser tests
2079    // =========================================================================
2080
2081    #[test]
2082    fn test_parse_complete_event() {
2083        let mut parser = SseParser::new();
2084        let events = parser
2085            .feed("id: 1\nevent: message\ndata: {\"hello\":\"world\"}\n\n")
2086            .unwrap();
2087
2088        assert_eq!(events.len(), 1);
2089        assert_eq!(events[0].id, Some("1".to_string()));
2090        assert_eq!(events[0].data, "{\"hello\":\"world\"}");
2091    }
2092
2093    #[test]
2094    fn test_parse_multiple_events() {
2095        let mut parser = SseParser::new();
2096        let events = parser
2097            .feed("id: 1\ndata: first\n\nid: 2\ndata: second\n\nid: 3\ndata: third\n\n")
2098            .unwrap();
2099
2100        assert_eq!(events.len(), 3);
2101        assert_eq!(events[0].data, "first");
2102        assert_eq!(events[1].data, "second");
2103        assert_eq!(events[2].data, "third");
2104        assert_eq!(events[0].id, Some("1".to_string()));
2105        assert_eq!(events[1].id, Some("2".to_string()));
2106        assert_eq!(events[2].id, Some("3".to_string()));
2107    }
2108
2109    #[test]
2110    fn test_parse_partial_chunks() {
2111        let mut parser = SseParser::new();
2112
2113        // First chunk: partial event
2114        let events = parser.feed("id: 1\nda").unwrap();
2115        assert!(events.is_empty());
2116
2117        // Second chunk: completes the event
2118        let events = parser.feed("ta: hello\n\n").unwrap();
2119        assert_eq!(events.len(), 1);
2120        assert_eq!(events[0].id, Some("1".to_string()));
2121        assert_eq!(events[0].data, "hello");
2122    }
2123
2124    #[test]
2125    fn test_parse_multiline_data() {
2126        let mut parser = SseParser::new();
2127        let events = parser
2128            .feed("id: 1\ndata: line1\ndata: line2\ndata: line3\n\n")
2129            .unwrap();
2130
2131        assert_eq!(events.len(), 1);
2132        assert_eq!(events[0].data, "line1\nline2\nline3");
2133    }
2134
2135    #[test]
2136    fn test_parse_comment_lines() {
2137        let mut parser = SseParser::new();
2138        let events = parser.feed(": keep-alive\nid: 1\ndata: hello\n\n").unwrap();
2139
2140        assert_eq!(events.len(), 1);
2141        assert_eq!(events[0].data, "hello");
2142    }
2143
2144    #[test]
2145    fn test_parse_event_without_id() {
2146        let mut parser = SseParser::new();
2147        let events = parser.feed("data: no-id-event\n\n").unwrap();
2148
2149        assert_eq!(events.len(), 1);
2150        assert_eq!(events[0].id, None);
2151        assert_eq!(events[0].data, "no-id-event");
2152    }
2153
2154    #[test]
2155    fn test_empty_data_no_event() {
2156        let mut parser = SseParser::new();
2157        let events = parser.feed("id: 1\n\n").unwrap();
2158
2159        // No data lines = no event produced
2160        assert!(events.is_empty());
2161    }
2162
2163    #[test]
2164    fn test_parse_crlf_line_endings() {
2165        let mut parser = SseParser::new();
2166        let events = parser.feed("id: 1\r\ndata: crlf\r\n\r\n").unwrap();
2167
2168        assert_eq!(events.len(), 1);
2169        assert_eq!(events[0].data, "crlf");
2170    }
2171
2172    #[test]
2173    fn test_parse_json_data() {
2174        let mut parser = SseParser::new();
2175        let json = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{"token":"t1","progress":50}}"#;
2176        let input = format!("id: 42\nevent: message\ndata: {}\n\n", json);
2177        let events = parser.feed(&input).unwrap();
2178
2179        assert_eq!(events.len(), 1);
2180        assert_eq!(events[0].id, Some("42".to_string()));
2181
2182        // Verify it's valid JSON
2183        let parsed: serde_json::Value = serde_json::from_str(&events[0].data).unwrap();
2184        assert_eq!(parsed["method"], "notifications/progress");
2185    }
2186
2187    #[test]
2188    fn test_event_exceeding_limit_is_rejected() {
2189        let mut parser = SseParser::with_limit(64);
2190
2191        // An unterminated data line larger than the limit trips the cap.
2192        let big = "data: ".to_string() + &"x".repeat(128);
2193        let err = parser.feed(&big).unwrap_err();
2194        match err {
2195            Error::SseEventTooLarge { size, limit } => {
2196                assert!(size > 64, "size {} should exceed limit", size);
2197                assert_eq!(limit, 64);
2198            }
2199            other => panic!("expected SseEventTooLarge, got {:?}", other),
2200        }
2201    }
2202
2203    #[test]
2204    fn test_accumulated_data_lines_count_toward_limit() {
2205        let mut parser = SseParser::with_limit(64);
2206
2207        // Many complete data lines belonging to one unterminated event.
2208        let mut result = Ok(Vec::new());
2209        for _ in 0..10 {
2210            result = parser.feed("data: 0123456789\n");
2211            if result.is_err() {
2212                break;
2213            }
2214        }
2215        assert!(matches!(result, Err(Error::SseEventTooLarge { .. })));
2216    }
2217
2218    #[test]
2219    fn test_events_within_limit_pass() {
2220        let mut parser = SseParser::with_limit(64);
2221        let events = parser.feed("data: hello\n\ndata: world\n\n").unwrap();
2222        assert_eq!(events.len(), 2);
2223    }
2224
2225    // =========================================================================
2226    // Config tests
2227    // =========================================================================
2228
2229    #[test]
2230    fn test_default_config() {
2231        let config = HttpClientConfig::default();
2232        assert!(config.auto_sse);
2233        assert_eq!(config.channel_capacity, 256);
2234        assert_eq!(config.request_timeout, Duration::from_secs(30));
2235        assert!(config.sse_reconnect);
2236        assert_eq!(config.sse_reconnect_delay, Duration::from_secs(1));
2237        assert_eq!(config.max_sse_reconnect_attempts, 5);
2238        assert!(config.headers.is_empty());
2239    }
2240
2241    // =========================================================================
2242    // Transport constructor tests
2243    // =========================================================================
2244
2245    #[test]
2246    fn test_new_transport() {
2247        let transport = HttpClientTransport::new("http://localhost:3000");
2248        assert_eq!(transport.url, "http://localhost:3000");
2249        assert!(transport.session_id.is_none());
2250        assert!(transport.protocol_version.is_none());
2251        assert!(transport.is_connected());
2252    }
2253
2254    #[test]
2255    fn test_with_config() {
2256        let config = HttpClientConfig {
2257            request_timeout: Duration::from_secs(60),
2258            sse_reconnect: false,
2259            ..Default::default()
2260        };
2261        let transport = HttpClientTransport::with_config("http://example.com", config);
2262        assert_eq!(transport.url, "http://example.com");
2263        assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2264        assert!(!transport.config.sse_reconnect);
2265    }
2266
2267    #[test]
2268    fn test_with_client() {
2269        let client = reqwest::Client::new();
2270        let transport = HttpClientTransport::with_client("http://example.com", client);
2271        assert_eq!(transport.url, "http://example.com");
2272        assert!(transport.is_connected());
2273    }
2274
2275    // =========================================================================
2276    // Auth builder tests
2277    // =========================================================================
2278
2279    #[test]
2280    fn test_bearer_token() {
2281        let transport =
2282            HttpClientTransport::new("http://localhost:3000").bearer_token("sk-test-token");
2283        assert_eq!(
2284            transport.config.headers.get("Authorization").unwrap(),
2285            "Bearer sk-test-token"
2286        );
2287    }
2288
2289    #[test]
2290    fn test_api_key() {
2291        let transport = HttpClientTransport::new("http://localhost:3000").api_key("sk-api-key-123");
2292        assert_eq!(
2293            transport.config.headers.get("Authorization").unwrap(),
2294            "Bearer sk-api-key-123"
2295        );
2296    }
2297
2298    #[test]
2299    fn test_api_key_header() {
2300        let transport =
2301            HttpClientTransport::new("http://localhost:3000").api_key_header("X-API-Key", "my-key");
2302        assert_eq!(transport.config.headers.get("X-API-Key").unwrap(), "my-key");
2303        assert!(!transport.config.headers.contains_key("Authorization"));
2304    }
2305
2306    #[test]
2307    fn test_basic_auth() {
2308        let transport =
2309            HttpClientTransport::new("http://localhost:3000").basic_auth("admin", "secret");
2310        let header = transport.config.headers.get("Authorization").unwrap();
2311        assert!(header.starts_with("Basic "));
2312        use base64::Engine;
2313        let decoded = base64::engine::general_purpose::STANDARD
2314            .decode(header.strip_prefix("Basic ").unwrap())
2315            .unwrap();
2316        assert_eq!(String::from_utf8(decoded).unwrap(), "admin:secret");
2317    }
2318
2319    #[test]
2320    fn test_custom_header() {
2321        let transport = HttpClientTransport::new("http://localhost:3000")
2322            .header("X-Custom", "value1")
2323            .header("X-Another", "value2");
2324        assert_eq!(transport.config.headers.get("X-Custom").unwrap(), "value1");
2325        assert_eq!(transport.config.headers.get("X-Another").unwrap(), "value2");
2326    }
2327
2328    #[test]
2329    fn test_chaining_with_config() {
2330        let config = HttpClientConfig {
2331            request_timeout: Duration::from_secs(60),
2332            ..Default::default()
2333        };
2334        let transport =
2335            HttpClientTransport::with_config("http://localhost:3000", config).bearer_token("tk");
2336        assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2337        assert_eq!(
2338            transport.config.headers.get("Authorization").unwrap(),
2339            "Bearer tk"
2340        );
2341    }
2342
2343    #[test]
2344    fn test_last_auth_wins() {
2345        let transport = HttpClientTransport::new("http://localhost:3000")
2346            .bearer_token("token1")
2347            .basic_auth("user", "pass");
2348        let header = transport.config.headers.get("Authorization").unwrap();
2349        assert!(header.starts_with("Basic "));
2350    }
2351
2352    #[test]
2353    fn test_config_bearer_token() {
2354        let config = HttpClientConfig::default().bearer_token("tk-123");
2355        assert_eq!(
2356            config.headers.get("Authorization").unwrap(),
2357            "Bearer tk-123"
2358        );
2359    }
2360
2361    #[test]
2362    fn test_config_header() {
2363        let config = HttpClientConfig::default().header("X-Foo", "bar");
2364        assert_eq!(config.headers.get("X-Foo").unwrap(), "bar");
2365    }
2366
2367    #[test]
2368    fn test_config_api_key_header() {
2369        let config = HttpClientConfig::default().api_key_header("X-Key", "secret");
2370        assert_eq!(config.headers.get("X-Key").unwrap(), "secret");
2371    }
2372
2373    #[test]
2374    fn test_config_basic_auth() {
2375        let config = HttpClientConfig::default().basic_auth("user", "pw");
2376        let header = config.headers.get("Authorization").unwrap();
2377        assert!(header.starts_with("Basic "));
2378    }
2379
2380    #[test]
2381    fn sep_2243_encodes_only_unsafe_values() {
2382        assert_eq!(encode_header_value("us west 1"), "us west 1");
2383        assert_eq!(encode_header_value(""), "");
2384        assert_eq!(encode_header_value(" padded "), "=?base64?IHBhZGRlZCA=?=");
2385        assert_eq!(
2386            encode_header_value("Hello, 世界"),
2387            "=?base64?SGVsbG8sIOS4lueVjA==?="
2388        );
2389    }
2390
2391    #[test]
2392    fn oauth_error_body_is_not_misclassified_as_jsonrpc() {
2393        assert!(!is_jsonrpc_error_response(&serde_json::json!({
2394            "error": "insufficient_scope",
2395            "error_description": "Token has insufficient scope"
2396        })));
2397        assert!(is_jsonrpc_error_response(&serde_json::json!({
2398            "jsonrpc": "2.0",
2399            "id": 1,
2400            "error": {
2401                "code": -32022,
2402                "message": "Unsupported protocol version"
2403            }
2404        })));
2405    }
2406
2407    #[test]
2408    fn sep_2243_validates_custom_header_annotations() {
2409        let mappings = custom_header_mappings(&serde_json::json!({
2410            "type": "object",
2411            "properties": {
2412                "region": {"type": "string", "x-mcp-header": "Region"},
2413                "priority": {"type": "integer", "x-mcp-header": "Priority"},
2414                "ratio": {"type": "number", "x-mcp-header": "Ratio"}
2415            }
2416        }))
2417        .unwrap();
2418        assert_eq!(mappings.len(), 3);
2419
2420        for invalid in [
2421            serde_json::json!({
2422                "type": "object",
2423                "properties": {"value": {"type": "object", "x-mcp-header": "Value"}}
2424            }),
2425            serde_json::json!({
2426                "type": "object",
2427                "properties": {
2428                    "a": {"type": "string", "x-mcp-header": "Region"},
2429                    "b": {"type": "string", "x-mcp-header": "region"}
2430                }
2431            }),
2432            serde_json::json!({
2433                "type": "object",
2434                "properties": {"value": {"type": "string", "x-mcp-header": "Bad Header"}}
2435            }),
2436        ] {
2437            assert!(custom_header_mappings(&invalid).is_err());
2438        }
2439    }
2440
2441    #[test]
2442    fn sep_2243_filters_invalid_tools_and_caches_valid_mappings() {
2443        let mut transport = HttpClientTransport::new("http://localhost:3000");
2444        transport.protocol_version = Some(crate::protocol::PROTOCOL_VERSION_2026_07_28.to_string());
2445        let normalized = transport.normalize_incoming_message(
2446            serde_json::json!({
2447                "jsonrpc": "2.0",
2448                "id": 1,
2449                "result": {
2450                    "tools": [
2451                        {
2452                            "name": "valid",
2453                            "inputSchema": {
2454                                "type": "object",
2455                                "properties": {
2456                                    "region": {"type": "string", "x-mcp-header": "Region"}
2457                                }
2458                            }
2459                        },
2460                        {
2461                            "name": "invalid",
2462                            "inputSchema": {
2463                                "type": "object",
2464                                "properties": {
2465                                    "value": {"type": "array", "x-mcp-header": "Value"}
2466                                }
2467                            }
2468                        }
2469                    ]
2470                }
2471            })
2472            .to_string(),
2473        );
2474        let parsed: serde_json::Value = serde_json::from_str(&normalized).unwrap();
2475        assert_eq!(parsed["result"]["tools"].as_array().unwrap().len(), 1);
2476        assert!(transport.tool_header_mappings.contains_key("valid"));
2477        assert!(!transport.tool_header_mappings.contains_key("invalid"));
2478    }
2479}