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 {
ES256 {
#[serde(rename = "pub")]
public: String,
},
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
}
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)),
}
}
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 })
}
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 })
}
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 })
}
}
}
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,
),
})
}
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 {
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,
),
})
}
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!");
}
}