| 1 | use std::collections::VecDeque; |
| 2 | |
| 3 | use anyhow::{Context, Result}; |
| 4 | use reqwest::StatusCode; |
| 5 | use reqwest::header::CONTENT_TYPE; |
| 6 | |
| 7 | use super::http_client::McpHttpClient; |
| 8 | use super::wire::{ |
| 9 | MAX_MCP_RESPONSE_BYTES, is_streamable_http_incompatible_status, |
| 10 | is_streamable_http_stale_session_status, parse_sse_message_data, |
| 11 | }; |
| 12 | use super::{ERROR_BODY_PREVIEW_BYTES, bounded_body_excerpt, mask_url_secrets}; |
| 13 | |
| 14 | pub(super) struct StreamableHttpTransport { |
| 15 | pub(super) client: McpHttpClient, |
| 16 | pub(super) url: String, |
| 17 | pending_messages: VecDeque<Vec<u8>>, |
| 18 | /// Per-spec MCP session identifier returned by the server in the |
| 19 | /// first response (typically the `initialize` response). Attached |
| 20 | /// as the `Mcp-Session-Id` header on every subsequent outbound |
| 21 | /// request so the server can correlate messages within the same |
| 22 | /// session. |
| 23 | pub(super) session_id: Option<String>, |
| 24 | /// Protocol revision negotiated at `initialize`. Attached as the |
| 25 | /// `MCP-Protocol-Version` header on every subsequent outbound request |
| 26 | /// per the Streamable HTTP spec (absent means the server assumes |
| 27 | /// the 2025-03-26 default, so the negotiated value is always sent). |
| 28 | protocol_version: Option<String>, |
| 29 | } |
| 30 | |
| 31 | #[derive(Debug)] |
| 32 | pub(super) enum StreamableSendError { |
| 33 | Incompatible(String), |
| 34 | StaleSession(String), |
| 35 | Other(anyhow::Error), |
| 36 | } |
| 37 | |
| 38 | impl StreamableHttpTransport { |
| 39 | pub(super) fn new(client: McpHttpClient, url: String) -> Self { |
| 40 | Self { |
| 41 | client, |
| 42 | url, |
| 43 | pending_messages: VecDeque::new(), |
| 44 | session_id: None, |
| 45 | protocol_version: None, |
| 46 | } |
| 47 | } |
| 48 | |
| 49 | pub(super) fn set_protocol_version(&mut self, version: &str) { |
| 50 | self.protocol_version = Some(version.to_string()); |
| 51 | } |
| 52 | |
| 53 | pub(super) async fn send( |
| 54 | &mut self, |
| 55 | msg: Vec<u8>, |
| 56 | ) -> std::result::Result<(), StreamableSendError> { |
| 57 | let mut request = self.client.post(&self.url).body(msg); |
| 58 | if let Some(ref sid) = self.session_id { |
| 59 | request = request.header("Mcp-Session-Id", sid.as_str()); |
| 60 | } |
| 61 | if let Some(ref version) = self.protocol_version { |
| 62 | request = request.header("MCP-Protocol-Version", version.as_str()); |
| 63 | } |
| 64 | let client = self.client.clone(); |
| 65 | let response = client |
| 66 | .send_mcp_request( |
| 67 | request, |
| 68 | true, |
| 69 | false, |
| 70 | true, |
| 71 | || Ok(()), |
| 72 | |response| { |
| 73 | if let Some(sid) = response |
| 74 | .headers() |
| 75 | .get("Mcp-Session-Id") |
| 76 | .and_then(|v| v.to_str().ok()) |
| 77 | { |
| 78 | self.session_id = Some(sid.to_string()); |
| 79 | } |
| 80 | Ok(()) |
| 81 | }, |
| 82 | ) |
| 83 | .await |
| 84 | .map_err(StreamableSendError::Other)?; |
| 85 | let status = response.status(); |
| 86 | if status == StatusCode::ACCEPTED || status == StatusCode::NO_CONTENT { |
| 87 | return Ok(()); |
| 88 | } |
| 89 | if !status.is_success() { |
| 90 | let body_excerpt = bounded_body_excerpt(response, ERROR_BODY_PREVIEW_BYTES).await; |
| 91 | let stale_session = self.session_id.is_some() |
| 92 | && is_streamable_http_stale_session_status(status, &body_excerpt); |
| 93 | let body_excerpt = self.client.server_error_preview(&body_excerpt); |
| 94 | if stale_session { |
| 95 | return Err(StreamableSendError::StaleSession(format!( |
| 96 | "status={status} body={body_excerpt}" |
| 97 | ))); |
| 98 | } |
| 99 | if is_streamable_http_incompatible_status(status) { |
| 100 | return Err(StreamableSendError::Incompatible(format!( |
| 101 | "status={status} body={body_excerpt}" |
| 102 | ))); |
| 103 | } |
| 104 | return Err(StreamableSendError::Other(anyhow::anyhow!( |
| 105 | "MCP Streamable HTTP rejected (transport=http url={} status={}): {}", |
| 106 | mask_url_secrets(&self.url), |
| 107 | status, |
| 108 | body_excerpt, |
| 109 | ))); |
| 110 | } |
| 111 | |
| 112 | let content_type = response |
| 113 | .headers() |
| 114 | .get(CONTENT_TYPE) |
| 115 | .and_then(|value| value.to_str().ok()) |
| 116 | .map(str::to_string); |
| 117 | // Reject an over-large declared body before reading anything (fast |
| 118 | // path), then bound the read itself so chunked / length-less |
| 119 | // responses cannot OOM us either — Content-Length alone does not |
| 120 | // protect against a server that streams without declaring a length. |
| 121 | if let Some(len) = response.content_length() |
| 122 | && len > MAX_MCP_RESPONSE_BYTES as u64 |
| 123 | { |
| 124 | return Err(StreamableSendError::Other(anyhow::anyhow!( |
| 125 | "MCP response Content-Length {len} exceeds {} bytes — aborting", |
| 126 | MAX_MCP_RESPONSE_BYTES |
| 127 | ))); |
| 128 | } |
| 129 | let body = read_body_capped(response, MAX_MCP_RESPONSE_BYTES) |
| 130 | .await |
| 131 | .map_err(StreamableSendError::Other)?; |
| 132 | self.store_response_body(content_type.as_deref(), &body) |
| 133 | .map_err(StreamableSendError::Other) |
| 134 | } |
| 135 | |
| 136 | pub(super) async fn recv(&mut self) -> Result<Vec<u8>> { |
| 137 | self.pending_messages |
| 138 | .pop_front() |
| 139 | .context("MCP Streamable HTTP response queue is empty") |
| 140 | } |
| 141 | |
| 142 | fn store_response_body(&mut self, content_type: Option<&str>, body: &str) -> Result<()> { |
| 143 | if body.trim().is_empty() { |
| 144 | return Ok(()); |
| 145 | } |
| 146 | |
| 147 | let is_event_stream = content_type |
| 148 | .map(|value| value.to_ascii_lowercase().contains("text/event-stream")) |
| 149 | .unwrap_or(false) |
| 150 | || body.trim_start().starts_with("event:") |
| 151 | || body.trim_start().starts_with("data:"); |
| 152 | |
| 153 | if is_event_stream { |
| 154 | for msg in parse_sse_message_data(body) { |
| 155 | self.pending_messages.push_back(msg); |
| 156 | } |
| 157 | return Ok(()); |
| 158 | } |
| 159 | |
| 160 | self.pending_messages.push_back(body.as_bytes().to_vec()); |
| 161 | Ok(()) |
| 162 | } |
| 163 | } |
| 164 | |
| 165 | /// Read a response body through the byte stream, failing as soon as it |
| 166 | /// exceeds `max_bytes`. This bounds chunked and missing-Content-Length |
| 167 | /// responses exactly like declared ones (the declared-length fast path in |
| 168 | /// `send` only covers servers honest enough to announce their size). |
| 169 | /// MCP bodies are JSON or SSE, so lossy UTF-8 matches `.text()` behavior. |
| 170 | pub(super) async fn read_body_capped( |
| 171 | response: reqwest::Response, |
| 172 | max_bytes: usize, |
| 173 | ) -> Result<String> { |
| 174 | let buf = crate::utils::read_response_body_capped(response, max_bytes) |
| 175 | .await |
| 176 | .map_err(|error| anyhow::anyhow!("MCP {error:#}"))?; |
| 177 | Ok(String::from_utf8_lossy(&buf).into_owned()) |
| 178 | } |
| 179 | |
| 180 | /// TUI recovery for a rejected OAuth session. Settings recovery and the |
| 181 | /// Streamable HTTP path share the oauth helpers so `/mcp login <name>` stays |
| 182 | /// the only advertised command; `/mcp auth` is not a command. |
| 183 | #[cfg(test)] |
| 184 | fn oauth_refresh_failed_hint() -> &'static str { |
| 185 | super::oauth::tui_reauth_refresh_failed_hint() |
| 186 | } |
| 187 | |
| 188 | /// TUI recovery for a rejected OAuth session. `oauth_configured` is the |
| 189 | /// server's configured auth path ([`McpHttpClient::oauth_configured`]), not the |
| 190 | /// presence of a cached token, so a first-run OAuth server — a 401 with |
| 191 | /// nothing stored yet — is still pointed at `/mcp login <name>` rather than at |
| 192 | /// a bearer token it never had (#6030). Servers where a bearer credential is |
| 193 | /// genuinely configured (or that are plugin-contributed, where OAuth login is |
| 194 | /// disabled) keep the bearer-token copy. |
| 195 | #[cfg(test)] |
| 196 | fn unauthorized_session_hint(oauth_configured: bool) -> &'static str { |
| 197 | if oauth_configured { |
| 198 | super::oauth::tui_reauth_hint() |
| 199 | } else { |
| 200 | "Check the configured bearer token (or its environment variable)." |
| 201 | } |
| 202 | } |
| 203 | |
| 204 | #[cfg(test)] |
| 205 | mod tests { |
| 206 | use super::{oauth_refresh_failed_hint, unauthorized_session_hint}; |
| 207 | use crate::mcp::McpServerConfig; |
| 208 | use crate::mcp::http_client::McpHttpAuth; |
| 209 | |
| 210 | fn server_config(json: serde_json::Value) -> McpServerConfig { |
| 211 | serde_json::from_value(json).expect("MCP server config fixture") |
| 212 | } |
| 213 | |
| 214 | #[test] |
| 215 | fn oauth_configured_server_without_a_cached_token_names_login() { |
| 216 | // The OAuth fields are optional in MCP config, so a URL-based server |
| 217 | // with no manual bearer configuration is OAuth's to claim — including |
| 218 | // before the first login, when there is no runtime to observe. |
| 219 | let auth = McpHttpAuth::from_config( |
| 220 | "remote", |
| 221 | &server_config(serde_json::json!({ "url": "https://example.invalid/mcp" })), |
| 222 | None, |
| 223 | ); |
| 224 | assert!(auth.oauth.is_none(), "precondition: no cached credential"); |
| 225 | assert!(auth.oauth_configured, "a URL server is OAuth-servable"); |
| 226 | assert!( |
| 227 | unauthorized_session_hint(auth.oauth_configured).contains("/mcp login <name>"), |
| 228 | "a first-run OAuth 401 must name the login command, not a bearer token" |
| 229 | ); |
| 230 | |
| 231 | // A server whose bearer token is genuinely expected keeps that copy. |
| 232 | let bearer = McpHttpAuth::from_config( |
| 233 | "remote", |
| 234 | &server_config(serde_json::json!({ |
| 235 | "url": "https://example.invalid/mcp", |
| 236 | "bearer_token_env_var": "EXAMPLE_MCP_TOKEN", |
| 237 | })), |
| 238 | None, |
| 239 | ); |
| 240 | assert!(!bearer.oauth_configured); |
| 241 | assert!(unauthorized_session_hint(bearer.oauth_configured).contains("bearer token")); |
| 242 | } |
| 243 | |
| 244 | #[test] |
| 245 | fn unauthorized_oauth_hints_name_the_login_command() { |
| 246 | for hint in [oauth_refresh_failed_hint(), unauthorized_session_hint(true)] { |
| 247 | assert!( |
| 248 | hint.contains("/mcp login <name>"), |
| 249 | "OAuth recovery must name the implemented command" |
| 250 | ); |
| 251 | assert!( |
| 252 | !hint.contains("/mcp auth"), |
| 253 | "OAuth recovery must not advertise a missing /mcp auth command" |
| 254 | ); |
| 255 | } |
| 256 | assert!( |
| 257 | !unauthorized_session_hint(false).contains("/mcp"), |
| 258 | "bearer-token recovery should not send the user to OAuth login" |
| 259 | ); |
| 260 | } |
| 261 | } |
| 262 |