feat: per-request OAuth token injection with mid-session refresh for HTTP MCP servers
Replace the spawn-time static Authorization header for OAuth-managed HTTP MCP servers with McpOAuthClient, a custom implementation of rmcp's StreamableHttpClient trait that resolves the bearer token on every request via load_or_refresh_mcp_token. Tokens that expire mid-session now refresh transparently instead of failing tool calls until restart. On a 401 for an injected token, the wrapper force-refreshes (identity- aware: a still-unexpired copy of the rejected token is not trusted) and retries exactly once, matching Claude Code / official SDK semantics. Both *_with_max_sse_event_size trait methods are overridden to preserve the inner client's SSE size enforcement, and the inner reqwest client mirrors rmcp's default (pool_max_idle_per_host(0), no redirects). SSE, stdio, and static-header HTTP paths are unchanged; startup warning semantics (McpAuthRequired reasons) are preserved. Verified live: mid-session backdated token refreshed transparently during an active atlassian session.
This commit is contained in:
@@ -1,7 +1,6 @@
|
|||||||
use crate::mcp::oauth::McpTokenStatus;
|
|
||||||
use crate::mcp::{
|
use crate::mcp::{
|
||||||
ConnectedServer, JsonField, McpAuthReason, McpAuthRequired, McpServer, McpTransportType,
|
ConnectedServer, JsonField, McpAuthRequired, McpServer, McpTransportType,
|
||||||
is_auth_required_error, oauth, spawn_mcp_server,
|
is_auth_required_error, resolve_http_auth, spawn_mcp_server,
|
||||||
};
|
};
|
||||||
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
@@ -103,13 +102,8 @@ impl McpFactory {
|
|||||||
return Ok(existing);
|
return Ok(existing);
|
||||||
}
|
}
|
||||||
|
|
||||||
let token_status = if spec.is_remote() {
|
let (auth, auth_reason) = resolve_http_auth(name, spec).await;
|
||||||
oauth::load_or_refresh_mcp_token(name).await
|
let handle = spawn_mcp_server(spec, log_path, auth)
|
||||||
} else {
|
|
||||||
McpTokenStatus::NotAuthenticated
|
|
||||||
};
|
|
||||||
let auth_reason = McpAuthReason::from_token_status(&token_status);
|
|
||||||
let handle = spawn_mcp_server(spec, log_path, token_status.into_token())
|
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
if is_auth_required_error(&e) {
|
if is_auth_required_error(&e) {
|
||||||
|
|||||||
@@ -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<C = reqwest::Client> {
|
||||||
|
inner: C,
|
||||||
|
server: Arc<str>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<C> McpOAuthClient<C> {
|
||||||
|
pub fn new(inner: C, server: &str) -> Self {
|
||||||
|
Self {
|
||||||
|
inner,
|
||||||
|
server: Arc::from(server),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<C: StreamableHttpClient + Sync> McpOAuthClient<C> {
|
||||||
|
/// 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<String>,
|
||||||
|
) -> Result<(Option<String>, bool), StreamableHttpError<C::Error>> {
|
||||||
|
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<C::Error> {
|
||||||
|
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<F, Fut>(
|
||||||
|
&self,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
mut post: F,
|
||||||
|
) -> Result<StreamableHttpPostResponse, StreamableHttpError<C::Error>>
|
||||||
|
where
|
||||||
|
F: FnMut(Option<String>) -> Fut,
|
||||||
|
Fut: Future<Output = Result<StreamableHttpPostResponse, StreamableHttpError<C::Error>>>,
|
||||||
|
{
|
||||||
|
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<C: StreamableHttpClient + Sync> StreamableHttpClient for McpOAuthClient<C> {
|
||||||
|
type Error = C::Error;
|
||||||
|
|
||||||
|
async fn post_message(
|
||||||
|
&self,
|
||||||
|
uri: Arc<str>,
|
||||||
|
message: ClientJsonRpcMessage,
|
||||||
|
session_id: Option<Arc<str>>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
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<str>,
|
||||||
|
message: ClientJsonRpcMessage,
|
||||||
|
session_id: Option<Arc<str>>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
max_sse_event_size: usize,
|
||||||
|
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
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<str>,
|
||||||
|
session_id: Arc<str>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
) -> Result<(), StreamableHttpError<Self::Error>> {
|
||||||
|
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<str>,
|
||||||
|
session_id: Option<Arc<str>>,
|
||||||
|
last_event_id: Option<String>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
) -> Result<BoxedSseResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
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<str>,
|
||||||
|
session_id: Option<Arc<str>>,
|
||||||
|
last_event_id: Option<String>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
max_sse_event_size: usize,
|
||||||
|
) -> Result<BoxedSseResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
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<Mutex<Vec<(&'static str, Option<String>)>>>;
|
||||||
|
|
||||||
|
type AfterFirstCallAction = Option<Box<dyn FnOnce() + Send + 'static>>;
|
||||||
|
|
||||||
|
/// 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<AtomicUsize>,
|
||||||
|
transport_error_times: Arc<AtomicUsize>,
|
||||||
|
after_first_call: Arc<Mutex<AfterFirstCallAction>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl FakeInner {
|
||||||
|
fn record(
|
||||||
|
&self,
|
||||||
|
method: &'static str,
|
||||||
|
auth: Option<String>,
|
||||||
|
) -> Result<(), StreamableHttpError<Infallible>> {
|
||||||
|
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<Infallible> {
|
||||||
|
StreamableHttpError::AuthRequired(AuthRequiredError::new(
|
||||||
|
"Bearer error=\"invalid_token\"".to_string(),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl StreamableHttpClient for FakeInner {
|
||||||
|
type Error = Infallible;
|
||||||
|
|
||||||
|
async fn post_message(
|
||||||
|
&self,
|
||||||
|
_uri: Arc<str>,
|
||||||
|
_message: ClientJsonRpcMessage,
|
||||||
|
_session_id: Option<Arc<str>>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
_custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
self.record("post_message", auth_header)?;
|
||||||
|
Ok(StreamableHttpPostResponse::Accepted)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn post_message_with_max_sse_event_size(
|
||||||
|
&self,
|
||||||
|
_uri: Arc<str>,
|
||||||
|
_message: ClientJsonRpcMessage,
|
||||||
|
_session_id: Option<Arc<str>>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
_custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
_max_sse_event_size: usize,
|
||||||
|
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
self.record("post_message_with_max_sse_event_size", auth_header)?;
|
||||||
|
Ok(StreamableHttpPostResponse::Accepted)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn delete_session(
|
||||||
|
&self,
|
||||||
|
_uri: Arc<str>,
|
||||||
|
_session_id: Arc<str>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
_custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
) -> Result<(), StreamableHttpError<Self::Error>> {
|
||||||
|
self.record("delete_session", auth_header)?;
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn get_stream(
|
||||||
|
&self,
|
||||||
|
_uri: Arc<str>,
|
||||||
|
_session_id: Option<Arc<str>>,
|
||||||
|
_last_event_id: Option<String>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
_custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
) -> Result<BoxedSseResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
self.record("get_stream", auth_header)?;
|
||||||
|
Ok(futures_util::stream::empty().boxed())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn get_stream_with_max_sse_event_size(
|
||||||
|
&self,
|
||||||
|
_uri: Arc<str>,
|
||||||
|
_session_id: Option<Arc<str>>,
|
||||||
|
_last_event_id: Option<String>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
_custom_headers: HashMap<HeaderName, HeaderValue>,
|
||||||
|
_max_sse_event_size: usize,
|
||||||
|
) -> Result<BoxedSseResponse, StreamableHttpError<Self::Error>> {
|
||||||
|
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<FakeInner>,
|
||||||
|
auth_header: Option<String>,
|
||||||
|
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Infallible>> {
|
||||||
|
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());
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
+211
-12
@@ -1,3 +1,4 @@
|
|||||||
|
mod auth_client;
|
||||||
pub(crate) mod manage;
|
pub(crate) mod manage;
|
||||||
pub(crate) mod oauth;
|
pub(crate) mod oauth;
|
||||||
mod sse_transport;
|
mod sse_transport;
|
||||||
@@ -9,6 +10,7 @@ use crate::vault::Vault;
|
|||||||
use crate::vault::interpolate_secrets;
|
use crate::vault::interpolate_secrets;
|
||||||
use anyhow::Error;
|
use anyhow::Error;
|
||||||
use anyhow::{Context, Result, anyhow};
|
use anyhow::{Context, Result, anyhow};
|
||||||
|
use auth_client::McpOAuthClient;
|
||||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||||
use http::{HeaderName, HeaderValue};
|
use http::{HeaderName, HeaderValue};
|
||||||
use indexmap::IndexMap;
|
use indexmap::IndexMap;
|
||||||
@@ -328,16 +330,9 @@ impl McpRegistry {
|
|||||||
.and_then(|c| c.mcp_servers.get(&id))
|
.and_then(|c| c.mcp_servers.get(&id))
|
||||||
.with_context(|| format!("MCP server not found in config: {id}"))?;
|
.with_context(|| format!("MCP server not found in config: {id}"))?;
|
||||||
|
|
||||||
let token_status = if spec.is_remote() {
|
let (auth, auth_reason) = resolve_http_auth(&id, spec).await;
|
||||||
oauth::load_or_refresh_mcp_token(&id).await
|
|
||||||
} else {
|
|
||||||
oauth::McpTokenStatus::NotAuthenticated
|
|
||||||
};
|
|
||||||
let auth_reason = McpAuthReason::from_token_status(&token_status);
|
|
||||||
|
|
||||||
let service =
|
let service = match spawn_mcp_server(spec, self.log_path.as_deref(), auth).await {
|
||||||
match spawn_mcp_server(spec, self.log_path.as_deref(), token_status.into_token()).await
|
|
||||||
{
|
|
||||||
Ok(s) => s,
|
Ok(s) => s,
|
||||||
Err(e) if is_auth_required_error(&e) => {
|
Err(e) if is_auth_required_error(&e) => {
|
||||||
warn!(
|
warn!(
|
||||||
@@ -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", &"<redacted>")
|
||||||
|
.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(
|
pub(crate) async fn spawn_mcp_server(
|
||||||
spec: &McpServer,
|
spec: &McpServer,
|
||||||
log_path: Option<&Path>,
|
log_path: Option<&Path>,
|
||||||
bearer_token: Option<String>,
|
auth: HttpAuth,
|
||||||
) -> Result<Arc<ConnectedServer>> {
|
) -> Result<Arc<ConnectedServer>> {
|
||||||
match spec.transport_type {
|
match spec.transport_type {
|
||||||
McpTransportType::Http => {
|
McpTransportType::Http => {
|
||||||
let url = spec.url.as_deref().expect("validated: http spec has url");
|
let url = spec.url.as_deref().expect("validated: http spec has url");
|
||||||
let headers = merge_bearer_token(spec.headers.as_ref(), bearer_token);
|
match auth {
|
||||||
spawn_http_mcp_server(url, headers.as_ref()).await
|
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 => {
|
McpTransportType::Sse => {
|
||||||
let url = spec.url.as_deref().expect("validated: sse spec has url");
|
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);
|
let headers = merge_bearer_token(spec.headers.as_ref(), bearer_token);
|
||||||
spawn_sse_mcp_server(url, headers.as_ref()).await
|
spawn_sse_mcp_server(url, headers.as_ref()).await
|
||||||
}
|
}
|
||||||
@@ -460,6 +512,7 @@ fn merge_bearer_token(
|
|||||||
}
|
}
|
||||||
(Some(h), Some(token)) => {
|
(Some(h), Some(token)) => {
|
||||||
let mut m = h.clone();
|
let mut m = h.clone();
|
||||||
|
m.retain(|k, _| !k.eq_ignore_ascii_case("authorization"));
|
||||||
m.insert("Authorization".to_string(), format!("Bearer {token}"));
|
m.insert("Authorization".to_string(), format!("Bearer {token}"));
|
||||||
Some(m)
|
Some(m)
|
||||||
}
|
}
|
||||||
@@ -551,6 +604,66 @@ async fn spawn_http_mcp_server(
|
|||||||
Ok(service)
|
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<String, String>>,
|
||||||
|
) -> Result<HashMap<HeaderName, HeaderValue>> {
|
||||||
|
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::<HeaderName>()
|
||||||
|
.with_context(|| format!("Invalid header name: {k}"))?;
|
||||||
|
let value = v
|
||||||
|
.parse::<HeaderValue>()
|
||||||
|
.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<String, String>>,
|
||||||
|
) -> Result<Arc<ConnectedServer>> {
|
||||||
|
// 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(
|
async fn spawn_sse_mcp_server(
|
||||||
url: &str,
|
url: &str,
|
||||||
headers: Option<&IndexMap<String, String>>,
|
headers: Option<&IndexMap<String, String>>,
|
||||||
@@ -1109,6 +1222,92 @@ mod tests {
|
|||||||
assert_eq!(result["X-Custom"], "keep");
|
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("<redacted>"));
|
||||||
|
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]
|
#[test]
|
||||||
fn is_auth_required_error_matches_rmcp_message() {
|
fn is_auth_required_error_matches_rmcp_message() {
|
||||||
let e = anyhow!("Auth required, when send initialize request");
|
let e = anyhow!("Auth required, when send initialize request");
|
||||||
|
|||||||
+191
-8
@@ -10,6 +10,7 @@ use log::{debug, warn};
|
|||||||
use reqwest::Client;
|
use reqwest::Client;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
|
use std::fmt;
|
||||||
use std::fs;
|
use std::fs;
|
||||||
use std::net::TcpListener;
|
use std::net::TcpListener;
|
||||||
use std::sync::{Arc, OnceLock};
|
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
|
run_oauth_flow(&provider, &mcp_token_key(server_name)).await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, PartialEq, Eq)]
|
#[derive(PartialEq, Eq)]
|
||||||
pub enum McpTokenStatus {
|
pub enum McpTokenStatus {
|
||||||
Token(String),
|
Token(String),
|
||||||
NotAuthenticated,
|
NotAuthenticated,
|
||||||
RefreshFailed,
|
RefreshFailed,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl fmt::Debug for McpTokenStatus {
|
||||||
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
|
match self {
|
||||||
|
Self::Token(_) => f.write_str("Token(<redacted>)"),
|
||||||
|
Self::NotAuthenticated => f.write_str("NotAuthenticated"),
|
||||||
|
Self::RefreshFailed => f.write_str("RefreshFailed"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl McpTokenStatus {
|
impl McpTokenStatus {
|
||||||
pub fn into_token(self) -> Option<String> {
|
pub fn into_token(self) -> Option<String> {
|
||||||
match self {
|
match self {
|
||||||
@@ -219,11 +230,28 @@ impl McpTokenStatus {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_or_refresh_mcp_token(server_name: &str) -> 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<String> {
|
||||||
|
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 key = mcp_token_key(server_name);
|
||||||
let Some(tokens) = load_oauth_tokens(&key) else {
|
let Some(tokens) = load_oauth_tokens(&key) else {
|
||||||
return McpTokenStatus::NotAuthenticated;
|
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);
|
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 lock = refresh_lock(server_name);
|
||||||
let _guard = lock.lock().await;
|
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 {
|
let Some(tokens) = load_oauth_tokens(&key) else {
|
||||||
return McpTokenStatus::NotAuthenticated;
|
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);
|
return McpTokenStatus::Token(tokens.access_token);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -589,10 +621,8 @@ fn extract_base_url(url: &str) -> Result<String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
pub(crate) mod test_support {
|
||||||
use super::*;
|
|
||||||
use crate::utils::get_env_name;
|
use crate::utils::get_env_name;
|
||||||
use serial_test::serial;
|
|
||||||
use std::{
|
use std::{
|
||||||
env,
|
env,
|
||||||
ffi::OsString,
|
ffi::OsString,
|
||||||
@@ -601,7 +631,7 @@ mod tests {
|
|||||||
time::{self, SystemTime},
|
time::{self, SystemTime},
|
||||||
};
|
};
|
||||||
|
|
||||||
fn with_temp_cache<F: FnOnce()>(f: F) {
|
pub(crate) fn with_temp_cache<F: FnOnce()>(f: F) {
|
||||||
struct Restore {
|
struct Restore {
|
||||||
key: String,
|
key: String,
|
||||||
prev: Option<OsString>,
|
prev: Option<OsString>,
|
||||||
@@ -637,6 +667,14 @@ mod tests {
|
|||||||
};
|
};
|
||||||
f();
|
f();
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::test_support::with_temp_cache;
|
||||||
|
use super::*;
|
||||||
|
use serial_test::serial;
|
||||||
|
use std::fs;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn extract_base_url_strips_path_and_query() {
|
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(<redacted>)"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
format!("{:?}", McpTokenStatus::NotAuthenticated),
|
||||||
|
"NotAuthenticated"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
format!("{:?}", McpTokenStatus::RefreshFailed),
|
||||||
|
"RefreshFailed"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn token_status_into_token_extracts_only_token_variant() {
|
fn token_status_into_token_extracts_only_token_variant() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
|
|||||||
Reference in New Issue
Block a user