use crate::http_utils::HttpRequestError;
use crate::{
http::HttpClient,
jwt::{DecodeJwtError, DecodedJwt, decode_jwt},
};
use base64::Engine;
use jsonwebtoken::{self as jwt, Algorithm, DecodingKey, Validation};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashSet;
use thiserror::Error;
#[derive(Debug, Clone)]
#[allow(clippy::struct_excessive_bools)]
pub(crate) struct SsaValidationConfig {
pub required_claims: HashSet<String>,
pub allowed_algorithms: HashSet<Algorithm>,
pub validate_expiration: bool,
pub validate_issued_at: bool,
pub validate_not_before: bool,
pub validate_audience: bool,
pub validate_grant_types_array: bool,
pub validate_software_roles_array: bool,
}
impl Default for SsaValidationConfig {
fn default() -> Self {
Self {
required_claims: HashSet::from([
"software_id".to_string(),
"grant_types".to_string(),
"org_id".to_string(),
"iss".to_string(),
"software_roles".to_string(),
"exp".to_string(),
"iat".to_string(),
"jti".to_string(),
]),
allowed_algorithms: HashSet::new(), validate_expiration: true,
validate_issued_at: true,
validate_not_before: false,
validate_audience: false,
validate_grant_types_array: true,
validate_software_roles_array: true,
}
}
}
#[derive(Debug, Deserialize, Serialize, PartialEq)]
pub(crate) struct SsaClaims {
pub software_id: String,
pub grant_types: Vec<String>,
pub org_id: String,
pub iss: String,
pub software_roles: Vec<String>,
pub exp: u64,
pub iat: u64,
pub jti: String,
#[serde(flatten)]
pub additional_claims: Value,
}
pub(crate) async fn validate_ssa_jwt(
ssa_jwt: &str,
jwks_uri: &str,
client: &HttpClient,
) -> Result<SsaClaims, SsaValidationError> {
let config = SsaValidationConfig {
..Default::default()
};
validate_ssa_jwt_with_config(ssa_jwt, jwks_uri, &config, client).await
}
pub(crate) async fn validate_ssa_jwt_with_config(
ssa_jwt: &str,
jwks_uri: &str,
config: &SsaValidationConfig,
client: &HttpClient,
) -> Result<SsaClaims, SsaValidationError> {
let decoded_jwt = decode_jwt(ssa_jwt)?;
validate_ssa_structure_with_config(&decoded_jwt, config)?;
if !config.allowed_algorithms.is_empty()
&& !config.allowed_algorithms.contains(&decoded_jwt.header.alg)
{
return Err(SsaValidationError::AlgorithmNotAllowed(
decoded_jwt.header.alg,
));
}
let jwks = fetch_jwks(jwks_uri, client).await?;
let decoding_key = find_decoding_key(&decoded_jwt, &jwks)?;
let claims = validate_ssa_signature_and_claims_with_config(
ssa_jwt,
&decoding_key,
decoded_jwt.header.alg,
config,
)?;
Ok(claims)
}
pub(crate) fn validate_ssa_structure_with_config(
decoded_jwt: &DecodedJwt,
config: &SsaValidationConfig,
) -> Result<(), SsaValidationError> {
let claims = &decoded_jwt.claims.inner;
let mut missing_claims = Vec::new();
for claim in &config.required_claims {
if claims.get(claim).is_none() {
missing_claims.push(claim.clone());
}
}
if !missing_claims.is_empty() {
return Err(SsaValidationError::MissingRequiredClaims(missing_claims));
}
if config.validate_grant_types_array
&& let Some(grant_types) = claims.get("grant_types")
&& !grant_types.is_array()
{
return Err(SsaValidationError::InvalidGrantTypes);
}
if config.validate_software_roles_array
&& let Some(software_roles) = claims.get("software_roles")
&& !software_roles.is_array()
{
return Err(SsaValidationError::InvalidSoftwareRoles);
}
if config.validate_expiration
&& let Some(exp) = claims.get("exp")
&& !exp.is_number()
{
return Err(SsaValidationError::InvalidExpirationTime);
}
if config.validate_issued_at
&& let Some(iat) = claims.get("iat")
&& !iat.is_number()
{
return Err(SsaValidationError::InvalidIssuedAtTime);
}
Ok(())
}
async fn fetch_jwks(jwks_uri: &str, client: &HttpClient) -> Result<Value, SsaValidationError> {
let jwks: Value = client
.get_json(jwks_uri)
.await
.map_err(SsaValidationError::JwksFetchError)?;
Ok(jwks)
}
fn find_decoding_key(
decoded_jwt: &DecodedJwt,
jwks: &Value,
) -> Result<DecodingKey, SsaValidationError> {
let kid = decoded_jwt
.header
.kid
.as_ref()
.ok_or(SsaValidationError::MissingKeyId)?;
if let Some(keys) = jwks.get("keys").and_then(|k| k.as_array()) {
for key in keys {
if let Some(key_kid) = key.get("kid").and_then(|k| k.as_str())
&& key_kid == kid
{
if let Some(jwk_alg) = key.get("alg").and_then(|a| a.as_str()) {
let jwt_alg_str = match decoded_jwt.header.alg {
Algorithm::HS256 => "HS256",
Algorithm::HS384 => "HS384",
Algorithm::HS512 => "HS512",
Algorithm::RS256 => "RS256",
Algorithm::RS384 => "RS384",
Algorithm::RS512 => "RS512",
Algorithm::ES256 => "ES256",
Algorithm::ES384 => "ES384",
Algorithm::PS256 => "PS256",
Algorithm::PS384 => "PS384",
Algorithm::PS512 => "PS512",
Algorithm::EdDSA => "EdDSA",
};
if jwk_alg != jwt_alg_str {
return Err(SsaValidationError::AlgorithmMismatch {
jwt_alg: jwt_alg_str.to_string(),
jwk_alg: jwk_alg.to_string(),
});
}
}
return create_decoding_key(key, decoded_jwt.header.alg);
}
}
}
Err(SsaValidationError::KeyNotFound(kid.clone()))
}
fn create_decoding_key(
jwk: &Value,
algorithm: Algorithm,
) -> Result<DecodingKey, SsaValidationError> {
match algorithm {
Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512 => {
if let Some(k) = jwk.get("k").and_then(|k| k.as_str()) {
let key_data = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(k)
.map_err(SsaValidationError::KeyDecodeError)?;
Ok(DecodingKey::from_secret(&key_data))
} else {
Err(SsaValidationError::InvalidKeyFormat)
}
},
Algorithm::RS256 | Algorithm::RS384 | Algorithm::RS512 => {
if let (Some(n), Some(e)) = (
jwk.get("n").and_then(|n| n.as_str()),
jwk.get("e").and_then(|e| e.as_str()),
) {
Ok(DecodingKey::from_rsa_components(n, e).map_err(SsaValidationError::JwtError)?)
} else {
Err(SsaValidationError::InvalidKeyFormat)
}
},
_ => Err(SsaValidationError::UnsupportedAlgorithm),
}
}
fn validate_ssa_signature_and_claims_with_config(
ssa_jwt: &str,
decoding_key: &DecodingKey,
algorithm: Algorithm,
config: &SsaValidationConfig,
) -> Result<SsaClaims, SsaValidationError> {
let mut validation = Validation::new(algorithm);
validation.validate_exp = config.validate_expiration;
validation.validate_nbf = config.validate_not_before;
validation.validate_aud = config.validate_audience;
validation.required_spec_claims.clear();
let token_data = jwt::decode::<SsaClaims>(ssa_jwt, decoding_key, &validation)
.map_err(SsaValidationError::JwtError)?;
Ok(token_data.claims)
}
#[derive(Debug, Error)]
pub(crate) enum SsaValidationError {
#[error("failed to decode JWT: {0}")]
DecodeJwt(#[from] DecodeJwtError),
#[error("missing required claims: {0:?}")]
MissingRequiredClaims(Vec<String>),
#[error("grant_types must be an array")]
InvalidGrantTypes,
#[error("software_roles must be an array")]
InvalidSoftwareRoles,
#[error("exp must be a number")]
InvalidExpirationTime,
#[error("iat must be a number")]
InvalidIssuedAtTime,
#[error("failed to fetch JWKS (network/HTTP error): {0}")]
JwksFetchError(HttpRequestError),
#[error("missing key ID (kid) in JWT header")]
MissingKeyId,
#[error("key not found in JWKS: {0}")]
KeyNotFound(String),
#[error("failed to decode key: {0}")]
KeyDecodeError(#[from] base64::DecodeError),
#[error("invalid key format")]
InvalidKeyFormat,
#[error("unsupported algorithm")]
UnsupportedAlgorithm,
#[error("algorithm not allowed by configuration: {0:?}")]
AlgorithmNotAllowed(Algorithm),
#[error("algorithm mismatch: JWT header specifies {jwt_alg:?} but JWK specifies {jwk_alg:?}")]
AlgorithmMismatch {
jwt_alg: String,
jwk_alg: String,
},
#[error("JWT validation error: {0}")]
JwtError(#[from] jwt::errors::Error),
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_validate_ssa_structure_valid() {
let claims = json!({
"software_id": "test_software",
"grant_types": ["client_credentials"],
"org_id": "test_org",
"iss": "https://test.issuer.com",
"software_roles": ["cedarling"],
"exp": 1_735_689_600,
"iat": 1_735_603_200,
"jti": "test-jti-123"
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let config = SsaValidationConfig::default();
let result = validate_ssa_structure_with_config(&decoded_jwt, &config);
assert!(result.is_ok());
}
#[test]
fn test_validate_ssa_structure_missing_claims() {
let claims = json!({
"software_id": "test_software",
"grant_types": ["client_credentials"],
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let config = SsaValidationConfig::default();
let result = validate_ssa_structure_with_config(&decoded_jwt, &config);
assert!(matches!(
result,
Err(SsaValidationError::MissingRequiredClaims(_))
));
}
#[test]
fn test_validate_ssa_structure_invalid_grant_types() {
let claims = json!({
"software_id": "test_software",
"grant_types": "client_credentials", "org_id": "test_org",
"iss": "https://test.issuer.com",
"software_roles": ["cedarling"],
"exp": 1_735_689_600,
"iat": 1_735_603_200,
"jti": "test-jti-123"
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let config = SsaValidationConfig::default();
let result = validate_ssa_structure_with_config(&decoded_jwt, &config);
assert!(matches!(result, Err(SsaValidationError::InvalidGrantTypes)));
}
#[test]
fn test_algorithm_mismatch_detection() {
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: json!({}) },
};
let jwks = json!({
"keys": [
{
"kid": "test-kid",
"alg": "HS256",
"k": "dGVzdC1rZXk=" }
]
});
let result = find_decoding_key(&decoded_jwt, &jwks);
assert!(
matches!(result, Err(SsaValidationError::AlgorithmMismatch { jwt_alg, jwk_alg }) if jwt_alg == "RS256" && jwk_alg == "HS256")
);
}
#[test]
fn test_algorithm_match_success() {
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: json!({}) },
};
let jwks = json!({
"keys": [
{
"kid": "test-kid",
"alg": "RS256",
"n": "test-n",
"e": "AQAB"
}
]
});
let result = find_decoding_key(&decoded_jwt, &jwks);
assert!(!matches!(
result,
Err(SsaValidationError::AlgorithmMismatch { .. })
));
}
#[test]
fn test_validate_ssa_structure_invalid_exp_type() {
let claims = json!({
"software_id": "test_software",
"grant_types": ["client_credentials"],
"org_id": "test_org",
"iss": "https://test.issuer.com",
"software_roles": ["cedarling"],
"exp": "not_a_number", "iat": 1_735_603_200,
"jti": "test-jti-123"
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let result =
validate_ssa_structure_with_config(&decoded_jwt, &SsaValidationConfig::default());
assert!(matches!(
result,
Err(SsaValidationError::InvalidExpirationTime)
));
}
#[test]
fn test_validate_ssa_structure_invalid_iat_type() {
let claims = json!({
"software_id": "test_software",
"grant_types": ["client_credentials"],
"org_id": "test_org",
"iss": "https://test.issuer.com",
"software_roles": ["cedarling"],
"exp": 1_735_689_600,
"iat": ["not", "a", "number"], "jti": "test-jti-123"
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let result =
validate_ssa_structure_with_config(&decoded_jwt, &SsaValidationConfig::default());
assert!(matches!(
result,
Err(SsaValidationError::InvalidIssuedAtTime)
));
}
#[test]
fn test_validate_ssa_structure_invalid_software_roles() {
let claims = json!({
"software_id": "test_software",
"grant_types": ["client_credentials"],
"org_id": "test_org",
"iss": "https://test.issuer.com",
"software_roles": "cedarling", "exp": 1_735_689_600,
"iat": 1_735_603_200,
"jti": "test-jti-123"
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let result =
validate_ssa_structure_with_config(&decoded_jwt, &SsaValidationConfig::default());
assert!(matches!(
result,
Err(SsaValidationError::InvalidSoftwareRoles)
));
}
#[test]
fn test_create_decoding_key_rsa_missing_n() {
let jwk = json!({
"kid": "test-kid",
"alg": "RS256",
"e": "AQAB"
});
let result = create_decoding_key(&jwk, Algorithm::RS256);
assert!(matches!(result, Err(SsaValidationError::InvalidKeyFormat)));
}
#[test]
fn test_create_decoding_key_rsa_missing_e() {
let jwk = json!({
"kid": "test-kid",
"alg": "RS256",
"n": "test-n"
});
let result = create_decoding_key(&jwk, Algorithm::RS256);
assert!(matches!(result, Err(SsaValidationError::InvalidKeyFormat)));
}
#[test]
fn test_create_decoding_key_hmac_missing_k() {
let jwk = json!({
"kid": "test-kid",
"alg": "HS256"
});
let result = create_decoding_key(&jwk, Algorithm::HS256);
assert!(matches!(result, Err(SsaValidationError::InvalidKeyFormat)));
}
#[test]
fn test_create_decoding_key_unsupported_algorithm() {
let jwk = json!({
"kid": "test-kid",
"alg": "ES256",
"x": "test-x",
"y": "test-y"
});
let result = create_decoding_key(&jwk, Algorithm::ES256);
assert!(matches!(
result,
Err(SsaValidationError::UnsupportedAlgorithm)
));
}
#[test]
fn test_find_decoding_key_missing_kid() {
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: None, },
claims: crate::jwt::DecodedJwtClaims { inner: json!({}) },
};
let jwks = json!({
"keys": [
{
"kid": "test-kid",
"alg": "RS256",
"n": "test-n",
"e": "AQAB"
}
]
});
let result = find_decoding_key(&decoded_jwt, &jwks);
assert!(matches!(result, Err(SsaValidationError::MissingKeyId)));
}
#[test]
fn test_find_decoding_key_kid_not_found() {
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("different-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: json!({}) },
};
let jwks = json!({
"keys": [
{
"kid": "test-kid",
"alg": "RS256",
"n": "test-n",
"e": "AQAB"
}
]
});
let result = find_decoding_key(&decoded_jwt, &jwks);
assert!(
matches!(result, Err(SsaValidationError::KeyNotFound(kid)) if kid == "different-kid")
);
}
#[test]
fn test_find_decoding_key_empty_keys_array() {
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: json!({}) },
};
let jwks = json!({
"keys": [] });
let result = find_decoding_key(&decoded_jwt, &jwks);
assert!(matches!(result, Err(SsaValidationError::KeyNotFound(kid)) if kid == "test-kid"));
}
#[test]
fn test_find_decoding_key_missing_keys_field() {
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: json!({}) },
};
let jwks = json!({
});
let result = find_decoding_key(&decoded_jwt, &jwks);
assert!(matches!(result, Err(SsaValidationError::KeyNotFound(kid)) if kid == "test-kid"));
}
#[test]
fn test_config_default() {
let config = SsaValidationConfig::default();
assert!(config.validate_expiration);
assert!(config.validate_issued_at);
assert!(!config.validate_not_before);
assert!(!config.validate_audience);
assert!(config.validate_grant_types_array);
assert!(config.validate_software_roles_array);
assert!(config.allowed_algorithms.is_empty()); assert!(config.required_claims.contains("software_id"));
assert!(config.required_claims.contains("exp"));
}
#[test]
fn test_validate_ssa_structure_with_config_skip_exp() {
let claims = json!({
"software_id": "test_software",
"grant_types": ["client_credentials"],
"org_id": "test_org",
"iss": "https://test.issuer.com",
"software_roles": ["cedarling"],
"jti": "test-jti-123"
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let mut config = SsaValidationConfig::default();
config.required_claims.remove("exp");
config.required_claims.remove("iat");
config.validate_expiration = false;
config.validate_issued_at = false;
let result = validate_ssa_structure_with_config(&decoded_jwt, &config);
assert!(result.is_ok());
}
#[test]
fn test_validate_ssa_structure_with_config_skip_array_validation() {
let claims = json!({
"software_id": "test_software",
"grant_types": "client_credentials", "org_id": "test_org",
"iss": "https://test.issuer.com",
"software_roles": "cedarling", "exp": 1_735_689_600,
"iat": 1_735_603_200,
"jti": "test-jti-123"
});
let decoded_jwt = DecodedJwt {
header: crate::jwt::DecodedJwtHeader {
typ: Some("JWT".to_string()),
alg: Algorithm::RS256,
cty: None,
kid: Some("test-kid".to_string()),
},
claims: crate::jwt::DecodedJwtClaims { inner: claims },
};
let config = SsaValidationConfig {
validate_grant_types_array: false,
validate_software_roles_array: false,
..Default::default()
};
let result = validate_ssa_structure_with_config(&decoded_jwt, &config);
assert!(result.is_ok());
}
}