use std::fmt;
use jsonwebtoken::jwk::{AlgorithmParameters, EllipticCurve, KeyAlgorithm};
pub const DEFAULT_ALGORITHMS: &[&str] = &[
"RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "EdDSA",
];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
#[allow(clippy::upper_case_acronyms)]
pub enum Algorithm {
RS256,
RS384,
RS512,
PS256,
PS384,
PS512,
ES256,
ES384,
EdDSA,
}
impl Algorithm {
pub fn as_str(self) -> &'static str {
match self {
Algorithm::RS256 => "RS256",
Algorithm::RS384 => "RS384",
Algorithm::RS512 => "RS512",
Algorithm::PS256 => "PS256",
Algorithm::PS384 => "PS384",
Algorithm::PS512 => "PS512",
Algorithm::ES256 => "ES256",
Algorithm::ES384 => "ES384",
Algorithm::EdDSA => "EdDSA",
}
}
pub(crate) fn to_jwt(self) -> jsonwebtoken::Algorithm {
match self {
Algorithm::RS256 => jsonwebtoken::Algorithm::RS256,
Algorithm::RS384 => jsonwebtoken::Algorithm::RS384,
Algorithm::RS512 => jsonwebtoken::Algorithm::RS512,
Algorithm::PS256 => jsonwebtoken::Algorithm::PS256,
Algorithm::PS384 => jsonwebtoken::Algorithm::PS384,
Algorithm::PS512 => jsonwebtoken::Algorithm::PS512,
Algorithm::ES256 => jsonwebtoken::Algorithm::ES256,
Algorithm::ES384 => jsonwebtoken::Algorithm::ES384,
Algorithm::EdDSA => jsonwebtoken::Algorithm::EdDSA,
}
}
pub(crate) fn from_jwt(alg: jsonwebtoken::Algorithm) -> Option<Self> {
Some(match alg {
jsonwebtoken::Algorithm::RS256 => Algorithm::RS256,
jsonwebtoken::Algorithm::RS384 => Algorithm::RS384,
jsonwebtoken::Algorithm::RS512 => Algorithm::RS512,
jsonwebtoken::Algorithm::PS256 => Algorithm::PS256,
jsonwebtoken::Algorithm::PS384 => Algorithm::PS384,
jsonwebtoken::Algorithm::PS512 => Algorithm::PS512,
jsonwebtoken::Algorithm::ES256 => Algorithm::ES256,
jsonwebtoken::Algorithm::ES384 => Algorithm::ES384,
jsonwebtoken::Algorithm::EdDSA => Algorithm::EdDSA,
_ => return None,
})
}
}
impl fmt::Display for Algorithm {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum AlgorithmError {
#[error("\"{name}\" — an unsigned token is never acceptable")]
#[non_exhaustive]
Unsigned {
name: String,
},
#[error(
"\"{name}\" — HMAC algorithms verify with a shared secret, which a resource \
server must never hold, and accepting one alongside a public key set is the \
classic key-confusion attack (a token signed with the PUBLIC key as the HMAC \
secret)"
)]
#[non_exhaustive]
Hmac {
name: String,
},
#[error("\"{name}\" — not a JWS algorithm this server can verify (supported: {supported})")]
#[non_exhaustive]
Unsupported {
name: String,
supported: String,
},
}
pub fn parse_algorithm(name: &str) -> Result<Algorithm, AlgorithmError> {
let name = name.trim();
if name.eq_ignore_ascii_case("none") {
return Err(AlgorithmError::Unsigned {
name: name.to_string(),
});
}
match name.parse::<jsonwebtoken::Algorithm>() {
Ok(
jsonwebtoken::Algorithm::HS256
| jsonwebtoken::Algorithm::HS384
| jsonwebtoken::Algorithm::HS512,
) => Err(AlgorithmError::Hmac {
name: name.to_string(),
}),
Ok(alg) => Algorithm::from_jwt(alg).ok_or_else(|| unsupported(name)),
Err(_) => Err(unsupported(name)),
}
}
fn unsupported(name: &str) -> AlgorithmError {
AlgorithmError::Unsupported {
name: name.to_string(),
supported: DEFAULT_ALGORITHMS.join(", "),
}
}
pub(crate) fn key_algorithms(params: &AlgorithmParameters) -> Option<Vec<Algorithm>> {
use Algorithm::*;
match params {
AlgorithmParameters::RSA(_) => Some(vec![RS256, RS384, RS512, PS256, PS384, PS512]),
AlgorithmParameters::EllipticCurve(p) => match p.curve {
EllipticCurve::P256 => Some(vec![ES256]),
EllipticCurve::P384 => Some(vec![ES384]),
_ => None,
},
AlgorithmParameters::OctetKeyPair(p) => match p.curve {
EllipticCurve::Ed25519 => Some(vec![EdDSA]),
_ => None,
},
AlgorithmParameters::OctetKey(_) => None,
}
}
pub(crate) fn signing_algorithm(alg: &KeyAlgorithm) -> Option<Algorithm> {
Some(match alg {
KeyAlgorithm::RS256 => Algorithm::RS256,
KeyAlgorithm::RS384 => Algorithm::RS384,
KeyAlgorithm::RS512 => Algorithm::RS512,
KeyAlgorithm::PS256 => Algorithm::PS256,
KeyAlgorithm::PS384 => Algorithm::PS384,
KeyAlgorithm::PS512 => Algorithm::PS512,
KeyAlgorithm::ES256 => Algorithm::ES256,
KeyAlgorithm::ES384 => Algorithm::ES384,
KeyAlgorithm::EdDSA => Algorithm::EdDSA,
_ => return None,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hmac_and_none_can_never_be_configured() {
for bad in [
"HS256", "HS384", "HS512", "none", "None", "ES512", "rs256", "",
] {
assert!(parse_algorithm(bad).is_err(), "{bad:?} must be refused");
}
assert!(
parse_algorithm("HS256")
.unwrap_err()
.to_string()
.contains("key-confusion")
);
for good in DEFAULT_ALGORITHMS {
let alg = parse_algorithm(good).expect("every default parses");
assert_eq!(alg.as_str(), *good);
assert_eq!(alg.to_string(), *good);
assert_eq!(format!("{alg:?}"), *good);
assert_eq!(Algorithm::from_jwt(alg.to_jwt()), Some(alg));
}
}
#[test]
fn refusals_are_typed_and_keep_their_text() {
assert_eq!(
parse_algorithm(" none "),
Err(AlgorithmError::Unsigned {
name: "none".into()
})
);
assert_eq!(
parse_algorithm("none").unwrap_err().to_string(),
"\"none\" — an unsigned token is never acceptable"
);
assert!(matches!(
parse_algorithm("HS384"),
Err(AlgorithmError::Hmac { .. })
));
let err = parse_algorithm("ES512").unwrap_err();
assert!(matches!(err, AlgorithmError::Unsupported { .. }));
assert_eq!(
err.to_string(),
"\"ES512\" — not a JWS algorithm this server can verify (supported: RS256, RS384, \
RS512, PS256, PS384, PS512, ES256, ES384, EdDSA)"
);
let _: Box<dyn std::error::Error + Send + Sync> = Box::new(err);
}
#[test]
fn hmac_has_no_crate_algorithm() {
for hmac in [
jsonwebtoken::Algorithm::HS256,
jsonwebtoken::Algorithm::HS384,
jsonwebtoken::Algorithm::HS512,
] {
assert_eq!(Algorithm::from_jwt(hmac), None);
}
}
}