use std::fmt;
use base64::Engine as _;
use jsonwebtoken::{
Algorithm, DecodingKey, Validation, decode, decode_header,
jwk::{Jwk, JwkSet},
};
use subtle::ConstantTimeEq as _;
use crate::{
claims::OidcClaims,
error::{ClaimsError, VerifyError},
};
const DEFAULT_ALGORITHMS: &[Algorithm] = &[
Algorithm::RS256,
Algorithm::RS384,
Algorithm::RS512,
Algorithm::PS256,
Algorithm::PS384,
Algorithm::PS512,
Algorithm::ES256,
Algorithm::ES384,
Algorithm::EdDSA,
];
const DEFAULT_LEEWAY_SECS: u64 = 60;
#[derive(Debug)]
pub struct IdTokenVerifier {
issuer: String,
audience: String,
algorithms: Vec<Algorithm>,
leeway_secs: u64,
jwks: JwkSet,
}
impl IdTokenVerifier {
pub fn new(issuer: impl Into<String>, audience: impl Into<String>, jwks: JwkSet) -> Self {
Self {
issuer: issuer.into(),
audience: audience.into(),
algorithms: DEFAULT_ALGORITHMS.to_vec(),
leeway_secs: DEFAULT_LEEWAY_SECS,
jwks,
}
}
pub fn algorithms(mut self, algorithms: &[Algorithm]) -> Self {
self.algorithms = algorithms.to_vec();
self
}
pub fn leeway_secs(mut self, leeway_secs: u64) -> Self {
self.leeway_secs = leeway_secs;
self
}
pub fn verify(&self, id_token: &str, nonce: &str) -> Result<VerifiedIdToken, VerifyError> {
let header = decode_header(id_token).map_err(|_| VerifyError::MalformedToken)?;
if !self.algorithms.contains(&header.alg) {
return Err(VerifyError::AlgorithmNotAllowed);
}
let jwk = self.select_key(header.kid.as_deref())?;
let key = DecodingKey::from_jwk(jwk).map_err(|_| VerifyError::InvalidKey)?;
let mut validation = Validation::new(header.alg);
validation.algorithms = vec![header.alg];
validation.leeway = self.leeway_secs;
validation.validate_exp = true;
validation.set_issuer(&[&self.issuer]);
validation.set_audience(&[&self.audience]);
let data = decode::<VerifiedClaims>(id_token, &key, &validation).map_err(map_jwt_error)?;
self.check_authorized_party(&data.claims)?;
let token_nonce = data.claims.nonce.unwrap_or_default();
if !bool::from(token_nonce.as_bytes().ct_eq(nonce.as_bytes())) {
return Err(VerifyError::NonceMismatch);
}
let payload = decode_payload(id_token)?;
Ok(VerifiedIdToken {
id_token: id_token.to_owned(),
payload,
})
}
fn check_authorized_party(&self, claims: &VerifiedClaims) -> Result<(), VerifyError> {
match &claims.azp {
Some(azp) if azp != &self.audience => Err(VerifyError::AuthorizedPartyMismatch),
None if claims.aud.len() > 1 => Err(VerifyError::AuthorizedPartyMissing),
_ => Ok(()),
}
}
fn select_key(&self, kid: Option<&str>) -> Result<&Jwk, VerifyError> {
match kid {
Some(kid) => self.jwks.find(kid).ok_or(VerifyError::UnknownKey),
None => match self.jwks.keys.as_slice() {
[only] => Ok(only),
_ => Err(VerifyError::UnknownKey),
},
}
}
}
pub struct VerifiedIdToken {
id_token: String,
payload: Vec<u8>,
}
impl VerifiedIdToken {
pub fn id_token(&self) -> &str {
&self.id_token
}
pub fn claims(&self) -> Result<OidcClaims<'_>, ClaimsError> {
OidcClaims::from_payload(&self.payload)
}
}
impl fmt::Debug for VerifiedIdToken {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("VerifiedIdToken")
.field("id_token", &"[redacted]")
.finish_non_exhaustive()
}
}
#[derive(serde::Deserialize)]
struct VerifiedClaims {
#[serde(default)]
nonce: Option<String>,
#[serde(default)]
azp: Option<String>,
aud: Audiences,
}
#[derive(serde::Deserialize)]
#[serde(untagged)]
enum Audiences {
One(#[allow(dead_code)] String),
Many(Vec<String>),
}
impl Audiences {
fn len(&self) -> usize {
match self {
Audiences::One(_) => 1,
Audiences::Many(auds) => auds.len(),
}
}
}
fn decode_payload(token: &str) -> Result<Vec<u8>, VerifyError> {
let payload = token.split('.').nth(1).ok_or(VerifyError::MalformedToken)?;
base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(payload)
.map_err(|_| VerifyError::MalformedToken)
}
fn map_jwt_error(err: jsonwebtoken::errors::Error) -> VerifyError {
use jsonwebtoken::errors::ErrorKind;
match err.kind() {
ErrorKind::InvalidSignature => VerifyError::SignatureInvalid,
ErrorKind::InvalidIssuer => VerifyError::IssuerMismatch,
ErrorKind::InvalidAudience => VerifyError::AudienceMismatch,
ErrorKind::ExpiredSignature => VerifyError::Expired,
ErrorKind::ImmatureSignature => VerifyError::ImmatureToken,
ErrorKind::InvalidAlgorithm => VerifyError::AlgorithmNotAllowed,
_ => VerifyError::InvalidToken,
}
}
#[cfg(test)]
mod tests {
use std::time::{SystemTime, UNIX_EPOCH};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
use serde_json::{Value, json};
use super::*;
const KID: &str = "test-key";
const ISSUER: &str = "https://issuer.example";
const AUDIENCE: &str = "test-client-id";
const N: &str = include_str!("testdata/rsa_modulus_b64u.txt");
const PRIV_PEM: &str = include_str!("testdata/rsa_priv_pem.txt");
fn now() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs()
}
fn jwks_with(kid: &str, n: &str) -> JwkSet {
let json = format!(
r#"{{"keys":[{{"kty":"RSA","use":"sig","kid":"{kid}","alg":"RS256","n":"{n}","e":"AQAB"}}]}}"#
);
serde_json::from_str(&json).expect("valid JWKS")
}
fn jwks() -> JwkSet {
jwks_with(KID, N)
}
fn sign(claims: &Value, kid: &str) -> String {
let mut header = Header::new(Algorithm::RS256);
header.kid = Some(kid.to_owned());
let key = EncodingKey::from_rsa_pem(PRIV_PEM.as_bytes()).expect("valid PEM");
encode(&header, claims, &key).expect("encode")
}
fn claims(nonce: &str, aud: &str, exp: u64) -> Value {
json!({
"iss": ISSUER,
"aud": aud,
"sub": "user-123",
"exp": exp,
"iat": now(),
"nonce": nonce,
"email": "user@example.com",
"email_verified": true,
"hd": "example.com",
})
}
fn valid_token(nonce: &str) -> String {
sign(&claims(nonce, AUDIENCE, now() + 3600), KID)
}
#[test]
fn verifies_a_valid_token() {
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let verified = verifier
.verify(&valid_token("the-nonce"), "the-nonce")
.unwrap();
let claims = verified.claims().unwrap();
assert_eq!(claims.issuer, ISSUER);
assert_eq!(claims.subject, "user-123");
assert_eq!(claims.audiences, ["test-client-id"]);
assert_eq!(claims.email, Some("user@example.com"));
assert_eq!(claims.email_verified, Some(true));
assert_eq!(
claims.extra.get("hd").and_then(Value::as_str),
Some("example.com")
);
}
#[test]
fn rejects_wrong_nonce() {
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let err = verifier
.verify(&valid_token("the-nonce"), "other-nonce")
.unwrap_err();
assert!(matches!(err, VerifyError::NonceMismatch), "{err:?}");
}
#[test]
fn rejects_wrong_audience() {
let token = sign(&claims("n", "some-other-client", now() + 3600), KID);
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let err = verifier.verify(&token, "n").unwrap_err();
assert!(matches!(err, VerifyError::AudienceMismatch), "{err:?}");
}
#[test]
fn accepts_matching_azp() {
let token = sign(
&json!({
"iss": ISSUER, "aud": AUDIENCE, "sub": "user-123",
"exp": now() + 3600, "iat": now(), "nonce": "n", "azp": AUDIENCE,
}),
KID,
);
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
assert!(verifier.verify(&token, "n").is_ok());
}
#[test]
fn rejects_mismatched_azp() {
let token = sign(
&json!({
"iss": ISSUER, "aud": AUDIENCE, "sub": "user-123",
"exp": now() + 3600, "iat": now(), "nonce": "n",
"azp": "another-client",
}),
KID,
);
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let err = verifier.verify(&token, "n").unwrap_err();
assert!(
matches!(err, VerifyError::AuthorizedPartyMismatch),
"{err:?}"
);
}
#[test]
fn accepts_multiple_audiences_with_matching_azp() {
let token = sign(
&json!({
"iss": ISSUER, "aud": [AUDIENCE, "other-client"], "sub": "user-123",
"exp": now() + 3600, "iat": now(), "nonce": "n", "azp": AUDIENCE,
}),
KID,
);
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
assert!(verifier.verify(&token, "n").is_ok());
}
#[test]
fn rejects_multiple_audiences_without_azp() {
let token = sign(
&json!({
"iss": ISSUER, "aud": [AUDIENCE, "other-client"], "sub": "user-123",
"exp": now() + 3600, "iat": now(), "nonce": "n",
}),
KID,
);
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let err = verifier.verify(&token, "n").unwrap_err();
assert!(
matches!(err, VerifyError::AuthorizedPartyMissing),
"{err:?}"
);
}
#[test]
fn rejects_wrong_issuer() {
let verifier = IdTokenVerifier::new("https://attacker.example", AUDIENCE, jwks());
let err = verifier.verify(&valid_token("n"), "n").unwrap_err();
assert!(matches!(err, VerifyError::IssuerMismatch), "{err:?}");
}
#[test]
fn rejects_expired_token() {
let token = sign(&claims("n", AUDIENCE, now() - 120), KID);
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let err = verifier.verify(&token, "n").unwrap_err();
assert!(matches!(err, VerifyError::Expired), "{err:?}");
}
#[test]
fn rejects_unknown_kid() {
let token = sign(&claims("n", AUDIENCE, now() + 3600), "rotated-away");
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let err = verifier.verify(&token, "n").unwrap_err();
assert!(matches!(err, VerifyError::UnknownKey), "{err:?}");
}
#[test]
fn rejects_disallowed_algorithm() {
let verifier =
IdTokenVerifier::new(ISSUER, AUDIENCE, jwks()).algorithms(&[Algorithm::ES256]);
let err = verifier.verify(&valid_token("n"), "n").unwrap_err();
assert!(matches!(err, VerifyError::AlgorithmNotAllowed), "{err:?}");
}
#[test]
fn rejects_bad_signature() {
let wrong_n = format!("y{}", &N[1..]);
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks_with(KID, &wrong_n));
let err = verifier.verify(&valid_token("n"), "n").unwrap_err();
assert!(matches!(err, VerifyError::SignatureInvalid), "{err:?}");
}
#[test]
fn rejects_malformed_token() {
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let err = verifier.verify("not-a-jwt", "n").unwrap_err();
assert!(matches!(err, VerifyError::MalformedToken), "{err:?}");
}
#[test]
fn debug_redacts_raw_token() {
let verifier = IdTokenVerifier::new(ISSUER, AUDIENCE, jwks());
let token = valid_token("n");
let verified = verifier.verify(&token, "n").unwrap();
let debug = format!("{verified:?}");
assert!(!debug.contains(&token), "{debug}");
assert!(debug.contains("[redacted]"), "{debug}");
}
}