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:
parent
3c5ce2b0d7
commit
de898dd842
2 changed files with 579 additions and 127 deletions
|
|
@ -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")]
|
||||||
|
|
|
||||||
|
|
@ -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(®istration_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(®istration_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));
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue