use core::fmt;
use super::decode::{SegmentDecodeError, decode_segment};
#[doc(alias = "algorithm")]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum JwtAlgorithm {
HS256,
HS512,
#[cfg(feature = "asym-jwt")]
EdDSA,
#[cfg(feature = "asym-jwt")]
ES256,
#[cfg(feature = "asym-jwt")]
RS256,
#[cfg(feature = "asym-jwt")]
RS512,
None,
}
impl std::str::FromStr for JwtAlgorithm {
type Err = JwtHeaderError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"HS256" => Ok(Self::HS256),
"HS512" => Ok(Self::HS512),
#[cfg(feature = "asym-jwt")]
"EdDSA" => Ok(Self::EdDSA),
#[cfg(feature = "asym-jwt")]
"ES256" => Ok(Self::ES256),
#[cfg(feature = "asym-jwt")]
"RS256" => Ok(Self::RS256),
#[cfg(feature = "asym-jwt")]
"RS512" => Ok(Self::RS512),
"none" => Ok(Self::None),
_ => Err(JwtHeaderError {
kind: JwtHeaderErrorKind::UnsupportedAlgorithm,
}),
}
}
}
impl fmt::Display for JwtAlgorithm {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::HS256 => write!(f, "HS256"),
Self::HS512 => write!(f, "HS512"),
#[cfg(feature = "asym-jwt")]
Self::EdDSA => write!(f, "EdDSA"),
#[cfg(feature = "asym-jwt")]
Self::ES256 => write!(f, "ES256"),
#[cfg(feature = "asym-jwt")]
Self::RS256 => write!(f, "RS256"),
#[cfg(feature = "asym-jwt")]
Self::RS512 => write!(f, "RS512"),
Self::None => write!(f, "none"),
}
}
}
#[doc(alias = "jose_header")]
#[derive(Debug, Clone)]
pub struct JwtHeader {
alg: JwtAlgorithm,
typ: Option<String>,
kid: Option<String>,
}
impl JwtHeader {
#[must_use]
#[inline]
pub fn alg(&self) -> JwtAlgorithm {
self.alg
}
#[must_use]
#[inline]
pub fn typ(&self) -> Option<&str> {
self.typ.as_deref()
}
#[must_use]
#[inline]
pub fn kid(&self) -> Option<&str> {
self.kid.as_deref()
}
}
impl JwtHeader {
#[must_use = "parsing may fail; handle the Result"]
pub fn parse(header_b64: &str) -> Result<Self, JwtHeaderError> {
let value = decode_segment(header_b64).map_err(|e| JwtHeaderError {
kind: match e {
SegmentDecodeError::InvalidBase64 => JwtHeaderErrorKind::InvalidBase64,
SegmentDecodeError::InvalidUtf8 => JwtHeaderErrorKind::InvalidUtf8,
SegmentDecodeError::InvalidJson => JwtHeaderErrorKind::InvalidJson,
},
})?;
if value.get("crit").is_some() {
return Err(JwtHeaderError {
kind: JwtHeaderErrorKind::CriticalHeaderUnsupported,
});
}
let alg_str = value.get_str("alg").ok_or(JwtHeaderError {
kind: JwtHeaderErrorKind::MissingAlgorithm,
})?;
let alg = alg_str.parse::<JwtAlgorithm>()?;
let typ = value.get_str("typ").map(String::from);
let kid = value.get_str("kid").map(String::from);
Ok(Self { alg, typ, kid })
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum JwtHeaderErrorKind {
InvalidBase64,
InvalidUtf8,
InvalidJson,
MissingAlgorithm,
UnsupportedAlgorithm,
CriticalHeaderUnsupported,
}
#[doc(alias = "header_error")]
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct JwtHeaderError {
kind: JwtHeaderErrorKind,
}
impl fmt::Display for JwtHeaderError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.kind {
JwtHeaderErrorKind::InvalidBase64 => {
write!(f, "jwt header: invalid base64url encoding")
}
JwtHeaderErrorKind::InvalidUtf8 => {
write!(f, "jwt header: decoded bytes are not valid UTF-8")
}
JwtHeaderErrorKind::InvalidJson => {
write!(f, "jwt header: decoded content is not valid JSON")
}
JwtHeaderErrorKind::MissingAlgorithm => {
write!(f, "jwt header: missing 'alg' field")
}
JwtHeaderErrorKind::UnsupportedAlgorithm => {
write!(f, "jwt header: unsupported algorithm")
}
JwtHeaderErrorKind::CriticalHeaderUnsupported => {
write!(f, "jwt header: unsupported 'crit' header parameter")
}
}
}
}
impl std::error::Error for JwtHeaderError {}
#[cfg(test)]
mod tests {
use super::*;
use crate::encoding::base64url_encode;
#[test]
fn parse_hs256_header() {
let json = r#"{"alg":"HS256","typ":"JWT"}"#;
let b64 = base64url_encode(json.as_bytes());
let header = JwtHeader::parse(&b64).unwrap();
assert_eq!(header.alg(), JwtAlgorithm::HS256);
assert_eq!(header.typ(), Some("JWT"));
assert_eq!(header.kid(), None);
}
#[test]
fn parse_hs512_header() {
let json = r#"{"alg":"HS512"}"#;
let b64 = base64url_encode(json.as_bytes());
let header = JwtHeader::parse(&b64).unwrap();
assert_eq!(header.alg(), JwtAlgorithm::HS512);
assert_eq!(header.typ(), None);
}
#[test]
fn parse_none_algorithm() {
let json = r#"{"alg":"none"}"#;
let b64 = base64url_encode(json.as_bytes());
let header = JwtHeader::parse(&b64).unwrap();
assert_eq!(header.alg(), JwtAlgorithm::None);
}
#[test]
fn parse_header_with_kid() {
let json = r#"{"alg":"HS256","typ":"JWT","kid":"key-1"}"#;
let b64 = base64url_encode(json.as_bytes());
let header = JwtHeader::parse(&b64).unwrap();
assert_eq!(header.kid(), Some("key-1"));
}
#[test]
fn reject_invalid_base64() {
let err = JwtHeader::parse("!!!invalid!!!").unwrap_err();
assert!(err.to_string().contains("base64url"));
}
#[test]
fn reject_invalid_json() {
let b64 = base64url_encode(b"not json");
let err = JwtHeader::parse(&b64).unwrap_err();
assert!(err.to_string().contains("JSON"));
}
#[test]
fn reject_missing_alg() {
let json = r#"{"typ":"JWT"}"#;
let b64 = base64url_encode(json.as_bytes());
let err = JwtHeader::parse(&b64).unwrap_err();
assert!(err.to_string().contains("alg"));
}
#[test]
fn reject_unsupported_algorithm() {
let json = r#"{"alg":"PS256"}"#;
let b64 = base64url_encode(json.as_bytes());
let err = JwtHeader::parse(&b64).unwrap_err();
assert!(err.to_string().contains("unsupported"));
}
#[test]
fn reject_crit_header() {
let json = r#"{"alg":"HS256","crit":["b64"],"b64":false}"#;
let b64 = base64url_encode(json.as_bytes());
let err = JwtHeader::parse(&b64).unwrap_err();
assert!(err.to_string().contains("crit"));
}
#[test]
fn reject_empty_crit_header() {
let json = r#"{"alg":"HS256","crit":[]}"#;
let b64 = base64url_encode(json.as_bytes());
let err = JwtHeader::parse(&b64).unwrap_err();
assert!(err.to_string().contains("crit"));
}
#[cfg(feature = "asym-jwt")]
#[test]
fn parse_rs256_header() {
let json = r#"{"alg":"RS256","typ":"JWT","kid":"g-1"}"#;
let b64 = base64url_encode(json.as_bytes());
let header = JwtHeader::parse(&b64).unwrap();
assert_eq!(header.alg(), JwtAlgorithm::RS256);
assert_eq!(header.kid(), Some("g-1"));
assert_eq!(JwtAlgorithm::RS256.to_string(), "RS256");
assert_eq!(JwtAlgorithm::RS512.to_string(), "RS512");
}
#[test]
fn algorithm_display() {
assert_eq!(JwtAlgorithm::HS256.to_string(), "HS256");
assert_eq!(JwtAlgorithm::HS512.to_string(), "HS512");
assert_eq!(JwtAlgorithm::None.to_string(), "none");
}
#[test]
fn error_implements_std_error() {
let err: Box<dyn std::error::Error> = Box::new(JwtHeaderError {
kind: JwtHeaderErrorKind::MissingAlgorithm,
});
let _ = err.to_string();
}
}