diff --git a/src/config/mcp_factory.rs b/src/config/mcp_factory.rs index 4a46a35..0afc715 100644 --- a/src/config/mcp_factory.rs +++ b/src/config/mcp_factory.rs @@ -1,7 +1,6 @@ -use crate::mcp::oauth::McpTokenStatus; use crate::mcp::{ - ConnectedServer, JsonField, McpAuthReason, McpAuthRequired, McpServer, McpTransportType, - is_auth_required_error, oauth, spawn_mcp_server, + ConnectedServer, JsonField, McpAuthRequired, McpServer, McpTransportType, + is_auth_required_error, resolve_http_auth, spawn_mcp_server, }; use anyhow::Result; @@ -103,13 +102,8 @@ impl McpFactory { return Ok(existing); } - let token_status = if spec.is_remote() { - oauth::load_or_refresh_mcp_token(name).await - } else { - McpTokenStatus::NotAuthenticated - }; - let auth_reason = McpAuthReason::from_token_status(&token_status); - let handle = spawn_mcp_server(spec, log_path, token_status.into_token()) + let (auth, auth_reason) = resolve_http_auth(name, spec).await; + let handle = spawn_mcp_server(spec, log_path, auth) .await .map_err(|e| { if is_auth_required_error(&e) { diff --git a/src/mcp/auth_client.rs b/src/mcp/auth_client.rs new file mode 100644 index 0000000..d337f23 --- /dev/null +++ b/src/mcp/auth_client.rs @@ -0,0 +1,651 @@ +use crate::mcp::oauth::{force_refresh_mcp_token, load_or_refresh_mcp_token}; +use http::{HeaderName, HeaderValue}; +use log::debug; +use rmcp::model::ClientJsonRpcMessage; +use rmcp::transport::common::client_side_sse::BoxedSseResponse; +use rmcp::transport::streamable_http_client::{ + AuthRequiredError, StreamableHttpClient, StreamableHttpError, StreamableHttpPostResponse, +}; +use std::collections::HashMap; +use std::future::Future; +use std::sync::Arc; + +/// [`StreamableHttpClient`] wrapper that injects the OAuth bearer token for an +/// MCP server on every request instead of pinning it at spawn time, so tokens +/// refreshed mid-session take effect without reconnecting. +/// +/// A caller-supplied `auth_header` always passes through untouched; only a +/// `None` header is filled from the stored token. When the wrapper injected +/// the token and a POST comes back 401, it forces a token refresh and retries +/// exactly once (see [`Self::post_with_retry`]). +#[derive(Clone)] +pub struct McpOAuthClient { + inner: C, + server: Arc, +} + +impl McpOAuthClient { + pub fn new(inner: C, server: &str) -> Self { + Self { + inner, + server: Arc::from(server), + } + } +} + +impl McpOAuthClient { + /// Resolves the effective auth header. Caller-supplied values pass through + /// untouched; `None` is filled from the stored token for this server. + /// Returns the header plus whether the wrapper injected it. Errors with + /// [`StreamableHttpError::AuthRequired`] (without contacting the server) + /// when no usable token exists. + async fn resolve_auth( + &self, + auth_header: Option, + ) -> Result<(Option, bool), StreamableHttpError> { + if auth_header.is_some() { + return Ok((auth_header, false)); + } + match load_or_refresh_mcp_token(&self.server).await.into_token() { + Some(token) => Ok((Some(token), true)), + None => Err(self.auth_required()), + } + } + + fn auth_required(&self) -> StreamableHttpError { + StreamableHttpError::AuthRequired(AuthRequiredError::new(format!( + "no valid OAuth token for MCP server '{server}'; \ + run `.mcp auth {server}` to re-authenticate", + server = self.server + ))) + } + + async fn post_with_retry( + &self, + auth_header: Option, + mut post: F, + ) -> Result> + where + F: FnMut(Option) -> Fut, + Fut: Future>>, + { + let (auth, injected) = self.resolve_auth(auth_header).await?; + let (original, rejected) = match (post(auth.clone()).await, auth) { + (Err(err @ StreamableHttpError::AuthRequired(_)), Some(rejected)) if injected => { + (err, rejected) + } + (result, _) => return result, + }; + + debug!( + "MCP server '{}' rejected the injected token; forcing a refresh and retrying once", + self.server + ); + + let Some(token) = force_refresh_mcp_token(&self.server, &rejected).await else { + return Err(original); + }; + + match post(Some(token)).await { + Err(StreamableHttpError::AuthRequired(_)) => { + debug!( + "Retry after forced token refresh was rejected again by MCP server '{}'", + self.server + ); + Err(original) + } + result => result, + } + } +} + +impl StreamableHttpClient for McpOAuthClient { + type Error = C::Error; + + async fn post_message( + &self, + uri: Arc, + message: ClientJsonRpcMessage, + session_id: Option>, + auth_header: Option, + custom_headers: HashMap, + ) -> Result> { + self.post_with_retry(auth_header, |auth| { + self.inner.post_message( + uri.clone(), + message.clone(), + session_id.clone(), + auth, + custom_headers.clone(), + ) + }) + .await + } + + /// Overridden rather than left to the trait default: the default impl + /// delegates to [`Self::post_message`], silently dropping the + /// transport-wide SSE event size limit. Delegating to the inner client's + /// size-enforcing variant keeps the limit applied at the raw byte layer. + async fn post_message_with_max_sse_event_size( + &self, + uri: Arc, + message: ClientJsonRpcMessage, + session_id: Option>, + auth_header: Option, + custom_headers: HashMap, + max_sse_event_size: usize, + ) -> Result> { + self.post_with_retry(auth_header, |auth| { + self.inner.post_message_with_max_sse_event_size( + uri.clone(), + message.clone(), + session_id.clone(), + auth, + custom_headers.clone(), + max_sse_event_size, + ) + }) + .await + } + + async fn delete_session( + &self, + uri: Arc, + session_id: Arc, + auth_header: Option, + custom_headers: HashMap, + ) -> Result<(), StreamableHttpError> { + let (auth, _) = self.resolve_auth(auth_header).await?; + self.inner + .delete_session(uri, session_id, auth, custom_headers) + .await + } + + async fn get_stream( + &self, + uri: Arc, + session_id: Option>, + last_event_id: Option, + auth_header: Option, + custom_headers: HashMap, + ) -> Result> { + let (auth, _) = self.resolve_auth(auth_header).await?; + self.inner + .get_stream(uri, session_id, last_event_id, auth, custom_headers) + .await + } + + /// Overridden for the same reason as + /// [`Self::post_message_with_max_sse_event_size`]: the trait default + /// bypasses SSE event size enforcement. + async fn get_stream_with_max_sse_event_size( + &self, + uri: Arc, + session_id: Option>, + last_event_id: Option, + auth_header: Option, + custom_headers: HashMap, + max_sse_event_size: usize, + ) -> Result> { + let (auth, _) = self.resolve_auth(auth_header).await?; + self.inner + .get_stream_with_max_sse_event_size( + uri, + session_id, + last_event_id, + auth, + custom_headers, + max_sse_event_size, + ) + .await + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::paths; + use crate::mcp::oauth::test_support::with_temp_cache; + use futures_util::StreamExt; + use parking_lot::Mutex; + use serial_test::serial; + use std::convert::Infallible; + use std::fs; + use std::sync::atomic::{AtomicUsize, Ordering}; + + const FRESH: i64 = 9999999999; + + type Calls = Arc)>>>; + + type AfterFirstCallAction = Option>; + + /// Inner client that records `(method, auth_header)` per call, rejects the + /// first `reject_times` POSTs/streams with `AuthRequired` (then the next + /// `transport_error_times` with a non-auth error), and runs an optional + /// side effect after the first call (to mutate token files between the + /// initial attempt and the retry). + #[derive(Clone, Default)] + struct FakeInner { + calls: Calls, + reject_times: Arc, + transport_error_times: Arc, + after_first_call: Arc>, + } + + impl FakeInner { + fn record( + &self, + method: &'static str, + auth: Option, + ) -> Result<(), StreamableHttpError> { + self.calls.lock().push((method, auth)); + if let Some(f) = self.after_first_call.lock().take() { + f(); + } + if self.reject_times.load(Ordering::SeqCst) > 0 { + self.reject_times.fetch_sub(1, Ordering::SeqCst); + return Err(Self::rejection()); + } + if self.transport_error_times.load(Ordering::SeqCst) > 0 { + self.transport_error_times.fetch_sub(1, Ordering::SeqCst); + return Err(StreamableHttpError::UnexpectedServerResponse( + "connection reset".into(), + )); + } + Ok(()) + } + + fn rejection() -> StreamableHttpError { + StreamableHttpError::AuthRequired(AuthRequiredError::new( + "Bearer error=\"invalid_token\"".to_string(), + )) + } + } + + impl StreamableHttpClient for FakeInner { + type Error = Infallible; + + async fn post_message( + &self, + _uri: Arc, + _message: ClientJsonRpcMessage, + _session_id: Option>, + auth_header: Option, + _custom_headers: HashMap, + ) -> Result> { + self.record("post_message", auth_header)?; + Ok(StreamableHttpPostResponse::Accepted) + } + + async fn post_message_with_max_sse_event_size( + &self, + _uri: Arc, + _message: ClientJsonRpcMessage, + _session_id: Option>, + auth_header: Option, + _custom_headers: HashMap, + _max_sse_event_size: usize, + ) -> Result> { + self.record("post_message_with_max_sse_event_size", auth_header)?; + Ok(StreamableHttpPostResponse::Accepted) + } + + async fn delete_session( + &self, + _uri: Arc, + _session_id: Arc, + auth_header: Option, + _custom_headers: HashMap, + ) -> Result<(), StreamableHttpError> { + self.record("delete_session", auth_header)?; + Ok(()) + } + + async fn get_stream( + &self, + _uri: Arc, + _session_id: Option>, + _last_event_id: Option, + auth_header: Option, + _custom_headers: HashMap, + ) -> Result> { + self.record("get_stream", auth_header)?; + Ok(futures_util::stream::empty().boxed()) + } + + async fn get_stream_with_max_sse_event_size( + &self, + _uri: Arc, + _session_id: Option>, + _last_event_id: Option, + auth_header: Option, + _custom_headers: HashMap, + _max_sse_event_size: usize, + ) -> Result> { + self.record("get_stream_with_max_sse_event_size", auth_header)?; + Ok(futures_util::stream::empty().boxed()) + } + } + + fn write_token_file(server: &str, access_token: &str, expires_at: i64) { + fs::create_dir_all(paths::oauth_tokens_dir()).unwrap(); + fs::write( + paths::token_file(&format!("mcp_{server}")), + format!( + r#"{{"access_token":"{access_token}","refresh_token":"r","expires_at":{expires_at}}}"# + ), + ) + .unwrap(); + } + + fn ping() -> ClientJsonRpcMessage { + serde_json::from_value(serde_json::json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "ping" + })) + .unwrap() + } + + fn rt() -> tokio::runtime::Runtime { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap() + } + + fn post( + client: &McpOAuthClient, + auth_header: Option, + ) -> Result> { + rt().block_on(client.post_message( + Arc::from("http://mcp.test/mcp"), + ping(), + None, + auth_header, + HashMap::new(), + )) + } + + #[test] + #[serial] + fn injects_token_from_disk_when_auth_header_none() { + with_temp_cache(|| { + write_token_file("wrapper-inject", "tok-live", FRESH); + let inner = FakeInner::default(); + let client = McpOAuthClient::new(inner.clone(), "wrapper-inject"); + + let result = post(&client, None); + + assert!(result.is_ok()); + assert_eq!( + *inner.calls.lock(), + vec![("post_message", Some("tok-live".to_string()))] + ); + }); + } + + #[test] + #[serial] + fn caller_supplied_auth_header_passes_through() { + with_temp_cache(|| { + write_token_file("wrapper-passthrough", "tok-disk", FRESH); + let inner = FakeInner::default(); + let client = McpOAuthClient::new(inner.clone(), "wrapper-passthrough"); + + let result = post(&client, Some("caller-tok".to_string())); + + assert!(result.is_ok()); + assert_eq!( + *inner.calls.lock(), + vec![("post_message", Some("caller-tok".to_string()))] + ); + }); + } + + #[test] + #[serial] + fn caller_supplied_header_rejection_propagates_without_refresh() { + with_temp_cache(|| { + // A fresh, different token sits on disk: if the injected guard + // were dropped, the wrapper would refresh and retry with it. + write_token_file("wrapper-caller-401", "tok-disk", FRESH); + let inner = FakeInner::default(); + inner.reject_times.store(1, Ordering::SeqCst); + let client = McpOAuthClient::new(inner.clone(), "wrapper-caller-401"); + + let result = post(&client, Some("caller-tok".to_string())); + + assert!(matches!(result, Err(StreamableHttpError::AuthRequired(_)))); + assert_eq!( + *inner.calls.lock(), + vec![("post_message", Some("caller-tok".to_string()))] + ); + }); + } + + #[test] + #[serial] + fn missing_token_returns_auth_required_without_calling_inner() { + with_temp_cache(|| { + let inner = FakeInner::default(); + let client = McpOAuthClient::new(inner.clone(), "wrapper-no-token"); + + let result = post(&client, None); + + assert!(matches!(result, Err(StreamableHttpError::AuthRequired(_)))); + assert!(inner.calls.lock().is_empty()); + }); + } + + #[test] + #[serial] + fn rejected_token_forces_refresh_and_retries_once() { + with_temp_cache(|| { + write_token_file("wrapper-retry", "tok-a", FRESH); + let inner = FakeInner::default(); + inner.reject_times.store(1, Ordering::SeqCst); + // Simulate a concurrent refresh landing between the rejection and + // the forced refresh: the retry must carry the new token. + *inner.after_first_call.lock() = Some(Box::new(|| { + write_token_file("wrapper-retry", "tok-b", FRESH); + })); + let client = McpOAuthClient::new(inner.clone(), "wrapper-retry"); + + let result = post(&client, None); + + assert!(result.is_ok()); + assert_eq!( + *inner.calls.lock(), + vec![ + ("post_message", Some("tok-a".to_string())), + ("post_message", Some("tok-b".to_string())), + ] + ); + }); + } + + #[test] + #[serial] + fn failed_force_refresh_propagates_original_error_after_one_call() { + with_temp_cache(|| { + write_token_file("wrapper-refresh-fail", "tok-a", FRESH); + let inner = FakeInner::default(); + inner.reject_times.store(1, Ordering::SeqCst); + // Token file gone by refresh time: force refresh yields nothing. + *inner.after_first_call.lock() = Some(Box::new(|| { + fs::remove_file(paths::token_file("mcp_wrapper-refresh-fail")).unwrap(); + })); + let client = McpOAuthClient::new(inner.clone(), "wrapper-refresh-fail"); + + let result = post(&client, None); + + assert!(matches!(result, Err(StreamableHttpError::AuthRequired(_)))); + assert_eq!(inner.calls.lock().len(), 1); + }); + } + + #[test] + #[serial] + fn second_rejection_after_retry_propagates_original_error() { + with_temp_cache(|| { + write_token_file("wrapper-double-401", "tok-a", FRESH); + let inner = FakeInner::default(); + inner.reject_times.store(2, Ordering::SeqCst); + // A changed token appears before the forced refresh, so the retry + // actually runs (an unchanged token would trigger a real refresh + // attempt, which fails without a cached registration). + *inner.after_first_call.lock() = Some(Box::new(|| { + write_token_file("wrapper-double-401", "tok-b", FRESH); + })); + let client = McpOAuthClient::new(inner.clone(), "wrapper-double-401"); + + let result = post(&client, None); + + assert!(matches!(result, Err(StreamableHttpError::AuthRequired(_)))); + assert_eq!(inner.calls.lock().len(), 2); + }); + } + + #[test] + #[serial] + fn non_auth_retry_error_propagates_as_is() { + with_temp_cache(|| { + write_token_file("wrapper-retry-transport", "tok-a", FRESH); + let inner = FakeInner::default(); + inner.reject_times.store(1, Ordering::SeqCst); + inner.transport_error_times.store(1, Ordering::SeqCst); + *inner.after_first_call.lock() = Some(Box::new(|| { + write_token_file("wrapper-retry-transport", "tok-b", FRESH); + })); + let client = McpOAuthClient::new(inner.clone(), "wrapper-retry-transport"); + + let result = post(&client, None); + + assert!(matches!( + result, + Err(StreamableHttpError::UnexpectedServerResponse(_)) + )); + assert_eq!(inner.calls.lock().len(), 2); + }); + } + + #[test] + #[serial] + fn sized_post_delegates_to_inner_sized_variant() { + with_temp_cache(|| { + write_token_file("wrapper-sized-post", "tok-live", FRESH); + let inner = FakeInner::default(); + let client = McpOAuthClient::new(inner.clone(), "wrapper-sized-post"); + + let result = rt().block_on(client.post_message_with_max_sse_event_size( + Arc::from("http://mcp.test/mcp"), + ping(), + None, + None, + HashMap::new(), + 4096, + )); + + assert!(result.is_ok()); + assert_eq!( + *inner.calls.lock(), + vec![( + "post_message_with_max_sse_event_size", + Some("tok-live".to_string()) + )] + ); + }); + } + + #[test] + #[serial] + fn sized_get_stream_delegates_to_inner_sized_variant() { + with_temp_cache(|| { + write_token_file("wrapper-sized-get", "tok-live", FRESH); + let inner = FakeInner::default(); + let client = McpOAuthClient::new(inner.clone(), "wrapper-sized-get"); + + let result = rt().block_on(client.get_stream_with_max_sse_event_size( + Arc::from("http://mcp.test/mcp"), + None, + None, + None, + HashMap::new(), + 4096, + )); + + assert!(result.is_ok()); + assert_eq!( + *inner.calls.lock(), + vec![( + "get_stream_with_max_sse_event_size", + Some("tok-live".to_string()) + )] + ); + }); + } + + #[test] + #[serial] + fn get_stream_does_not_retry_on_rejection() { + with_temp_cache(|| { + write_token_file("wrapper-get-401", "tok-live", FRESH); + let inner = FakeInner::default(); + inner.reject_times.store(1, Ordering::SeqCst); + let client = McpOAuthClient::new(inner.clone(), "wrapper-get-401"); + + let result = rt().block_on(client.get_stream( + Arc::from("http://mcp.test/mcp"), + None, + None, + None, + HashMap::new(), + )); + + assert!(matches!(result, Err(StreamableHttpError::AuthRequired(_)))); + assert_eq!(inner.calls.lock().len(), 1); + }); + } + + #[test] + #[serial] + fn delete_session_injects_token() { + with_temp_cache(|| { + write_token_file("wrapper-delete", "tok-live", FRESH); + let inner = FakeInner::default(); + let client = McpOAuthClient::new(inner.clone(), "wrapper-delete"); + + let result = rt().block_on(client.delete_session( + Arc::from("http://mcp.test/mcp"), + Arc::from("session-1"), + None, + HashMap::new(), + )); + + assert!(result.is_ok()); + assert_eq!( + *inner.calls.lock(), + vec![("delete_session", Some("tok-live".to_string()))] + ); + }); + } + + #[test] + #[serial] + fn auth_required_error_contains_no_token_material() { + with_temp_cache(|| { + write_token_file("wrapper-redact", "stale-secret-token", 0); + let inner = FakeInner::default(); + let client = McpOAuthClient::new(inner.clone(), "wrapper-redact"); + + let err = post(&client, None).unwrap_err(); + + let display = format!("{err}"); + let debug = format!("{err:?}"); + assert!(!display.contains("stale-secret-token")); + assert!(!debug.contains("stale-secret-token")); + assert!(inner.calls.lock().is_empty()); + }); + } +} diff --git a/src/mcp/mod.rs b/src/mcp/mod.rs index 4b8baf4..e6b6222 100644 --- a/src/mcp/mod.rs +++ b/src/mcp/mod.rs @@ -1,3 +1,4 @@ +mod auth_client; pub(crate) mod manage; pub(crate) mod oauth; mod sse_transport; @@ -9,6 +10,7 @@ use crate::vault::Vault; use crate::vault::interpolate_secrets; use anyhow::Error; use anyhow::{Context, Result, anyhow}; +use auth_client::McpOAuthClient; use futures_util::{StreamExt, TryStreamExt, stream}; use http::{HeaderName, HeaderValue}; use indexmap::IndexMap; @@ -328,29 +330,22 @@ impl McpRegistry { .and_then(|c| c.mcp_servers.get(&id)) .with_context(|| format!("MCP server not found in config: {id}"))?; - let token_status = if spec.is_remote() { - oauth::load_or_refresh_mcp_token(&id).await - } else { - oauth::McpTokenStatus::NotAuthenticated - }; - let auth_reason = McpAuthReason::from_token_status(&token_status); + let (auth, auth_reason) = resolve_http_auth(&id, spec).await; - let service = - match spawn_mcp_server(spec, self.log_path.as_deref(), token_status.into_token()).await - { - Ok(s) => s, - Err(e) if is_auth_required_error(&e) => { - warn!( - "{}", - McpAuthRequired { - server: id, - reason: auth_reason, - } - ); - return Ok(None); - } - Err(e) => return Err(e), - }; + let service = match spawn_mcp_server(spec, self.log_path.as_deref(), auth).await { + Ok(s) => s, + Err(e) if is_auth_required_error(&e) => { + warn!( + "{}", + McpAuthRequired { + server: id, + reason: auth_reason, + } + ); + return Ok(None); + } + Err(e) => return Err(e), + }; let tools = service.list_tools(None).await?; debug!("Available tools for MCP server {id}: {tools:?}"); @@ -420,19 +415,76 @@ impl McpRegistry { } } +/// How a remote MCP server authenticates outgoing requests. +pub(crate) enum HttpAuth { + /// Only the static headers from the server spec; no OAuth token. + StaticOnly, + /// OAuth-managed: HTTP transports inject a fresh bearer token per request + /// via [`McpOAuthClient`] (ignoring the carried token); SSE transports + /// send the carried token as a static header. + Managed { server: String, token: String }, +} + +impl fmt::Debug for HttpAuth { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::StaticOnly => f.write_str("StaticOnly"), + Self::Managed { server, token: _ } => f + .debug_struct("Managed") + .field("server", server) + .field("token", &"") + .finish(), + } + } +} + +impl HttpAuth { + pub(crate) fn from_token_status(status: &oauth::McpTokenStatus, server: &str) -> Self { + match status { + oauth::McpTokenStatus::Token(token) => Self::Managed { + server: server.to_string(), + token: token.clone(), + }, + oauth::McpTokenStatus::NotAuthenticated | oauth::McpTokenStatus::RefreshFailed => { + Self::StaticOnly + } + } + } +} + +pub(crate) async fn resolve_http_auth(name: &str, spec: &McpServer) -> (HttpAuth, McpAuthReason) { + let token_status = if spec.is_remote() { + oauth::load_or_refresh_mcp_token(name).await + } else { + oauth::McpTokenStatus::NotAuthenticated + }; + ( + HttpAuth::from_token_status(&token_status, name), + McpAuthReason::from_token_status(&token_status), + ) +} + pub(crate) async fn spawn_mcp_server( spec: &McpServer, log_path: Option<&Path>, - bearer_token: Option, + auth: HttpAuth, ) -> Result> { match spec.transport_type { McpTransportType::Http => { let url = spec.url.as_deref().expect("validated: http spec has url"); - let headers = merge_bearer_token(spec.headers.as_ref(), bearer_token); - spawn_http_mcp_server(url, headers.as_ref()).await + match auth { + HttpAuth::Managed { server, token: _ } => { + spawn_oauth_http_mcp_server(url, &server, spec.headers.as_ref()).await + } + HttpAuth::StaticOnly => spawn_http_mcp_server(url, spec.headers.as_ref()).await, + } } McpTransportType::Sse => { let url = spec.url.as_deref().expect("validated: sse spec has url"); + let bearer_token = match auth { + HttpAuth::Managed { server: _, token } => Some(token), + HttpAuth::StaticOnly => None, + }; let headers = merge_bearer_token(spec.headers.as_ref(), bearer_token); spawn_sse_mcp_server(url, headers.as_ref()).await } @@ -460,6 +512,7 @@ fn merge_bearer_token( } (Some(h), Some(token)) => { let mut m = h.clone(); + m.retain(|k, _| !k.eq_ignore_ascii_case("authorization")); m.insert("Authorization".to_string(), format!("Bearer {token}")); Some(m) } @@ -551,6 +604,66 @@ async fn spawn_http_mcp_server( Ok(service) } +/// Builds the custom-header map for an OAuth-managed HTTP transport, dropping +/// any static `Authorization` entry case-insensitively: [`McpOAuthClient`] +/// owns that header, and a stale configured value must not collide with the +/// per-request token. +fn oauth_custom_headers( + headers: Option<&IndexMap>, +) -> Result> { + let mut custom = HashMap::new(); + let Some(hdrs) = headers else { + return Ok(custom); + }; + + for (k, v) in hdrs { + if k.eq_ignore_ascii_case("authorization") { + continue; + } + let name = k + .parse::() + .with_context(|| format!("Invalid header name: {k}"))?; + let value = v + .parse::() + .with_context(|| format!("Invalid header value for {k}"))?; + custom.insert(name, value); + } + + Ok(custom) +} + +async fn spawn_oauth_http_mcp_server( + url: &str, + server: &str, + headers: Option<&IndexMap>, +) -> Result> { + // Mirror rmcp's default_http_client, which `with_client` bypasses: + // idle pooling off avoids a documented TCP delayed-ACK stall, and + // redirects off keeps custom headers from being replayed to a redirect + // target. + let inner = reqwest::Client::builder() + .pool_max_idle_per_host(0) + .redirect(reqwest::redirect::Policy::none()) + .build() + .context("Failed to build HTTP client for OAuth-managed MCP transport")?; + let client = McpOAuthClient::new(inner, server); + // `auth_header` stays None so the wrapper injects a fresh token per + // request; `reinit_on_expired_session` defaults to true in rmcp 3.1.2 + // but is pinned explicitly because transparent session re-init is + // load-bearing for long-lived sessions. + let config = StreamableHttpClientTransportConfig::with_uri(url) + .custom_headers(oauth_custom_headers(headers)?) + .reinit_on_expired_session(true); + let transport = StreamableHttpClientTransport::with_client(client, config); + let service = Arc::new( + ().serve(transport) + .await + .with_context(|| format!("Failed to connect to HTTP MCP server: {url}"))?, + ); + + Ok(service) +} + async fn spawn_sse_mcp_server( url: &str, headers: Option<&IndexMap>, @@ -1109,6 +1222,92 @@ mod tests { assert_eq!(result["X-Custom"], "keep"); } + #[test] + fn merge_bearer_token_replaces_authorization_case_insensitively() { + let mut h = IndexMap::new(); + h.insert("authorization".to_string(), "Bearer stale-1".to_string()); + h.insert("AUTHORIZATION".to_string(), "Bearer stale-2".to_string()); + h.insert("X-Custom".to_string(), "keep".to_string()); + + let result = merge_bearer_token(Some(&h), Some("newtoken".to_string())).unwrap(); + + assert_eq!(result.len(), 2); + assert_eq!(result["Authorization"], "Bearer newtoken"); + assert_eq!(result["X-Custom"], "keep"); + assert!(!result.contains_key("authorization")); + assert!(!result.contains_key("AUTHORIZATION")); + } + + #[test] + fn http_auth_from_token_status_maps_token_to_managed() { + assert!(matches!( + HttpAuth::from_token_status(&oauth::McpTokenStatus::Token("tok".into()), "srv"), + HttpAuth::Managed { server, token } if server == "srv" && token == "tok" + )); + assert!(matches!( + HttpAuth::from_token_status(&oauth::McpTokenStatus::NotAuthenticated, "srv"), + HttpAuth::StaticOnly + )); + assert!(matches!( + HttpAuth::from_token_status(&oauth::McpTokenStatus::RefreshFailed, "srv"), + HttpAuth::StaticOnly + )); + } + + #[test] + fn http_auth_debug_redacts_token() { + let auth = HttpAuth::Managed { + server: "srv".into(), + token: "live-secret".into(), + }; + + let debug = format!("{auth:?}"); + + assert!(debug.contains("srv")); + assert!(debug.contains("")); + assert!(!debug.contains("live-secret")); + } + + #[test] + fn oauth_custom_headers_strips_authorization_case_insensitively() { + let mut h = IndexMap::new(); + h.insert("Authorization".to_string(), "Bearer stale-1".to_string()); + h.insert("authorization".to_string(), "Bearer stale-2".to_string()); + h.insert("AUTHORIZATION".to_string(), "Bearer stale-3".to_string()); + h.insert("X-Custom".to_string(), "keep".to_string()); + + let custom = oauth_custom_headers(Some(&h)).unwrap(); + + assert_eq!(custom.len(), 1); + assert_eq!(custom[&HeaderName::from_static("x-custom")], "keep"); + } + + #[test] + fn oauth_custom_headers_none_is_empty() { + assert!(oauth_custom_headers(None).unwrap().is_empty()); + } + + #[test] + fn oauth_custom_headers_rejects_invalid_header_name() { + let mut h = IndexMap::new(); + h.insert("bad header".to_string(), "v".to_string()); + + assert!(oauth_custom_headers(Some(&h)).is_err()); + } + + #[test] + fn oauth_custom_headers_keeps_non_authorization_headers() { + let mut h = IndexMap::new(); + h.insert("X-Api-Key".to_string(), "k".to_string()); + h.insert("X-Trace".to_string(), "t".to_string()); + + let custom = oauth_custom_headers(Some(&h)).unwrap(); + + assert_eq!(custom.len(), 2); + assert_eq!(custom[&HeaderName::from_static("x-api-key")], "k"); + assert_eq!(custom[&HeaderName::from_static("x-trace")], "t"); + } + #[test] fn is_auth_required_error_matches_rmcp_message() { let e = anyhow!("Auth required, when send initialize request"); diff --git a/src/mcp/oauth.rs b/src/mcp/oauth.rs index bd2f89e..a911f36 100644 --- a/src/mcp/oauth.rs +++ b/src/mcp/oauth.rs @@ -10,6 +10,7 @@ use log::{debug, warn}; use reqwest::Client; use serde::{Deserialize, Serialize}; use std::collections::HashMap; +use std::fmt; use std::fs; use std::net::TcpListener; use std::sync::{Arc, OnceLock}; @@ -202,13 +203,23 @@ pub async fn run_mcp_oauth_flow( run_oauth_flow(&provider, &mcp_token_key(server_name)).await } -#[derive(Debug, PartialEq, Eq)] +#[derive(PartialEq, Eq)] pub enum McpTokenStatus { Token(String), NotAuthenticated, RefreshFailed, } +impl fmt::Debug for McpTokenStatus { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Token(_) => f.write_str("Token()"), + Self::NotAuthenticated => f.write_str("NotAuthenticated"), + Self::RefreshFailed => f.write_str("RefreshFailed"), + } + } +} + impl McpTokenStatus { pub fn into_token(self) -> Option { match self { @@ -219,11 +230,28 @@ impl McpTokenStatus { } pub async fn load_or_refresh_mcp_token(server_name: &str) -> McpTokenStatus { + load_or_refresh_inner(server_name, None).await +} + +/// Re-acquires a token after the server rejected the current one mid-session +/// (HTTP 401). The rejection proves the stored token is bad regardless of its +/// expiry timestamp, so the unexpired fast-paths only short-circuit when the +/// stored token DIFFERS from `rejected_token` (a concurrent caller genuinely +/// refreshed while we waited); an unexpired copy of the rejected token is +/// refreshed anyway. The failure backoff and per-server single-flight lock +/// still apply. +pub async fn force_refresh_mcp_token(server_name: &str, rejected_token: &str) -> Option { + load_or_refresh_inner(server_name, Some(rejected_token)) + .await + .into_token() +} + +async fn load_or_refresh_inner(server_name: &str, rejected_token: Option<&str>) -> McpTokenStatus { let key = mcp_token_key(server_name); let Some(tokens) = load_oauth_tokens(&key) else { return McpTokenStatus::NotAuthenticated; }; - if Utc::now().timestamp() < tokens.expires_at { + if rejected_token.is_none() && Utc::now().timestamp() < tokens.expires_at { return McpTokenStatus::Token(tokens.access_token); } @@ -235,11 +263,15 @@ pub async fn load_or_refresh_mcp_token(server_name: &str) -> McpTokenStatus { let lock = refresh_lock(server_name); let _guard = lock.lock().await; - // A concurrent caller may have refreshed while we waited for the lock. + // A concurrent caller may have refreshed while we waited for the lock. An + // unexpired token is only trusted if it differs from the rejected one: + // the server already proved that exact token bad. let Some(tokens) = load_oauth_tokens(&key) else { return McpTokenStatus::NotAuthenticated; }; - if Utc::now().timestamp() < tokens.expires_at { + if Utc::now().timestamp() < tokens.expires_at + && rejected_token.is_none_or(|rejected| rejected != tokens.access_token) + { return McpTokenStatus::Token(tokens.access_token); } @@ -589,10 +621,8 @@ fn extract_base_url(url: &str) -> Result { } #[cfg(test)] -mod tests { - use super::*; +pub(crate) mod test_support { use crate::utils::get_env_name; - use serial_test::serial; use std::{ env, ffi::OsString, @@ -601,7 +631,7 @@ mod tests { time::{self, SystemTime}, }; - fn with_temp_cache(f: F) { + pub(crate) fn with_temp_cache(f: F) { struct Restore { key: String, prev: Option, @@ -637,6 +667,14 @@ mod tests { }; f(); } +} + +#[cfg(test)] +mod tests { + use super::test_support::with_temp_cache; + use super::*; + use serial_test::serial; + use std::fs; #[test] fn extract_base_url_strips_path_and_query() { @@ -1034,6 +1072,151 @@ mod tests { }); } + #[test] + #[serial] + fn force_refresh_returns_concurrently_refreshed_token() { + with_temp_cache(|| { + fs::create_dir_all(paths::oauth_tokens_dir()).unwrap(); + fs::write( + paths::token_file("mcp_force-fresh"), + r#"{"access_token":"fresh-tok","refresh_token":"r","expires_at":9999999999}"#, + ) + .unwrap(); + + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let token = rt.block_on(force_refresh_mcp_token("force-fresh", "rejected-tok")); + + assert_eq!(token.as_deref(), Some("fresh-tok")); + }); + } + + #[test] + #[serial] + fn force_refresh_unexpired_rejected_token_attempts_real_refresh() { + with_temp_cache(|| { + fs::create_dir_all(paths::oauth_tokens_dir()).unwrap(); + fs::write( + paths::token_file("mcp_force-rejected"), + r#"{"access_token":"same-tok","refresh_token":"r","expires_at":9999999999}"#, + ) + .unwrap(); + + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let token = rt.block_on(force_refresh_mcp_token("force-rejected", "same-tok")); + + assert_eq!(token, None); + }); + } + + #[test] + #[serial] + fn force_refresh_missing_token_file_returns_none() { + with_temp_cache(|| { + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let token = rt.block_on(force_refresh_mcp_token( + "force-never-authed", + "rejected-tok", + )); + + assert_eq!(token, None); + }); + } + + #[test] + #[serial] + fn force_refresh_failed_refresh_returns_none() { + with_temp_cache(|| { + fs::create_dir_all(paths::oauth_tokens_dir()).unwrap(); + fs::write( + paths::token_file("mcp_force-fail"), + r#"{"access_token":"stale","refresh_token":"r","expires_at":0}"#, + ) + .unwrap(); + + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let token = rt.block_on(force_refresh_mcp_token("force-fail", "stale")); + + assert_eq!(token, None); + }); + } + + #[test] + #[serial] + fn force_refresh_concurrent_callers_complete_without_deadlock() { + with_temp_cache(|| { + fs::create_dir_all(paths::oauth_tokens_dir()).unwrap(); + fs::write( + paths::token_file("mcp_force-concurrent"), + r#"{"access_token":"same-tok","refresh_token":"r","expires_at":9999999999}"#, + ) + .unwrap(); + + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let (a, b) = rt.block_on(async { + tokio::join!( + force_refresh_mcp_token("force-concurrent", "same-tok"), + force_refresh_mcp_token("force-concurrent", "same-tok"), + ) + }); + + assert_eq!(a, None); + assert_eq!(b, None); + }); + } + + #[test] + #[serial] + fn force_refresh_respects_failure_backoff() { + with_temp_cache(|| { + fs::create_dir_all(paths::oauth_tokens_dir()).unwrap(); + fs::write( + paths::token_file("mcp_force-backoff"), + r#"{"access_token":"same-tok","refresh_token":"r","expires_at":9999999999}"#, + ) + .unwrap(); + note_refresh_failure("force-backoff"); + + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let token = rt.block_on(force_refresh_mcp_token("force-backoff", "same-tok")); + + assert_eq!(token, None); + }); + } + + #[test] + fn token_status_debug_redacts_token() { + assert_eq!( + format!("{:?}", McpTokenStatus::Token("live-secret".into())), + "Token()" + ); + assert_eq!( + format!("{:?}", McpTokenStatus::NotAuthenticated), + "NotAuthenticated" + ); + assert_eq!( + format!("{:?}", McpTokenStatus::RefreshFailed), + "RefreshFailed" + ); + } + #[test] fn token_status_into_token_extracts_only_token_variant() { assert_eq!(