Allow custom HTTP clients for OAuth (#908)

* feat: allow custom HTTP clients for OAuth

* fix(auth): preserve configured client for refresh

* fix(auth): harden OAuth HTTP adapter

* refactor(auth): simplify OAuth HTTP plumbing

* fix(auth): stop refresh token redirects by default

* fix(auth): re-export OAuth HTTP client types
This commit is contained in:
jif 2026-06-22 18:24:32 +01:00 committed by GitHub
parent 3c5ce2b0d7
commit de898dd842
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 579 additions and 127 deletions

View file

@ -101,8 +101,9 @@ pub use auth::JwtSigningAlgorithm;
pub use auth::{ pub use auth::{
AuthClient, AuthError, AuthorizationManager, AuthorizationSession, AuthorizedHttpClient, AuthClient, AuthError, AuthorizationManager, AuthorizationSession, AuthorizedHttpClient,
ClientCredentialsConfig, CredentialStore, EXTENSION_OAUTH_CLIENT_CREDENTIALS, ClientCredentialsConfig, CredentialStore, EXTENSION_OAUTH_CLIENT_CREDENTIALS,
InMemoryCredentialStore, InMemoryStateStore, ScopeUpgradeConfig, StateStore, InMemoryCredentialStore, InMemoryStateStore, OAuthHttpClient, OAuthHttpClientError,
StoredAuthorizationState, StoredCredentials, WWWAuthenticateParams, OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, ScopeUpgradeConfig,
StateStore, StoredAuthorizationState, StoredCredentials, WWWAuthenticateParams,
}; };
// #[cfg(feature = "transport-ws")] // #[cfg(feature = "transport-ws")]

View file

@ -1,19 +1,22 @@
use std::{ use std::{
collections::HashMap, collections::HashMap,
future::Future,
pin::Pin,
sync::Arc, sync::Arc,
time::{Duration, SystemTime, UNIX_EPOCH}, time::{Duration, SystemTime, UNIX_EPOCH},
}; };
use async_trait::async_trait; use async_trait::async_trait;
use futures::StreamExt;
use oauth2::{ use oauth2::{
AsyncHttpClient, AuthType, AuthUrl, AuthorizationCode, ClientId, ClientSecret, CsrfToken, AsyncHttpClient, AuthType, AuthUrl, AuthorizationCode, ClientId, ClientSecret, CsrfToken,
EmptyExtraTokenFields, ExtraTokenFields, HttpClientError, HttpRequest, HttpResponse, EmptyExtraTokenFields, ExtraTokenFields, HttpRequest, HttpResponse, PkceCodeChallenge,
PkceCodeChallenge, PkceCodeVerifier, RedirectUrl, RefreshToken, RequestTokenError, Scope, PkceCodeVerifier, RedirectUrl, RefreshToken, RequestTokenError, Scope, StandardTokenResponse,
StandardTokenResponse, TokenResponse, TokenUrl, basic::BasicTokenType, TokenResponse, TokenUrl, basic::BasicTokenType,
}; };
use reqwest::{ use reqwest::{
Client as HttpClient, IntoUrl, StatusCode, Url, Client as ReqwestClient, IntoUrl, StatusCode, Url,
header::{AUTHORIZATION, WWW_AUTHENTICATE}, header::{AUTHORIZATION, CONTENT_TYPE, WWW_AUTHENTICATE},
}; };
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::Value; use serde_json::Value;
@ -23,39 +26,152 @@ use tracing::{debug, warn};
use crate::transport::common::http_header::HEADER_MCP_PROTOCOL_VERSION; use crate::transport::common::http_header::HEADER_MCP_PROTOCOL_VERSION;
/// Owned wrapper around [`reqwest::Client`] that implements [`AsyncHttpClient`] for oauth2. const DEFAULT_HTTP_TIMEOUT: Duration = Duration::from_secs(30);
struct OAuthReqwestClient(HttpClient); const MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES: usize = 1024 * 1024;
impl<'c> AsyncHttpClient<'c> for OAuthReqwestClient { /// Redirect handling requested for an outbound OAuth HTTP operation.
type Error = HttpClientError<reqwest::Error>; #[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub enum OAuthHttpRedirectPolicy {
/// Follow redirects using the client's normal limits.
#[default]
Follow,
/// Return the redirect response without following its location.
Stop,
}
type Future = std::pin::Pin< /// A complete outbound HTTP operation requested by the OAuth implementation.
Box<dyn std::future::Future<Output = Result<HttpResponse, Self::Error>> + Send + Sync + 'c>, #[non_exhaustive]
>; pub struct OAuthHttpRequest {
/// HTTP request with an absolute URI and buffered body.
pub request: HttpRequest,
/// Redirect behavior required by the OAuth operation.
pub redirect_policy: OAuthHttpRedirectPolicy,
/// Suggested maximum duration for the operation, or no SDK-specified timeout.
/// Implementations with their own timeout policy may retain it instead.
pub timeout: Option<Duration>,
}
fn call(&'c self, request: HttpRequest) -> Self::Future { impl OAuthHttpRequest {
fn new(request: HttpRequest, redirect_policy: OAuthHttpRedirectPolicy) -> Self {
Self {
request,
redirect_policy,
timeout: Some(DEFAULT_HTTP_TIMEOUT),
}
}
}
/// Error returned by a custom OAuth HTTP client.
#[derive(Debug, Error)]
#[error("{message}")]
pub struct OAuthHttpClientError {
message: String,
}
impl OAuthHttpClientError {
/// Create an error from a transport-provided message.
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
}
/// Future returned by [`OAuthHttpClient::execute`].
pub type OAuthHttpClientFuture<'a> =
Pin<Box<dyn Future<Output = Result<HttpResponse, OAuthHttpClientError>> + Send + 'a>>;
/// Executes every outbound HTTP request made by the OAuth state machine.
///
/// Implementations may route requests through a remote execution environment.
/// They must honor the request's redirect policy and return the raw response
/// status, headers, and body.
pub trait OAuthHttpClient: Send + Sync {
/// Execute one OAuth HTTP operation.
fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_>;
}
struct ReqwestOAuthHttpClient {
follow_redirects: ReqwestClient,
stop_redirects: ReqwestClient,
}
impl ReqwestOAuthHttpClient {
fn new(follow_redirects: ReqwestClient) -> Result<Self, AuthError> {
let stop_redirects = ReqwestClient::builder()
.timeout(DEFAULT_HTTP_TIMEOUT)
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|error| AuthError::InternalError(error.to_string()))?;
Ok(Self {
follow_redirects,
stop_redirects,
})
}
}
impl OAuthHttpClient for ReqwestOAuthHttpClient {
fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> {
Box::pin(async move { Box::pin(async move {
let response = self let OAuthHttpRequest {
.0 request,
.execute(request.try_into().map_err(Box::new)?) redirect_policy,
..
} = request;
let client = match redirect_policy {
OAuthHttpRedirectPolicy::Follow => &self.follow_redirects,
OAuthHttpRedirectPolicy::Stop => &self.stop_redirects,
};
let request = reqwest::Request::try_from(request)
.map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
let response = client
.execute(request)
.await .await
.map_err(Box::new)?; .map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
let mut builder = oauth2::http::Response::builder() let mut builder = oauth2::http::Response::builder()
.status(response.status()) .status(response.status())
.version(response.version()); .version(response.version());
for (name, value) in response.headers() {
for (name, value) in response.headers().iter() {
builder = builder.header(name, value); builder = builder.header(name, value);
} }
let mut body = Vec::new();
let mut body_stream = response.bytes_stream();
while let Some(chunk) = body_stream.next().await {
let chunk = chunk.map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
if chunk.len() > MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES - body.len() {
return Err(OAuthHttpClientError::new(format!(
"OAuth HTTP response body exceeds {MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES} bytes"
)));
}
body.extend_from_slice(&chunk);
}
builder builder
.body(response.bytes().await.map_err(Box::new)?.to_vec()) .body(body)
.map_err(HttpClientError::Http) .map_err(|error| OAuthHttpClientError::new(error.to_string()))
}) })
} }
} }
struct OAuth2HttpClient<'a> {
client: &'a dyn OAuthHttpClient,
redirect_policy: OAuthHttpRedirectPolicy,
}
impl<'c> AsyncHttpClient<'c> for OAuth2HttpClient<'_> {
type Error = OAuthHttpClientError;
type Future = std::pin::Pin<
Box<dyn std::future::Future<Output = Result<HttpResponse, Self::Error>> + Send + 'c>,
>;
fn call(&'c self, request: HttpRequest) -> Self::Future {
self.client
.execute(OAuthHttpRequest::new(request, self.redirect_policy))
}
}
const DEFAULT_EXCHANGE_URL: &str = "http://localhost"; const DEFAULT_EXCHANGE_URL: &str = "http://localhost";
/// Default OIDC Dynamic Client Registration `application_type` (SEP-837) /// Default OIDC Dynamic Client Registration `application_type` (SEP-837)
@ -639,7 +755,9 @@ impl Default for ScopeUpgradeConfig {
/// oauth2 auth manager /// oauth2 auth manager
pub struct AuthorizationManager { pub struct AuthorizationManager {
http_client: HttpClient, http_client: Arc<dyn OAuthHttpClient>,
// Preserve legacy reqwest refresh behavior without weakening custom clients.
refresh_redirect_policy: OAuthHttpRedirectPolicy,
metadata: Option<AuthorizationMetadata>, metadata: Option<AuthorizationMetadata>,
oauth_client: Option<OAuthClient>, oauth_client: Option<OAuthClient>,
credential_store: Arc<dyn CredentialStore>, credential_store: Arc<dyn CredentialStore>,
@ -732,14 +850,36 @@ impl AuthorizationManager {
/// create new auth manager with base url /// create new auth manager with base url
pub async fn new<U: IntoUrl>(base_url: U) -> Result<Self, AuthError> { pub async fn new<U: IntoUrl>(base_url: U) -> Result<Self, AuthError> {
let base_url = base_url.into_url()?; let http_client = ReqwestClient::builder()
let http_client = HttpClient::builder() .timeout(DEFAULT_HTTP_TIMEOUT)
.timeout(Duration::from_secs(30))
.build() .build()
.map_err(|e| AuthError::InternalError(e.to_string()))?; .map_err(|e| AuthError::InternalError(e.to_string()))?;
Self::new_inner(
base_url,
Arc::new(ReqwestOAuthHttpClient::new(http_client)?),
OAuthHttpRedirectPolicy::Stop,
)
.await
}
/// Create an auth manager with a client used for every OAuth HTTP operation.
pub async fn new_with_oauth_http_client<U: IntoUrl>(
base_url: U,
http_client: Arc<dyn OAuthHttpClient>,
) -> Result<Self, AuthError> {
Self::new_inner(base_url, http_client, OAuthHttpRedirectPolicy::Stop).await
}
async fn new_inner<U: IntoUrl>(
base_url: U,
http_client: Arc<dyn OAuthHttpClient>,
refresh_redirect_policy: OAuthHttpRedirectPolicy,
) -> Result<Self, AuthError> {
let base_url = base_url.into_url()?;
let manager = Self { let manager = Self {
http_client, http_client,
refresh_redirect_policy,
metadata: None, metadata: None,
oauth_client: None, oauth_client: None,
credential_store: Arc::new(InMemoryCredentialStore::new()), credential_store: Arc::new(InMemoryCredentialStore::new()),
@ -804,8 +944,9 @@ impl AuthorizationManager {
Ok(false) Ok(false)
} }
pub fn with_client(&mut self, http_client: HttpClient) -> Result<(), AuthError> { pub fn with_client(&mut self, http_client: ReqwestClient) -> Result<(), AuthError> {
self.http_client = http_client; self.http_client = Arc::new(ReqwestOAuthHttpClient::new(http_client)?);
self.refresh_redirect_policy = OAuthHttpRedirectPolicy::Follow;
Ok(()) Ok(())
} }
@ -957,11 +1098,21 @@ impl AuthorizationManager {
application_type: application_type.clone(), application_type: application_type.clone(),
}; };
let request = oauth2::http::Request::builder()
.method("POST")
.uri(registration_url)
.header(CONTENT_TYPE, "application/json")
.body(
serde_json::to_vec(&registration_request)
.map_err(|error| AuthError::RegistrationFailed(error.to_string()))?,
)
.map_err(|error| AuthError::RegistrationFailed(error.to_string()))?;
let response = match self let response = match self
.http_client .http_client
.post(registration_url) .execute(OAuthHttpRequest::new(
.json(&registration_request) request,
.send() OAuthHttpRedirectPolicy::Follow,
))
.await .await
{ {
Ok(response) => response, Ok(response) => response,
@ -975,10 +1126,7 @@ impl AuthorizationManager {
if !response.status().is_success() { if !response.status().is_success() {
let status = response.status(); let status = response.status();
let error_text = match response.text().await { let error_text = String::from_utf8_lossy(response.body());
Ok(text) => text,
Err(_) => "cannot get error details".to_string(),
};
return Err(AuthError::RegistrationFailed(format!( return Err(AuthError::RegistrationFailed(format!(
"HTTP {}: {}", "HTTP {}: {}",
@ -986,16 +1134,17 @@ impl AuthorizationManager {
))); )));
} }
debug!("registration response: {:?}", response); debug!("registration response status: {:?}", response.status());
let reg_response = match response.json::<ClientRegistrationResponse>().await { let reg_response =
Ok(response) => response, match serde_json::from_slice::<ClientRegistrationResponse>(response.body()) {
Err(e) => { Ok(response) => response,
return Err(AuthError::RegistrationFailed(format!( Err(e) => {
"analyze response error: {}", return Err(AuthError::RegistrationFailed(format!(
e "analyze response error: {}",
))); e
} )));
}; }
};
let config = OAuthClientConfig { let config = OAuthClientConfig {
client_id: reg_response.client_id, client_id: reg_response.client_id,
@ -1287,10 +1436,6 @@ impl AuthorizationManager {
// Reconstruct the PKCE verifier // Reconstruct the PKCE verifier
let pkce_verifier = stored_state.into_pkce_verifier(); let pkce_verifier = stored_state.into_pkce_verifier();
let http_client = reqwest::ClientBuilder::new()
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::InternalError(e.to_string()))?;
debug!("client_id: {:?}", oauth_client.client_id()); debug!("client_id: {:?}", oauth_client.client_id());
// exchange token // exchange token
@ -1298,7 +1443,10 @@ impl AuthorizationManager {
.exchange_code(AuthorizationCode::new(code.to_string())) .exchange_code(AuthorizationCode::new(code.to_string()))
.set_pkce_verifier(pkce_verifier) .set_pkce_verifier(pkce_verifier)
.add_extra_param("resource", self.base_url.to_string()) .add_extra_param("resource", self.base_url.to_string())
.request_async(&OAuthReqwestClient(http_client)) .request_async(&OAuth2HttpClient {
client: self.http_client.as_ref(),
redirect_policy: OAuthHttpRedirectPolicy::Stop,
})
.await .await
{ {
Ok(token) => token, Ok(token) => token,
@ -1432,7 +1580,10 @@ impl AuthorizationManager {
refresh_request = refresh_request.add_scope(Scope::new(scope)); refresh_request = refresh_request.add_scope(Scope::new(scope));
} }
let token_result = refresh_request let token_result = refresh_request
.request_async(&OAuthReqwestClient(self.http_client.clone())) .request_async(&OAuth2HttpClient {
client: self.http_client.as_ref(),
redirect_policy: self.refresh_redirect_policy,
})
.await .await
.map_err(|e| AuthError::TokenRefreshFailed(e.to_string()))?; .map_err(|e| AuthError::TokenRefreshFailed(e.to_string()))?;
@ -1539,13 +1690,7 @@ impl AuthorizationManager {
discovery_url: &Url, discovery_url: &Url,
) -> Result<Option<AuthorizationMetadata>, AuthError> { ) -> Result<Option<AuthorizationMetadata>, AuthError> {
debug!("discovery url: {:?}", discovery_url); debug!("discovery url: {:?}", discovery_url);
let response = match self let response = match self.discovery_get(discovery_url).await {
.http_client
.get(discovery_url.clone())
.header(HEADER_MCP_PROTOCOL_VERSION, "2024-11-05")
.send()
.await
{
Ok(r) => r, Ok(r) => r,
Err(e) => { Err(e) => {
debug!("discovery request failed: {}", e); debug!("discovery request failed: {}", e);
@ -1558,8 +1703,7 @@ impl AuthorizationManager {
return Ok(None); return Ok(None);
} }
let body = response.text().await?; match serde_json::from_slice::<AuthorizationMetadata>(response.body()) {
match serde_json::from_str::<AuthorizationMetadata>(&body) {
Ok(metadata) => Ok(Some(metadata)), Ok(metadata) => Ok(Some(metadata)),
Err(err) => { Err(err) => {
debug!("Failed to parse metadata for {}: {}", discovery_url, err); debug!("Failed to parse metadata for {}: {}", discovery_url, err);
@ -1659,13 +1803,7 @@ impl AuthorizationManager {
/// Extract the resource metadata url from the WWW-Authenticate header value. /// Extract the resource metadata url from the WWW-Authenticate header value.
/// https://www.rfc-editor.org/rfc/rfc9728.html#name-use-of-www-authenticate-for /// https://www.rfc-editor.org/rfc/rfc9728.html#name-use-of-www-authenticate-for
async fn fetch_resource_metadata_url(&self, url: &Url) -> Result<Option<Url>, AuthError> { async fn fetch_resource_metadata_url(&self, url: &Url) -> Result<Option<Url>, AuthError> {
let response = match self let response = match self.discovery_get(url).await {
.http_client
.get(url.clone())
.header(HEADER_MCP_PROTOCOL_VERSION, "2024-11-05")
.send()
.await
{
Ok(r) => r, Ok(r) => r,
Err(e) => { Err(e) => {
debug!("resource metadata probe failed: {}", e); debug!("resource metadata probe failed: {}", e);
@ -1712,13 +1850,7 @@ impl AuthorizationManager {
"resource metadata discovery url: {:?}", "resource metadata discovery url: {:?}",
resource_metadata_url resource_metadata_url
); );
let response = match self let response = match self.discovery_get(resource_metadata_url).await {
.http_client
.get(resource_metadata_url.clone())
.header(HEADER_MCP_PROTOCOL_VERSION, "2024-11-05")
.send()
.await
{
Ok(r) => r, Ok(r) => r,
Err(e) => { Err(e) => {
debug!("resource metadata request failed: {}", e); debug!("resource metadata request failed: {}", e);
@ -1734,7 +1866,7 @@ impl AuthorizationManager {
return Ok(None); return Ok(None);
} }
let metadata = match response.json::<ResourceServerMetadata>().await { let metadata = match serde_json::from_slice::<ResourceServerMetadata>(response.body()) {
Ok(metadata) => metadata, Ok(metadata) => metadata,
Err(e) => { Err(e) => {
debug!("failed to parse resource metadata as JSON: {}", e); debug!("failed to parse resource metadata as JSON: {}", e);
@ -1744,6 +1876,21 @@ impl AuthorizationManager {
Ok(Some(metadata)) Ok(Some(metadata))
} }
async fn discovery_get(&self, url: &Url) -> Result<HttpResponse, OAuthHttpClientError> {
let request = oauth2::http::Request::builder()
.method("GET")
.uri(url.as_str())
.header(HEADER_MCP_PROTOCOL_VERSION, "2024-11-05")
.body(Vec::new())
.map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
self.http_client
.execute(OAuthHttpRequest::new(
request,
OAuthHttpRedirectPolicy::Follow,
))
.await
}
/// extract parameters from WWW-Authenticate header (resource_metadata and scope) /// extract parameters from WWW-Authenticate header (resource_metadata and scope)
fn extract_www_authenticate_params(header: &str, base_url: &Url) -> WWWAuthenticateParams { fn extract_www_authenticate_params(header: &str, base_url: &Url) -> WWWAuthenticateParams {
let mut params = WWWAuthenticateParams::default(); let mut params = WWWAuthenticateParams::default();
@ -2016,13 +2163,11 @@ impl AuthorizationManager {
request = request.add_extra_param("resource", resource); request = request.add_extra_param("resource", resource);
} }
let http_client = reqwest::ClientBuilder::new()
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::InternalError(e.to_string()))?;
let token_result = match request let token_result = match request
.request_async(&OAuthReqwestClient(http_client)) .request_async(&OAuth2HttpClient {
client: self.http_client.as_ref(),
redirect_policy: OAuthHttpRedirectPolicy::Stop,
})
.await .await
{ {
Ok(token) => token, Ok(token) => token,
@ -2135,28 +2280,28 @@ impl AuthorizationManager {
} }
let body_str = serializer.finish(); let body_str = serializer.finish();
let http_client = reqwest::ClientBuilder::new() let request = oauth2::http::Request::builder()
.redirect(reqwest::redirect::Policy::none()) .method("POST")
.build() .uri(token_endpoint_url.as_str())
.map_err(|e| AuthError::InternalError(e.to_string()))?; .header(CONTENT_TYPE, "application/x-www-form-urlencoded")
.body(body_str.into_bytes())
let response = http_client .map_err(|error| AuthError::ClientCredentialsError(error.to_string()))?;
.post(token_endpoint_url.as_str()) let response = self
.header("content-type", "application/x-www-form-urlencoded") .http_client
.body(body_str) .execute(OAuthHttpRequest::new(
.send() request,
OAuthHttpRedirectPolicy::Stop,
))
.await .await
.map_err(|e| { .map_err(|e| {
AuthError::ClientCredentialsError(format!("Token exchange request failed: {e}")) AuthError::ClientCredentialsError(format!("Token exchange request failed: {e}"))
})?; })?;
let status = response.status(); let status = response.status();
let body = response.bytes().await.map_err(|e| { let body = response.body();
AuthError::ClientCredentialsError(format!("Failed to read token response: {e}"))
})?;
if !status.is_success() { if !status.is_success() {
let msg = if let Ok(v) = serde_json::from_slice::<serde_json::Value>(&body) { let msg = if let Ok(v) = serde_json::from_slice::<serde_json::Value>(body) {
let error = v.get("error").and_then(|e| e.as_str()).unwrap_or("unknown"); let error = v.get("error").and_then(|e| e.as_str()).unwrap_or("unknown");
let desc = v let desc = v
.get("error_description") .get("error_description")
@ -2169,7 +2314,7 @@ impl AuthorizationManager {
return Err(AuthError::ClientCredentialsError(msg)); return Err(AuthError::ClientCredentialsError(msg));
} }
let token_result = serde_json::from_slice::<OAuthTokenResponse>(&body).map_err(|e| { let token_result = serde_json::from_slice::<OAuthTokenResponse>(body).map_err(|e| {
AuthError::ClientCredentialsError(format!("Failed to parse token response: {e}")) AuthError::ClientCredentialsError(format!("Failed to parse token response: {e}"))
})?; })?;
@ -2415,12 +2560,12 @@ impl AuthorizationSession {
/// http client extension, automatically add authorization header /// http client extension, automatically add authorization header
pub struct AuthorizedHttpClient { pub struct AuthorizedHttpClient {
auth_manager: Arc<AuthorizationManager>, auth_manager: Arc<AuthorizationManager>,
inner_client: HttpClient, inner_client: ReqwestClient,
} }
impl AuthorizedHttpClient { impl AuthorizedHttpClient {
/// create new authorized http client /// create new authorized http client
pub fn new(auth_manager: Arc<AuthorizationManager>, client: Option<HttpClient>) -> Self { pub fn new(auth_manager: Arc<AuthorizationManager>, client: Option<ReqwestClient>) -> Self {
let inner_client = client.unwrap_or_default(); let inner_client = client.unwrap_or_default();
Self { Self {
auth_manager, auth_manager,
@ -2467,10 +2612,34 @@ pub enum OAuthState {
} }
impl OAuthState { impl OAuthState {
fn oauth_http_client_config(&self) -> (Arc<dyn OAuthHttpClient>, OAuthHttpRedirectPolicy) {
let manager = match self {
OAuthState::Unauthorized(manager) | OAuthState::Authorized(manager) => manager,
OAuthState::Session(session) => &session.auth_manager,
OAuthState::AuthorizedHttpClient(client) => &client.auth_manager,
};
(
Arc::clone(&manager.http_client),
manager.refresh_redirect_policy,
)
}
async fn placeholder(&self) -> Result<Self, AuthError> {
let (http_client, refresh_redirect_policy) = self.oauth_http_client_config();
Ok(OAuthState::Unauthorized(
AuthorizationManager::new_inner(
DEFAULT_EXCHANGE_URL,
http_client,
refresh_redirect_policy,
)
.await?,
))
}
/// Create new OAuth state machine /// Create new OAuth state machine
pub async fn new<U: IntoUrl>( pub async fn new<U: IntoUrl>(
base_url: U, base_url: U,
client: Option<HttpClient>, client: Option<ReqwestClient>,
) -> Result<Self, AuthError> { ) -> Result<Self, AuthError> {
let mut manager = AuthorizationManager::new(base_url).await?; let mut manager = AuthorizationManager::new(base_url).await?;
if let Some(client) = client { if let Some(client) = client {
@ -2480,6 +2649,16 @@ impl OAuthState {
Ok(OAuthState::Unauthorized(manager)) Ok(OAuthState::Unauthorized(manager))
} }
/// Create an OAuth state machine that routes all OAuth HTTP operations
/// through the supplied client.
pub async fn new_with_oauth_http_client<U: IntoUrl>(
base_url: U,
client: Arc<dyn OAuthHttpClient>,
) -> Result<Self, AuthError> {
let manager = AuthorizationManager::new_with_oauth_http_client(base_url, client).await?;
Ok(OAuthState::Unauthorized(manager))
}
/// Get client_id and OAuth credentials /// Get client_id and OAuth credentials
pub async fn get_credentials(&self) -> Result<Credentials, AuthError> { pub async fn get_credentials(&self) -> Result<Credentials, AuthError> {
// return client_id and credentials // return client_id and credentials
@ -2500,10 +2679,13 @@ impl OAuthState {
credentials: OAuthTokenResponse, credentials: OAuthTokenResponse,
) -> Result<(), AuthError> { ) -> Result<(), AuthError> {
if let OAuthState::Unauthorized(manager) = self { if let OAuthState::Unauthorized(manager) = self {
let mut manager = std::mem::replace( let replacement = AuthorizationManager::new_inner(
manager, DEFAULT_EXCHANGE_URL,
AuthorizationManager::new(DEFAULT_EXCHANGE_URL).await?, Arc::clone(&manager.http_client),
); manager.refresh_redirect_policy,
)
.await?;
let mut manager = std::mem::replace(manager, replacement);
let granted_scopes: Vec<String> = credentials let granted_scopes: Vec<String> = credentials
.scopes() .scopes()
@ -2553,10 +2735,8 @@ impl OAuthState {
client_name: Option<&str>, client_name: Option<&str>,
client_metadata_url: Option<&str>, client_metadata_url: Option<&str>,
) -> Result<(), AuthError> { ) -> Result<(), AuthError> {
if let OAuthState::Unauthorized(mut manager) = std::mem::replace( let placeholder = self.placeholder().await?;
self, if let OAuthState::Unauthorized(mut manager) = std::mem::replace(self, placeholder) {
OAuthState::Unauthorized(AuthorizationManager::new(DEFAULT_EXCHANGE_URL).await?),
) {
debug!("start discovery"); debug!("start discovery");
let metadata = manager.discover_metadata().await?; let metadata = manager.discover_metadata().await?;
manager.metadata = Some(metadata); manager.metadata = Some(metadata);
@ -2588,10 +2768,8 @@ impl OAuthState {
/// complete authorization /// complete authorization
pub async fn complete_authorization(&mut self) -> Result<(), AuthError> { pub async fn complete_authorization(&mut self) -> Result<(), AuthError> {
if let OAuthState::Session(session) = std::mem::replace( let placeholder = self.placeholder().await?;
self, if let OAuthState::Session(session) = std::mem::replace(self, placeholder) {
OAuthState::Unauthorized(AuthorizationManager::new(DEFAULT_EXCHANGE_URL).await?),
) {
*self = OAuthState::Authorized(session.auth_manager); *self = OAuthState::Authorized(session.auth_manager);
Ok(()) Ok(())
} else { } else {
@ -2600,10 +2778,8 @@ impl OAuthState {
} }
/// covert to authorized http client /// covert to authorized http client
pub async fn to_authorized_http_client(&mut self) -> Result<(), AuthError> { pub async fn to_authorized_http_client(&mut self) -> Result<(), AuthError> {
if let OAuthState::Authorized(manager) = std::mem::replace( let placeholder = self.placeholder().await?;
self, if let OAuthState::Authorized(manager) = std::mem::replace(self, placeholder) {
OAuthState::Authorized(AuthorizationManager::new(DEFAULT_EXCHANGE_URL).await?),
) {
*self = OAuthState::AuthorizedHttpClient(AuthorizedHttpClient::new( *self = OAuthState::AuthorizedHttpClient(AuthorizedHttpClient::new(
Arc::new(manager), Arc::new(manager),
None, None,
@ -2622,8 +2798,7 @@ impl OAuthState {
required_scope: &str, required_scope: &str,
redirect_uri: &str, redirect_uri: &str,
) -> Result<String, AuthError> { ) -> Result<String, AuthError> {
let placeholder = let placeholder = self.placeholder().await?;
OAuthState::Authorized(AuthorizationManager::new(DEFAULT_EXCHANGE_URL).await?);
let old = std::mem::replace(self, placeholder); let old = std::mem::replace(self, placeholder);
let OAuthState::Authorized(manager) = old else { let OAuthState::Authorized(manager) = old else {
*self = old; *self = old;
@ -2755,10 +2930,8 @@ impl OAuthState {
&mut self, &mut self,
config: ClientCredentialsConfig, config: ClientCredentialsConfig,
) -> Result<(), AuthError> { ) -> Result<(), AuthError> {
let OAuthState::Unauthorized(mut manager) = std::mem::replace( let placeholder = self.placeholder().await?;
self, let OAuthState::Unauthorized(mut manager) = std::mem::replace(self, placeholder) else {
OAuthState::Unauthorized(AuthorizationManager::new(DEFAULT_EXCHANGE_URL).await?),
) else {
return Err(AuthError::InternalError( return Err(AuthError::InternalError(
"Client credentials flow requires Unauthorized state".to_string(), "Client credentials flow requires Unauthorized state".to_string(),
)); ));
@ -2784,18 +2957,228 @@ impl OAuthState {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use std::{collections::HashMap, sync::Arc}; use std::{
collections::{HashMap, VecDeque},
sync::{Arc, Mutex as StdMutex},
};
use oauth2::{AuthType, CsrfToken, PkceCodeVerifier}; use oauth2::{AuthType, CsrfToken, HttpResponse, PkceCodeVerifier};
use url::Url; use url::Url;
use super::{ use super::{
AuthError, AuthorizationCallback, AuthorizationManager, AuthorizationMetadata, AuthError, AuthorizationCallback, AuthorizationManager, AuthorizationMetadata,
InMemoryStateStore, OAuthClientConfig, ScopeUpgradeConfig, StateStore, InMemoryStateStore, OAuthClientConfig, OAuthHttpClient, OAuthHttpClientError,
StoredAuthorizationState, is_https_url, OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, ScopeUpgradeConfig,
StateStore, StoredAuthorizationState, is_https_url,
}; };
use crate::transport::auth::VendorExtraTokenFields; use crate::transport::auth::VendorExtraTokenFields;
#[derive(Clone, Debug, PartialEq, Eq)]
struct RecordedOAuthRequest {
method: String,
uri: String,
redirect_policy: OAuthHttpRedirectPolicy,
body: Vec<u8>,
}
#[derive(Clone, Default)]
struct RecordingOAuthHttpClient {
requests: Arc<StdMutex<Vec<RecordedOAuthRequest>>>,
responses: Arc<StdMutex<VecDeque<HttpResponse>>>,
}
impl RecordingOAuthHttpClient {
fn with_responses(responses: Vec<HttpResponse>) -> Self {
Self {
responses: Arc::new(StdMutex::new(responses.into())),
..Default::default()
}
}
fn requests(&self) -> Vec<RecordedOAuthRequest> {
self.requests.lock().unwrap().clone()
}
}
impl OAuthHttpClient for RecordingOAuthHttpClient {
fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> {
self.requests.lock().unwrap().push(RecordedOAuthRequest {
method: request.request.method().to_string(),
uri: request.request.uri().to_string(),
redirect_policy: request.redirect_policy,
body: request.request.body().clone(),
});
let response = self.responses.lock().unwrap().pop_front();
Box::pin(async move {
response.ok_or_else(|| OAuthHttpClientError::new("missing fake response"))
})
}
}
fn http_response(status: u16, body: serde_json::Value) -> HttpResponse {
oauth2::http::Response::builder()
.status(status)
.body(serde_json::to_vec(&body).unwrap())
.unwrap()
}
#[tokio::test]
async fn custom_http_client_handles_protected_resource_discovery() {
let challenge = oauth2::http::Response::builder()
.status(401)
.header(
"www-authenticate",
r#"Bearer resource_metadata="https://mcp.example.com/.well-known/oauth-protected-resource""#,
)
.body(Vec::new())
.unwrap();
let client = RecordingOAuthHttpClient::with_responses(vec![
challenge,
http_response(
200,
serde_json::json!({
"authorization_servers": ["https://auth.example.com"]
}),
),
http_response(
200,
serde_json::json!({
"authorization_endpoint": "https://auth.example.com/authorize",
"token_endpoint": "https://auth.example.com/token"
}),
),
]);
let manager = AuthorizationManager::new_with_oauth_http_client(
"https://mcp.example.com/mcp",
Arc::new(client.clone()),
)
.await
.unwrap();
let metadata = manager.discover_metadata().await.unwrap();
assert_eq!(metadata.token_endpoint, "https://auth.example.com/token");
assert_eq!(
client.requests(),
vec![
RecordedOAuthRequest {
method: "GET".to_string(),
uri: "https://mcp.example.com/mcp".to_string(),
redirect_policy: OAuthHttpRedirectPolicy::Follow,
body: Vec::new(),
},
RecordedOAuthRequest {
method: "GET".to_string(),
uri: "https://mcp.example.com/.well-known/oauth-protected-resource".to_string(),
redirect_policy: OAuthHttpRedirectPolicy::Follow,
body: Vec::new(),
},
RecordedOAuthRequest {
method: "GET".to_string(),
uri: "https://auth.example.com/.well-known/oauth-authorization-server"
.to_string(),
redirect_policy: OAuthHttpRedirectPolicy::Follow,
body: Vec::new(),
},
]
);
}
#[tokio::test]
async fn custom_http_client_handles_registration_exchange_and_refresh() {
let client = RecordingOAuthHttpClient::with_responses(vec![
http_response(
201,
serde_json::json!({
"client_id": "test-client",
"redirect_uris": ["http://localhost/callback"]
}),
),
http_response(
200,
serde_json::json!({
"access_token": "access-1",
"token_type": "bearer",
"refresh_token": "refresh-1",
"expires_in": 3600
}),
),
http_response(
200,
serde_json::json!({
"access_token": "access-2",
"token_type": "bearer",
"refresh_token": "refresh-2",
"expires_in": 3600
}),
),
]);
let mut manager = AuthorizationManager::new_with_oauth_http_client(
"https://mcp.example.com/mcp",
Arc::new(client.clone()),
)
.await
.unwrap();
manager.set_metadata(AuthorizationMetadata {
authorization_endpoint: "https://auth.example.com/authorize".to_string(),
token_endpoint: "https://auth.example.com/token".to_string(),
registration_endpoint: Some("https://auth.example.com/register".to_string()),
response_types_supported: Some(vec!["code".to_string()]),
..Default::default()
});
manager
.register_client(
"Codex",
"http://localhost/callback",
&["profile", "offline_access"],
)
.await
.unwrap();
let authorization_url = manager
.get_authorization_url(&["profile", "offline_access"])
.await
.unwrap();
let state = Url::parse(&authorization_url)
.unwrap()
.query_pairs()
.find(|(name, _)| name == "state")
.unwrap()
.1
.into_owned();
manager
.exchange_code_for_token("authorization-code", &state)
.await
.unwrap();
manager.refresh_token().await.unwrap();
let requests = client.requests();
let registration: serde_json::Value = serde_json::from_slice(&requests[0].body).unwrap();
assert_eq!(registration["scope"], "profile offline_access");
assert_eq!(
requests
.iter()
.map(|request| request.uri.as_str())
.collect::<Vec<_>>(),
vec![
"https://auth.example.com/register",
"https://auth.example.com/token",
"https://auth.example.com/token",
]
);
assert_eq!(
requests
.iter()
.map(|request| request.redirect_policy)
.collect::<Vec<_>>(),
vec![
OAuthHttpRedirectPolicy::Follow,
OAuthHttpRedirectPolicy::Stop,
OAuthHttpRedirectPolicy::Stop,
]
);
}
// -- url helpers -- // -- url helpers --
#[test] #[test]
@ -4420,6 +4803,74 @@ mod tests {
); );
} }
#[tokio::test]
async fn refresh_token_uses_client_configured_by_with_client() {
use axum::{Router, body::Body, http::Response, routing::post};
let received_header = Arc::new(std::sync::Mutex::new(None));
let received_header_clone = Arc::clone(&received_header);
let app = Router::new().route(
"/token",
post(move |headers: axum::http::HeaderMap| {
let received_header = Arc::clone(&received_header_clone);
async move {
*received_header.lock().unwrap() = headers
.get("x-custom-client")
.and_then(|value| value.to_str().ok())
.map(str::to_string);
Response::builder()
.status(200)
.header("content-type", "application/json")
.body(Body::from(
r#"{"access_token":"new-token","token_type":"Bearer","expires_in":3600}"#,
))
.unwrap()
}
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
let mut manager = manager_with_metadata(Some(AuthorizationMetadata {
authorization_endpoint: format!("http://{addr}/authorize"),
token_endpoint: format!("http://{addr}/token"),
..Default::default()
}))
.await;
let mut default_headers = reqwest::header::HeaderMap::new();
default_headers.insert("x-custom-client", "configured".parse().unwrap());
manager
.with_client(
reqwest::Client::builder()
.default_headers(default_headers)
.build()
.unwrap(),
)
.unwrap();
manager.configure_client(test_client_config()).unwrap();
manager
.credential_store
.save(StoredCredentials {
client_id: "my-client".to_string(),
token_response: Some(make_token_response_with_refresh(
"old-token",
"my-refresh-token",
)),
granted_scopes: vec![],
token_received_at: Some(AuthorizationManager::now_epoch_secs()),
})
.await
.unwrap();
manager.refresh_token().await.unwrap();
assert_eq!(
received_header.lock().unwrap().as_deref(),
Some("configured")
);
}
async fn start_token_server() -> (String, Arc<std::sync::Mutex<Option<String>>>) { async fn start_token_server() -> (String, Arc<std::sync::Mutex<Option<String>>>) {
use axum::{Router, body::Body, http::Response, routing::post}; use axum::{Router, body::Body, http::Response, routing::post};
let captured: Arc<std::sync::Mutex<Option<String>>> = Arc::new(std::sync::Mutex::new(None)); let captured: Arc<std::sync::Mutex<Option<String>>> = Arc::new(std::sync::Mutex::new(None));