use super::{JwksCache, JwksClient, StandardClaims};
use jsonwebtoken::{Algorithm, DecodingKey, TokenData, Validation, decode, decode_header};
use std::sync::Arc;
use std::time::{Duration, SystemTime};
use tokio::sync::OnceCell;
use tracing::{debug, error, warn};
use turbomcp_protocol::{Error as McpError, Result as McpResult};
#[derive(Debug, Clone)]
pub struct JwtValidationResult {
pub claims: StandardClaims,
pub algorithm: Algorithm,
pub key_id: Option<String>,
pub issued_at: Option<SystemTime>,
pub expires_at: Option<SystemTime>,
}
pub struct JwtValidator {
expected_issuer: String,
expected_audience: String,
jwks_client: Arc<JwksClient>,
clock_skew_leeway: Duration,
allowed_algorithms: Vec<Algorithm>,
discovered_jwks_uri: OnceCell<String>,
ssrf_validator: Option<Arc<crate::ssrf::SsrfValidator>>,
}
impl std::fmt::Debug for JwtValidator {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("JwtValidator")
.field("expected_issuer", &self.expected_issuer)
.field("expected_audience", &self.expected_audience)
.field("jwks_client", &self.jwks_client)
.field("clock_skew_leeway", &self.clock_skew_leeway)
.field("allowed_algorithms", &self.allowed_algorithms)
.field(
"discovered_jwks_uri",
&self.discovered_jwks_uri.get().map(|_| "<cached>"),
)
.field(
"ssrf_validator",
&self.ssrf_validator.as_ref().map(|_| "<SsrfValidator>"),
)
.finish()
}
}
impl JwtValidator {
#[cfg(feature = "mcp-oidc-discovery")]
async fn discover_jwks_client(
issuer: &str,
ssrf_validator: Arc<crate::ssrf::SsrfValidator>,
) -> McpResult<Arc<JwksClient>> {
let fetcher = crate::discovery::DiscoveryFetcher::new((*ssrf_validator).clone())
.map_err(|e| McpError::internal(format!("Failed to create discovery fetcher: {e}")))?;
let metadata = fetcher.fetch(issuer).await.map_err(|e| {
McpError::authentication(format!(
"Authorization server metadata discovery failed for issuer '{issuer}': {e}"
))
})?;
let jwks_uri = metadata.oauth2().jwks_uri.clone().ok_or_else(|| {
McpError::authentication(format!(
"Authorization server metadata for issuer '{issuer}' has no jwks_uri"
))
})?;
tracing::info!(issuer = issuer, jwks_uri = %jwks_uri, "Discovered JWKS URI via authorization server metadata");
Ok(Arc::new(JwksClient::with_ssrf_validator(
jwks_uri,
ssrf_validator,
)))
}
#[cfg(feature = "mcp-oidc-discovery")]
pub async fn new(expected_issuer: String, expected_audience: String) -> McpResult<Self> {
let ssrf_validator = Arc::new(crate::ssrf::SsrfValidator::default());
Self::new_with_ssrf(expected_issuer, expected_audience, ssrf_validator).await
}
#[cfg(feature = "mcp-oidc-discovery")]
pub async fn new_unchecked(
expected_issuer: String,
expected_audience: String,
) -> McpResult<Self> {
let permissive = crate::ssrf::SsrfPolicy::builder()
.allow_private_networks(true)
.allow_localhost(true)
.allow_link_local(true)
.require_https(false)
.build();
let ssrf_validator = Arc::new(crate::ssrf::SsrfValidator::new(permissive));
let jwks_client =
Self::discover_jwks_client(&expected_issuer, Arc::clone(&ssrf_validator)).await?;
let jwks_uri = jwks_client.jwks_uri().to_string();
Ok(Self {
expected_issuer,
expected_audience,
jwks_client,
clock_skew_leeway: Duration::from_secs(60),
allowed_algorithms: vec![Algorithm::ES256, Algorithm::RS256, Algorithm::PS256],
discovered_jwks_uri: OnceCell::new_with(Some(jwks_uri)),
ssrf_validator: None,
})
}
#[cfg(feature = "mcp-oidc-discovery")]
pub async fn new_with_ssrf(
expected_issuer: String,
expected_audience: String,
ssrf_validator: Arc<crate::ssrf::SsrfValidator>,
) -> McpResult<Self> {
let jwks_client =
Self::discover_jwks_client(&expected_issuer, Arc::clone(&ssrf_validator)).await?;
let jwks_uri = jwks_client.jwks_uri().to_string();
Ok(Self {
expected_issuer,
expected_audience,
jwks_client,
clock_skew_leeway: Duration::from_secs(60),
allowed_algorithms: vec![Algorithm::ES256, Algorithm::RS256, Algorithm::PS256],
discovered_jwks_uri: OnceCell::new_with(Some(jwks_uri)),
ssrf_validator: Some(ssrf_validator),
})
}
pub fn with_jwks_uri(
expected_issuer: String,
expected_audience: String,
jwks_uri: String,
) -> Self {
let jwks_client = Arc::new(JwksClient::new(jwks_uri.clone()));
Self {
expected_issuer,
expected_audience,
jwks_client,
clock_skew_leeway: Duration::from_secs(60),
allowed_algorithms: vec![Algorithm::ES256, Algorithm::RS256, Algorithm::PS256],
discovered_jwks_uri: OnceCell::new_with(Some(jwks_uri)),
ssrf_validator: None,
}
}
pub fn with_jwks_client(
expected_issuer: String,
expected_audience: String,
jwks_client: Arc<JwksClient>,
) -> Self {
Self {
expected_issuer,
expected_audience,
jwks_client,
clock_skew_leeway: Duration::from_secs(60),
allowed_algorithms: vec![Algorithm::ES256, Algorithm::RS256, Algorithm::PS256],
discovered_jwks_uri: OnceCell::new(), ssrf_validator: None,
}
}
pub fn with_ssrf_validator(mut self, ssrf_validator: Arc<crate::ssrf::SsrfValidator>) -> Self {
self.ssrf_validator = Some(ssrf_validator);
self
}
pub fn with_clock_skew(mut self, leeway: Duration) -> Self {
self.clock_skew_leeway = leeway;
self
}
pub fn with_algorithms(mut self, algorithms: Vec<Algorithm>) -> Self {
self.allowed_algorithms = algorithms;
self
}
pub async fn validate(&self, token: &str) -> McpResult<JwtValidationResult> {
let header = decode_header(token).map_err(|e| {
debug!(error = %e, "Failed to decode JWT header");
McpError::invalid_params(format!("Invalid JWT format: {e}"))
})?;
if !self.allowed_algorithms.contains(&header.alg) {
error!(
algorithm = ?header.alg,
allowed = ?self.allowed_algorithms,
"JWT algorithm not allowed"
);
return Err(McpError::invalid_params(format!(
"Algorithm {:?} not allowed",
header.alg
)));
}
let key_id = header.kid.clone().ok_or_else(|| {
error!("JWT missing kid (key ID) in header");
McpError::invalid_params("JWT must include kid (key ID) in header".to_string())
})?;
let decoding_key = self.get_decoding_key(&key_id, header.alg).await?;
let mut validation = Validation::new(header.alg);
validation.set_audience(&[&self.expected_audience]);
validation.set_issuer(&[&self.expected_issuer]);
validation.leeway = self.clock_skew_leeway.as_secs();
validation.validate_nbf = true;
let token_data: TokenData<StandardClaims> = decode(token, &decoding_key, &validation)
.map_err(|e| {
warn!(
error = %e,
issuer = %self.expected_issuer,
audience = %self.expected_audience,
"JWT validation failed"
);
McpError::invalid_params(format!("JWT validation failed: {e}"))
})?;
let issued_at = token_data
.claims
.iat
.and_then(crate::context::checked_system_time);
let expires_at = token_data
.claims
.exp
.and_then(crate::context::checked_system_time);
let sub_hash = match token_data.claims.sub.as_deref() {
Some(sub) => {
use sha2::{Digest, Sha256};
let digest = Sha256::digest(sub.as_bytes());
format!(
"sha256:{:02x}{:02x}{:02x}{:02x}",
digest[0], digest[1], digest[2], digest[3]
)
}
None => "<none>".to_string(),
};
debug!(
issuer = %self.expected_issuer,
audience = %self.expected_audience,
subject_hash = %sub_hash,
algorithm = ?header.alg,
"JWT validation successful"
);
Ok(JwtValidationResult {
claims: token_data.claims,
algorithm: header.alg,
key_id: Some(key_id),
issued_at,
expires_at,
})
}
pub async fn validate_with_refresh(&self, token: &str) -> McpResult<JwtValidationResult> {
match self.validate(token).await {
Ok(result) => Ok(result),
Err(first_error) => {
warn!(
error = %first_error,
"JWT validation failed, refreshing JWKS and retrying"
);
self.jwks_client.refresh().await?;
self.validate(token).await.map_err(|e| {
error!(error = %e, "JWT validation failed after JWKS refresh");
e
})
}
}
}
async fn get_decoding_key(
&self,
key_id: &str,
_algorithm: Algorithm,
) -> McpResult<DecodingKey> {
let jwks = self.jwks_client.get_jwks().await?;
let jwk = jwks.find(key_id).ok_or_else(|| {
error!(key_id = key_id, "Key ID not found in JWKS");
McpError::invalid_params(format!("Key ID '{key_id}' not found in JWKS"))
})?;
DecodingKey::from_jwk(jwk).map_err(|e| {
error!(key_id = key_id, error = %e, "Failed to create decoding key from JWK");
McpError::internal(format!("Invalid JWK: {e}"))
})
}
pub fn expected_issuer(&self) -> &str {
&self.expected_issuer
}
pub fn expected_audience(&self) -> &str {
&self.expected_audience
}
}
#[derive(Debug)]
pub struct MultiIssuerValidator {
expected_audience: String,
validators: std::collections::HashMap<String, Arc<JwtValidator>>,
#[allow(dead_code)]
jwks_cache: Arc<JwksCache>,
}
impl MultiIssuerValidator {
pub fn new(expected_audience: String) -> Self {
Self {
expected_audience,
validators: std::collections::HashMap::new(),
jwks_cache: Arc::new(JwksCache::new()),
}
}
#[cfg(feature = "mcp-oidc-discovery")]
pub async fn add_issuer(&mut self, issuer: String) -> McpResult<()> {
let ssrf_validator = Arc::new(crate::ssrf::SsrfValidator::default());
self.add_issuer_with_ssrf(issuer, ssrf_validator).await
}
#[cfg(feature = "mcp-oidc-discovery")]
pub async fn add_issuer_unchecked(&mut self, issuer: String) -> McpResult<()> {
let permissive = crate::ssrf::SsrfPolicy::builder()
.allow_private_networks(true)
.allow_localhost(true)
.allow_link_local(true)
.require_https(false)
.build();
let ssrf_validator = Arc::new(crate::ssrf::SsrfValidator::new(permissive));
let jwks_client = JwtValidator::discover_jwks_client(&issuer, ssrf_validator).await?;
let validator = Arc::new(JwtValidator::with_jwks_client(
issuer.clone(),
self.expected_audience.clone(),
jwks_client,
));
self.validators.insert(issuer, validator);
Ok(())
}
#[cfg(feature = "mcp-oidc-discovery")]
pub async fn add_issuer_with_ssrf(
&mut self,
issuer: String,
ssrf_validator: Arc<crate::ssrf::SsrfValidator>,
) -> McpResult<()> {
let jwks_client =
JwtValidator::discover_jwks_client(&issuer, Arc::clone(&ssrf_validator)).await?;
let validator = Arc::new(
JwtValidator::with_jwks_client(
issuer.clone(),
self.expected_audience.clone(),
jwks_client,
)
.with_ssrf_validator(ssrf_validator),
);
self.validators.insert(issuer, validator);
Ok(())
}
pub fn add_issuer_with_jwks_uri(&mut self, issuer: String, jwks_uri: String) {
let jwks_client = Arc::new(JwksClient::new(jwks_uri));
let validator = Arc::new(JwtValidator::with_jwks_client(
issuer.clone(),
self.expected_audience.clone(),
jwks_client,
));
self.validators.insert(issuer, validator);
}
pub async fn validate(&self, token: &str) -> McpResult<JwtValidationResult> {
let header = decode_header(token)
.map_err(|e| McpError::invalid_params(format!("Invalid JWT format: {e}")))?;
const ALLOWED_ALGORITHMS: &[Algorithm] = &[
Algorithm::ES256,
Algorithm::ES384,
Algorithm::RS256,
Algorithm::RS384,
Algorithm::RS512,
Algorithm::PS256,
Algorithm::PS384,
Algorithm::PS512,
];
if !ALLOWED_ALGORITHMS.contains(&header.alg) {
error!(algorithm = ?header.alg, "JWT algorithm not in allowlist");
return Err(McpError::invalid_params(format!(
"JWT algorithm {:?} not allowed. Only asymmetric algorithms (ES*, RS*, PS*) are permitted.",
header.alg
)));
}
let parts: Vec<&str> = token.split('.').collect();
if parts.len() != 3 {
return Err(McpError::invalid_params("Invalid JWT format".to_string()));
}
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
let payload = URL_SAFE_NO_PAD
.decode(parts[1])
.map_err(|e| McpError::invalid_params(format!("Invalid JWT payload encoding: {e}")))?;
let claims: StandardClaims = serde_json::from_slice(&payload)
.map_err(|e| McpError::invalid_params(format!("Invalid JWT claims: {e}")))?;
let issuer = claims.iss.ok_or_else(|| {
McpError::invalid_params("JWT missing iss (issuer) claim".to_string())
})?;
let validator = self.validators.get(&issuer).ok_or_else(|| {
error!(issuer = %issuer, "Unknown issuer");
McpError::invalid_params(format!("Issuer '{}' not supported", issuer))
})?;
validator.validate_with_refresh(token).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_jwt_validator_creation_with_jwks_uri() {
let validator = JwtValidator::with_jwks_uri(
"https://auth.example.com".to_string(),
"https://mcp.example.com".to_string(),
"https://auth.example.com/jwks".to_string(),
);
assert_eq!(validator.expected_issuer(), "https://auth.example.com");
assert_eq!(validator.expected_audience(), "https://mcp.example.com");
assert_eq!(validator.clock_skew_leeway, Duration::from_secs(60));
assert_eq!(validator.allowed_algorithms.len(), 3);
}
#[test]
fn test_jwt_validator_custom_clock_skew() {
let validator = JwtValidator::with_jwks_uri(
"https://auth.example.com".to_string(),
"https://mcp.example.com".to_string(),
"https://auth.example.com/jwks".to_string(),
)
.with_clock_skew(Duration::from_secs(30));
assert_eq!(validator.clock_skew_leeway, Duration::from_secs(30));
}
#[test]
fn test_jwt_validator_custom_algorithms() {
let validator = JwtValidator::with_jwks_uri(
"https://auth.example.com".to_string(),
"https://mcp.example.com".to_string(),
"https://auth.example.com/jwks".to_string(),
)
.with_algorithms(vec![Algorithm::ES256]);
assert_eq!(validator.allowed_algorithms, vec![Algorithm::ES256]);
}
#[test]
fn test_multi_issuer_validator_creation() {
let validator = MultiIssuerValidator::new("https://mcp.example.com".to_string());
assert_eq!(validator.expected_audience, "https://mcp.example.com");
assert_eq!(validator.validators.len(), 0);
}
#[test]
fn test_multi_issuer_validator_add_issuer() {
let mut validator = MultiIssuerValidator::new("https://mcp.example.com".to_string());
validator.add_issuer_with_jwks_uri(
"https://auth.example.com".to_string(),
"https://auth.example.com/jwks".to_string(),
);
assert_eq!(validator.validators.len(), 1);
assert!(
validator
.validators
.contains_key("https://auth.example.com")
);
}
}