basic-jwt 0.5.0

Basic JWT signing and verification library
Documentation
use const_oid::ObjectIdentifier;
use jsonwebtoken::{Algorithm, AlgorithmFamily, DecodingKey, EncodingKey, Validation};
use p384::ecdsa::signature::rand_core::OsRng;
use p384::elliptic_curve::zeroize::ZeroizeOnDrop;
use p384::pkcs8::EncodePublicKey;
use p384::pkcs8::{EncodePrivateKey, LineEnding};
use pkcs8::DecodePrivateKey;
use serde::Serialize;
use serde::de::DeserializeOwned;
use std::str::FromStr;
use zeroize::{Zeroize, Zeroizing};

#[derive(Debug, thiserror::Error)]
pub enum BasicJwtError {
    #[error("could not guess private key algorithm family!")]
    CouldNotGuessAlgorithmFamily,
    #[error("unsupported algorithm family: {0:?}!")]
    UnsupportedAlgorithmFamily(AlgorithmFamily),
    #[error("failed to parse elliptic curve private key: {0:?}")]
    ParseEcPrivateKey(#[source] pkcs8::Error),
    #[error("invalid elliptic curve algorithm: {0:?}")]
    InvalidEcAlgorithm(ObjectIdentifier),
    #[error("missing elliptic curve parameters in private key!")]
    MissingEcParameters,
    #[error("failed to decode ec params as OID! {0}")]
    DecodeEcParamsAsOID(#[source] pkcs8::der::Error),
    #[error("unsupported ec params: {0:?}")]
    UnsupportedEcParams(ObjectIdentifier),
    #[error("failed to encode key to pkcs#8: {0}")]
    EncodeKeyToPKCS8(#[source] p384::pkcs8::Error),
    #[error("failed to envelope pkcs#8-encoded key in pem: {0}")]
    EnvelopeEncodedKeyInPem(#[source] p384::pkcs8::der::Error),
    #[error("failed to parse signing key: {0}")]
    ParseSigningKey(#[source] jsonwebtoken::signature::Error),
    #[error("failed to parse encoding key: {0}")]
    ParseEncodingKey(#[source] jsonwebtoken::errors::Error),
    #[error("failed to encode jwt: {0}")]
    EncodeJWT(#[source] jsonwebtoken::errors::Error),
    #[error("failed to encode public key to pem: {0}")]
    EncodePublicKeyToPem(#[source] p384::pkcs8::spki::Error),
    #[error("failed to parse decoding key: {0}")]
    ParseDecodingKey(#[source] jsonwebtoken::errors::Error),
    #[error("failed to decode jwt: {0}")]
    DecodeJWT(#[source] jsonwebtoken::errors::Error),
}

pub type Res<T> = Result<T, BasicJwtError>;

#[derive(serde::Serialize, serde::Deserialize, Clone, Debug, Eq, PartialEq)]
#[serde(tag = "alg")]
pub enum JWTPublicKey {
    /// ECDSA with SHA2-256 variant
    ES256 {
        #[serde(rename = "pub")]
        public: String,
    },
    /// ECDSA with SHA2-384 variant
    ES384 {
        #[serde(rename = "pub")]
        public: String,
    },
}

#[derive(serde::Serialize, serde::Deserialize, Clone, Debug, Zeroize, ZeroizeOnDrop)]
#[serde(tag = "alg")]
pub enum JWTPrivateKey {
    ES256 { r#priv: String },
    ES384 { r#priv: String },
}

impl JWTPrivateKey {
    fn guess_key_family_from_pem(key: &str) -> Option<AlgorithmFamily> {
        match EncodingKey::from_ec_pem(key.as_bytes()) {
            Ok(_) => return Some(AlgorithmFamily::Ec),
            Err(e) => {
                tracing::trace!("Not a EC key: {e}");
            }
        }

        match EncodingKey::from_rsa_pem(key.as_bytes()) {
            Ok(_) => return Some(AlgorithmFamily::Rsa),
            Err(e) => {
                tracing::trace!("Not a RSA key: {e}");
            }
        }

        match EncodingKey::from_ed_pem(key.as_bytes()) {
            Ok(_) => return Some(AlgorithmFamily::Ed),
            Err(e) => {
                tracing::trace!("Not a Ecdsa key: {e}");
            }
        }

        None
    }

    /// Parse private key from given PEM
    pub fn parse_key_pem(key: &str) -> Res<Self> {
        match Self::guess_key_family_from_pem(key) {
            None => Err(BasicJwtError::CouldNotGuessAlgorithmFamily),
            Some(AlgorithmFamily::Ec) => {
                let pkey = pkcs8::PrivateKeyInfoOwned::from_pkcs8_pem(key)
                    .map_err(BasicJwtError::ParseEcPrivateKey)?;

                if pkey.algorithm.oid != const_oid::db::rfc5753::ID_EC_PUBLIC_KEY {
                    return Err(BasicJwtError::InvalidEcAlgorithm(pkey.algorithm.oid));
                }

                let Some(params) = pkey.algorithm.parameters.as_ref() else {
                    return Err(BasicJwtError::MissingEcParameters);
                };

                match params
                    .decode_as::<ObjectIdentifier>()
                    .map_err(BasicJwtError::DecodeEcParamsAsOID)?
                {
                    const_oid::db::rfc5912::SECP_256_R_1 => Ok(Self::ES256 {
                        r#priv: key.to_string(),
                    }),
                    const_oid::db::rfc5912::SECP_384_R_1 => Ok(Self::ES384 {
                        r#priv: key.to_string(),
                    }),
                    oid => Err(BasicJwtError::UnsupportedEcParams(oid)),
                }
            }
            Some(f) => Err(BasicJwtError::UnsupportedAlgorithmFamily(f)),
        }
    }

    /// Generate a new elliptic curve 256 signing key
    pub fn generate_ec256_signing_key() -> Res<Self> {
        let signing_key = p256::ecdsa::SigningKey::random(&mut OsRng);
        let priv_pem = signing_key
            .to_pkcs8_der()
            .map_err(BasicJwtError::EncodeKeyToPKCS8)?
            .to_pem("PRIVATE KEY", LineEnding::LF)
            .map_err(BasicJwtError::EnvelopeEncodedKeyInPem)?
            .to_string();

        Ok(Self::ES256 { r#priv: priv_pem })
    }

    /// Generate a new ES384 signing key
    pub fn generate_ec384_signing_key() -> Res<Self> {
        let signing_key = p384::ecdsa::SigningKey::random(&mut OsRng);
        let priv_pem = signing_key
            .to_pkcs8_der()
            .map_err(BasicJwtError::EncodeKeyToPKCS8)?
            .to_pem("PRIVATE KEY", LineEnding::LF)
            .map_err(BasicJwtError::EnvelopeEncodedKeyInPem)?
            .to_string();

        Ok(Self::ES384 { r#priv: priv_pem })
    }

    /// Get associated public key
    pub fn to_public_key(&self) -> Res<JWTPublicKey> {
        match self {
            JWTPrivateKey::ES256 { r#priv } => {
                let signing_key = p256::ecdsa::SigningKey::from_str(r#priv)
                    .map_err(BasicJwtError::ParseSigningKey)?;

                let pub_key = p256::ecdsa::VerifyingKey::from(signing_key);
                let pub_pem = pub_key
                    .to_public_key_pem(LineEnding::LF)
                    .map_err(BasicJwtError::EncodePublicKeyToPem)?;

                Ok(JWTPublicKey::ES256 { public: pub_pem })
            }
            JWTPrivateKey::ES384 { r#priv } => {
                let signing_key = p384::ecdsa::SigningKey::from_str(r#priv)
                    .map_err(BasicJwtError::ParseSigningKey)?;

                let pub_key = p384::ecdsa::VerifyingKey::from(signing_key);
                let pub_pem = pub_key
                    .to_public_key_pem(LineEnding::LF)
                    .map_err(BasicJwtError::EncodePublicKeyToPem)?;

                Ok(JWTPublicKey::ES384 { public: pub_pem })
            }
        }
    }

    /// Get the decoding key & algorithm associated with a private key
    pub fn get_encoding_key(&self) -> Res<(Zeroizing<EncodingKey>, Algorithm)> {
        Ok(match self {
            JWTPrivateKey::ES256 { r#priv } => (
                Zeroizing::new(
                    EncodingKey::from_ec_pem(r#priv.as_bytes())
                        .map_err(BasicJwtError::ParseEncodingKey)?,
                ),
                Algorithm::ES256,
            ),
            JWTPrivateKey::ES384 { r#priv } => (
                Zeroizing::new(
                    EncodingKey::from_ec_pem(r#priv.as_bytes())
                        .map_err(BasicJwtError::ParseEncodingKey)?,
                ),
                Algorithm::ES384,
            ),
        })
    }

    /// Sign a JWT
    pub fn sign_jwt<C: Serialize>(&self, claims: &C) -> Res<String> {
        let (encoding_key, algorithm) = self.get_encoding_key()?;

        jsonwebtoken::encode(
            &jsonwebtoken::Header::new(algorithm),
            &claims,
            &encoding_key,
        )
        .map_err(BasicJwtError::EncodeJWT)
    }
}

impl JWTPublicKey {
    /// Get the decoding key & algorithm associated with a public key
    pub fn get_decoding_key(&self) -> Res<(DecodingKey, Algorithm)> {
        Ok(match self {
            JWTPublicKey::ES256 { public } => (
                DecodingKey::from_ec_pem(public.as_bytes())
                    .map_err(BasicJwtError::ParseDecodingKey)?,
                Algorithm::ES256,
            ),
            JWTPublicKey::ES384 { public } => (
                DecodingKey::from_ec_pem(public.as_bytes())
                    .map_err(BasicJwtError::ParseDecodingKey)?,
                Algorithm::ES384,
            ),
        })
    }

    /// Validate a given JWT
    pub fn validate_jwt<E: DeserializeOwned + Clone>(&self, jwt: &str) -> Res<E> {
        let (decoding_key, algorithm) = self.get_decoding_key()?;

        let validation = Validation::new(algorithm);
        Ok(jsonwebtoken::decode::<E>(jwt, &decoding_key, &validation)
            .map_err(BasicJwtError::DecodeJWT)?
            .claims)
    }
}

#[cfg(test)]
mod test {
    use std::time::{SystemTime, UNIX_EPOCH};

    use crate::JWTPrivateKey;
    use serde::{Deserialize, Serialize};

    fn time() -> u64 {
        SystemTime::now()
            .duration_since(UNIX_EPOCH)
            .unwrap()
            .as_secs()
    }

    #[derive(Debug, Serialize, Deserialize, Eq, PartialEq, Clone)]
    pub struct Claims {
        sub: String,
        exp: u64,
    }

    impl Default for Claims {
        fn default() -> Self {
            Self {
                sub: "my-sub".to_string(),
                exp: time() + 100,
            }
        }
    }

    #[test]
    fn jwt_encode_sign_verify_valid_p256() {
        let priv_key = JWTPrivateKey::generate_ec256_signing_key().unwrap();
        let pub_key = priv_key.to_public_key().unwrap();

        let claims = Claims::default();
        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
        let claims_out = pub_key
            .validate_jwt::<Claims>(&jwt)
            .expect("Failed to validate JWT!");

        assert_eq!(claims, claims_out)
    }

    #[test]
    fn jwt_encode_sign_verify_valid_p384() {
        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
        let pub_key = priv_key.to_public_key().unwrap();

        let claims = Claims::default();
        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
        let claims_out = pub_key
            .validate_jwt::<Claims>(&jwt)
            .expect("Failed to validate JWT!");

        assert_eq!(claims, claims_out)
    }

    #[test]
    fn parse_keys() {
        const PEM_ONE: &str = r"-----BEGIN PRIVATE KEY-----
MIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQgktVspUQwcArgXoLx
hhGfw2yY6BB3/Cx8K2fIAck2KD2hRANCAAQ6PWADCFG5Ih6JMruZnuWKGmXEQtAf
0ii/D5ubYKRW+iqx63h/OcW7DAABuh9Go0mRUf85A3PdoGAbkTAZAQoT
-----END PRIVATE KEY-----
";
        let key_one = JWTPrivateKey::parse_key_pem(PEM_ONE).unwrap();

        assert!(matches!(key_one, JWTPrivateKey::ES256 { .. }));

        const PEM_TWO: &str = r"-----BEGIN PRIVATE KEY-----
MIG2AgEAMBAGByqGSM49AgEGBSuBBAAiBIGeMIGbAgEBBDDMAHDNXDQfIZD0n52g
ag4AxAJPK25TaPNK9TIbjx67Zs9U0JmvCbbbqFs6EiS08EyhZANiAAQZh4e/2BDk
pECHm6hsokTKIn9EAgOQ0RtrWh02CZkTBJKvHC58KdwNB1eWSUHUKPKsrE2+h3cW
Apd2mjEkBPSlCoTunNVAq+niutY+9LgcGZ3iFDTiI3GPDepQDtX8b6A=
-----END PRIVATE KEY-----
";
        let key_one = JWTPrivateKey::parse_key_pem(PEM_TWO).unwrap();

        assert!(matches!(key_one, JWTPrivateKey::ES384 { .. }));
    }

    #[test]
    fn jwt_encode_sign_verify_invalid_key() {
        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
        let pub_key_2 = JWTPrivateKey::generate_ec384_signing_key()
            .unwrap()
            .to_public_key()
            .unwrap();

        let claims = Claims::default();
        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
        pub_key_2
            .validate_jwt::<Claims>(&jwt)
            .expect_err("JWT should not have validated!");
    }

    #[test]
    fn jwt_verify_random_string() {
        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
        let pub_key = priv_key.to_public_key().unwrap();

        pub_key
            .validate_jwt::<Claims>("random_string")
            .expect_err("JWT should not have validated!");
    }

    #[test]
    fn jwt_expired() {
        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
        let pub_key = priv_key.to_public_key().unwrap();

        let claims = Claims {
            exp: time() - 100,
            ..Default::default()
        };
        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
        pub_key
            .validate_jwt::<Claims>(&jwt)
            .expect_err("JWT should not have validated!");
    }

    #[test]
    fn jwt_invalid_signature() {
        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
        let pub_key = priv_key.to_public_key().unwrap();

        let claims = Claims::default();
        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
        pub_key
            .validate_jwt::<Claims>(&format!("{jwt}bad"))
            .expect_err("JWT should not have validated!");
    }
}