use anda_core::{BoxError, Json};
use async_trait::async_trait;
use parking_lot::RwLock;
use reqwest::Client as ReqwestClient;
use rmcp::{
service::ClientInitializeError,
transport::{
AuthError, AuthorizationManager, ClientCredentialsConfig, CredentialStore,
StoredCredentials,
auth::{AuthorizationCallback, AuthorizationMetadataSource, OAuthClientConfig, OAuthState},
},
};
use serde::{Deserialize, Serialize};
use std::{
collections::HashMap,
sync::Arc,
time::{Duration, Instant},
};
use super::McpServerConfig;
use super::session::McpTransportConfig;
pub(crate) const CLIENT_CREDENTIALS_RENEW_BUFFER: Duration = Duration::from_secs(120);
pub(crate) const CLIENT_CREDENTIALS_MIN_RENEW_BUFFER: Duration = Duration::from_secs(31);
pub(crate) fn client_credentials_deadline(now: Instant, ttl: Duration) -> Option<Instant> {
let buffer = if ttl >= CLIENT_CREDENTIALS_RENEW_BUFFER.saturating_mul(2) {
CLIENT_CREDENTIALS_RENEW_BUFFER
} else {
(ttl / 2).max(CLIENT_CREDENTIALS_MIN_RENEW_BUFFER).min(ttl)
};
now.checked_add(ttl.saturating_sub(buffer))
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(tag = "flow", rename_all = "snake_case")]
pub enum McpOAuthConfig {
AuthorizationCode(OAuthAuthorizationCodeConfig),
ClientCredentials(OAuthClientCredentialsConfig),
}
impl McpOAuthConfig {
pub(crate) fn validate(&self) -> Result<(), BoxError> {
match self {
Self::AuthorizationCode(config) => config.validate(),
Self::ClientCredentials(config) => config.validate(),
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct OAuthAuthorizationCodeConfig {
pub redirect_uri: String,
#[serde(default)]
pub scopes: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub client_name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub client_id: Option<String>,
}
impl OAuthAuthorizationCodeConfig {
pub(crate) fn validate(&self) -> Result<(), BoxError> {
if self.redirect_uri.trim().is_empty() {
return Err("MCP OAuth authorization_code redirect_uri must not be empty".into());
}
Ok(())
}
}
#[derive(Clone, Deserialize, Serialize)]
pub struct OAuthClientCredentialsConfig {
pub client_id: String,
pub client_secret: String,
#[serde(default)]
pub scopes: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub resource: Option<String>,
}
impl std::fmt::Debug for OAuthClientCredentialsConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OAuthClientCredentialsConfig")
.field("client_id", &self.client_id)
.field("client_secret", &"[REDACTED]")
.field("scopes", &self.scopes)
.field("resource", &self.resource)
.finish()
}
}
impl OAuthClientCredentialsConfig {
pub(crate) fn validate(&self) -> Result<(), BoxError> {
if self.client_id.trim().is_empty() {
return Err("MCP OAuth client_credentials client_id must not be empty".into());
}
if self.client_secret.trim().is_empty() {
return Err("MCP OAuth client_credentials client_secret must not be empty".into());
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct McpOAuthMetadata {
pub scopes_supported: Vec<String>,
pub registration_supported: bool,
}
#[async_trait]
pub trait McpCredentialStore: Send + Sync {
async fn load(&self, server_id: &str) -> Result<Option<StoredCredentials>, BoxError>;
async fn save(&self, server_id: &str, credentials: StoredCredentials) -> Result<(), BoxError>;
async fn clear(&self, server_id: &str) -> Result<(), BoxError>;
async fn acquire_refresh_guard(
&self,
_server_id: &str,
) -> Result<Option<rmcp::transport::auth::CredentialRefreshGuard>, BoxError> {
Ok(None)
}
}
#[derive(Debug, Default)]
pub struct InMemoryMcpCredentialStore {
credentials: RwLock<HashMap<String, StoredCredentials>>,
refresh_locks: RwLock<HashMap<String, Arc<tokio::sync::Mutex<()>>>>,
}
impl InMemoryMcpCredentialStore {
pub fn new() -> Self {
Self::default()
}
}
#[async_trait]
impl McpCredentialStore for InMemoryMcpCredentialStore {
async fn acquire_refresh_guard(
&self,
server_id: &str,
) -> Result<Option<rmcp::transport::auth::CredentialRefreshGuard>, BoxError> {
let lock = self
.refresh_locks
.write()
.entry(server_id.to_string())
.or_default()
.clone();
Ok(Some(rmcp::transport::auth::CredentialRefreshGuard::new(
lock.lock_owned().await,
)))
}
async fn load(&self, server_id: &str) -> Result<Option<StoredCredentials>, BoxError> {
Ok(self.credentials.read().get(server_id).cloned())
}
async fn save(&self, server_id: &str, credentials: StoredCredentials) -> Result<(), BoxError> {
self.credentials
.write()
.insert(server_id.to_string(), credentials);
Ok(())
}
async fn clear(&self, server_id: &str) -> Result<(), BoxError> {
self.credentials.write().remove(server_id);
Ok(())
}
}
pub(crate) struct ScopedCredentialStore {
pub(crate) server_id: String,
pub(crate) inner: Arc<dyn McpCredentialStore>,
}
#[async_trait]
impl CredentialStore for ScopedCredentialStore {
async fn acquire_refresh_guard(
&self,
) -> Result<Option<rmcp::transport::auth::CredentialRefreshGuard>, AuthError> {
self.inner
.acquire_refresh_guard(&self.server_id)
.await
.map_err(|err| AuthError::CredentialStoreError(err.to_string()))
}
async fn load(&self) -> Result<Option<StoredCredentials>, AuthError> {
self.inner
.load(&self.server_id)
.await
.map_err(|err| AuthError::InternalError(err.to_string()))
}
async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError> {
self.inner
.save(&self.server_id, credentials)
.await
.map_err(|err| AuthError::InternalError(err.to_string()))
}
async fn clear(&self) -> Result<(), AuthError> {
self.inner
.clear(&self.server_id)
.await
.map_err(|err| AuthError::InternalError(err.to_string()))
}
}
#[derive(Debug, Clone)]
pub struct McpAuthorizationRequired {
pub server_id: String,
}
impl std::fmt::Display for McpAuthorizationRequired {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"MCP server {} requires interactive OAuth authorization; \
call begin_authorization/complete_authorization first",
self.server_id
)
}
}
impl std::error::Error for McpAuthorizationRequired {}
pub(crate) fn is_authorization_error(err: &(dyn std::error::Error + 'static)) -> bool {
use rmcp::transport::streamable_http_client::StreamableHttpError;
if let Some(rmcp::service::ServiceError::TransportSend(transport)) =
err.downcast_ref::<rmcp::service::ServiceError>()
{
return is_authorization_error(transport.error.as_ref());
}
if let Some(error) = err.downcast_ref::<ClientInitializeError>() {
match error {
ClientInitializeError::TransportError { error, .. } => {
return is_authorization_error(error.error.as_ref());
}
ClientInitializeError::LegacyFallbackFailed { fallback, .. } => {
return is_authorization_error(fallback.as_ref());
}
_ => {}
}
}
if let Some(error) = err.downcast_ref::<StreamableHttpError<super::http_client::HttpError>>() {
return matches!(
error,
StreamableHttpError::AuthRequired(_)
| StreamableHttpError::InsufficientScope(_)
| StreamableHttpError::Client(super::http_client::HttpError::Status(401 | 403))
| StreamableHttpError::Auth(
AuthError::AuthorizationRequired | AuthError::TokenRefreshRejected(_)
)
);
}
if err.is::<McpAuthorizationRequired>() {
return true;
}
if matches!(
err.downcast_ref::<AuthError>(),
Some(AuthError::AuthorizationRequired | AuthError::TokenRefreshRejected(_))
) {
return true;
}
err.downcast_ref::<ClientInitializeError>()
.is_some_and(ClientInitializeError::is_authorization_required)
}
pub(crate) fn authorization_required_hint(config: &McpServerConfig, err: BoxError) -> BoxError {
if err.is::<McpAuthorizationRequired>() || !is_authorization_error(err.as_ref()) {
return err;
}
let McpTransportConfig::StreamableHttp(http) = &config.transport else {
return err;
};
if !matches!(&http.auth, Some(McpOAuthConfig::AuthorizationCode(_))) {
return err;
}
log::warn!(
"MCP server {}: stored authorization is no longer usable ({err}); interactive \
authorization must run again",
config.id
);
McpAuthorizationRequired {
server_id: config.id.clone(),
}
.into()
}
pub(crate) async fn discover_http_oauth(url: &str) -> Result<Option<McpOAuthMetadata>, BoxError> {
let manager = AuthorizationManager::new(url).await?;
let resolution = manager.resolve_metadata().await?;
if resolution.source == AuthorizationMetadataSource::LegacyEndpointFallback {
return Ok(None);
}
let metadata = resolution.metadata;
Ok(Some(McpOAuthMetadata {
scopes_supported: metadata.scopes_supported.unwrap_or_default(),
registration_supported: metadata.registration_endpoint.is_some(),
}))
}
pub(crate) async fn begin_authorization_manager(
url: &str,
ac: &OAuthAuthorizationCodeConfig,
store: ScopedCredentialStore,
) -> Result<(AuthorizationManager, String), BoxError> {
let mut manager = AuthorizationManager::new(url).await?;
manager.set_credential_store(store);
let metadata = manager.resolve_metadata().await?.metadata;
manager.set_metadata(metadata);
let scope_refs: Vec<&str> = ac.scopes.iter().map(String::as_str).collect();
let client_config = match &ac.client_id {
Some(client_id) => {
let mut cfg = OAuthClientConfig::new(client_id.clone(), ac.redirect_uri.clone());
if !ac.scopes.is_empty() {
cfg = cfg.with_scopes(ac.scopes.clone());
}
cfg
}
None => {
manager
.register_client(
ac.client_name.as_deref().unwrap_or("Anda Engine MCP Host"),
&ac.redirect_uri,
&scope_refs,
)
.await?
}
};
manager.configure_client(client_config)?;
let auth_url = manager.get_authorization_url(&scope_refs).await?;
Ok((manager, auth_url))
}
pub(crate) async fn complete_authorization_exchange(
manager: AuthorizationManager,
redirect_url: &str,
) -> Result<(), BoxError> {
let callback = AuthorizationCallback::from_redirect_url(redirect_url)?;
manager
.exchange_code_for_token_with_issuer(
&callback.code,
&callback.csrf_token,
callback.issuer.as_deref(),
)
.await?;
Ok(())
}
pub(crate) async fn authorize_from_store(
server_id: &str,
url: &str,
store: ScopedCredentialStore,
) -> Result<AuthorizationManager, BoxError> {
let mut manager = AuthorizationManager::new(url).await?;
manager.set_credential_store(store);
if !manager.initialize_from_store().await? {
return Err(McpAuthorizationRequired {
server_id: server_id.to_string(),
}
.into());
}
Ok(manager)
}
pub(crate) async fn authorize_client_credentials(
url: &str,
config: &OAuthClientCredentialsConfig,
) -> Result<(AuthorizationManager, Option<Instant>), BoxError> {
let mut state = OAuthState::new(url, Some(ReqwestClient::new())).await?;
state
.authenticate_client_credentials(ClientCredentialsConfig::ClientSecret {
client_id: config.client_id.clone(),
client_secret: config.client_secret.clone(),
scopes: config.scopes.clone(),
resource: config.resource.clone(),
})
.await?;
let expires_at = match state.get_credentials().await {
Ok((_, Some(token))) => serde_json::to_value(&token)
.ok()
.and_then(|token| token.get("expires_in").and_then(Json::as_u64))
.and_then(|secs| {
client_credentials_deadline(Instant::now(), Duration::from_secs(secs))
}),
_ => None,
};
let manager = state
.into_authorization_manager()
.ok_or("MCP client_credentials authorization did not complete")?;
Ok((manager, expires_at))
}