Skip to main content

basic_jwt/
lib.rs

1use const_oid::ObjectIdentifier;
2use jsonwebtoken::{Algorithm, AlgorithmFamily, DecodingKey, EncodingKey, Validation};
3use p384::ecdsa::signature::rand_core::OsRng;
4use p384::elliptic_curve::zeroize::ZeroizeOnDrop;
5use p384::pkcs8::EncodePublicKey;
6use p384::pkcs8::{EncodePrivateKey, LineEnding};
7use pkcs8::DecodePrivateKey;
8use serde::Serialize;
9use serde::de::DeserializeOwned;
10use std::str::FromStr;
11use zeroize::{Zeroize, Zeroizing};
12
13#[derive(Debug, thiserror::Error)]
14pub enum BasicJwtError {
15    #[error("could not guess private key algorithm family!")]
16    CouldNotGuessAlgorithmFamily,
17    #[error("unsupported algorithm family: {0:?}!")]
18    UnsupportedAlgorithmFamily(AlgorithmFamily),
19    #[error("failed to parse elliptic curve private key: {0:?}")]
20    ParseEcPrivateKey(#[source] pkcs8::Error),
21    #[error("invalid elliptic curve algorithm: {0:?}")]
22    InvalidEcAlgorithm(ObjectIdentifier),
23    #[error("missing elliptic curve parameters in private key!")]
24    MissingEcParameters,
25    #[error("failed to decode ec params as OID! {0}")]
26    DecodeEcParamsAsOID(#[source] pkcs8::der::Error),
27    #[error("unsupported ec params: {0:?}")]
28    UnsupportedEcParams(ObjectIdentifier),
29    #[error("failed to encode key to pkcs#8: {0}")]
30    EncodeKeyToPKCS8(#[source] p384::pkcs8::Error),
31    #[error("failed to envelope pkcs#8-encoded key in pem: {0}")]
32    EnvelopeEncodedKeyInPem(#[source] p384::pkcs8::der::Error),
33    #[error("failed to parse signing key: {0}")]
34    ParseSigningKey(#[source] jsonwebtoken::signature::Error),
35    #[error("failed to parse encoding key: {0}")]
36    ParseEncodingKey(#[source] jsonwebtoken::errors::Error),
37    #[error("failed to encode jwt: {0}")]
38    EncodeJWT(#[source] jsonwebtoken::errors::Error),
39    #[error("failed to encode public key to pem: {0}")]
40    EncodePublicKeyToPem(#[source] p384::pkcs8::spki::Error),
41    #[error("failed to parse decoding key: {0}")]
42    ParseDecodingKey(#[source] jsonwebtoken::errors::Error),
43    #[error("failed to decode jwt: {0}")]
44    DecodeJWT(#[source] jsonwebtoken::errors::Error),
45}
46
47pub type Res<T> = Result<T, BasicJwtError>;
48
49#[derive(serde::Serialize, serde::Deserialize, Clone, Debug, Eq, PartialEq)]
50#[serde(tag = "alg")]
51pub enum JWTPublicKey {
52    /// ECDSA with SHA2-256 variant
53    ES256 {
54        #[serde(rename = "pub")]
55        public: String,
56    },
57    /// ECDSA with SHA2-384 variant
58    ES384 {
59        #[serde(rename = "pub")]
60        public: String,
61    },
62}
63
64#[derive(serde::Serialize, serde::Deserialize, Clone, Debug, Zeroize, ZeroizeOnDrop)]
65#[serde(tag = "alg")]
66pub enum JWTPrivateKey {
67    ES256 { r#priv: String },
68    ES384 { r#priv: String },
69}
70
71impl JWTPrivateKey {
72    fn guess_key_family_from_pem(key: &str) -> Option<AlgorithmFamily> {
73        match EncodingKey::from_ec_pem(key.as_bytes()) {
74            Ok(_) => return Some(AlgorithmFamily::Ec),
75            Err(e) => {
76                tracing::trace!("Not a EC key: {e}");
77            }
78        }
79
80        match EncodingKey::from_rsa_pem(key.as_bytes()) {
81            Ok(_) => return Some(AlgorithmFamily::Rsa),
82            Err(e) => {
83                tracing::trace!("Not a RSA key: {e}");
84            }
85        }
86
87        match EncodingKey::from_ed_pem(key.as_bytes()) {
88            Ok(_) => return Some(AlgorithmFamily::Ed),
89            Err(e) => {
90                tracing::trace!("Not a Ecdsa key: {e}");
91            }
92        }
93
94        None
95    }
96
97    /// Parse private key from given PEM
98    pub fn parse_key_pem(key: &str) -> Res<Self> {
99        match Self::guess_key_family_from_pem(key) {
100            None => Err(BasicJwtError::CouldNotGuessAlgorithmFamily),
101            Some(AlgorithmFamily::Ec) => {
102                let pkey = pkcs8::PrivateKeyInfoOwned::from_pkcs8_pem(key)
103                    .map_err(BasicJwtError::ParseEcPrivateKey)?;
104
105                if pkey.algorithm.oid != const_oid::db::rfc5753::ID_EC_PUBLIC_KEY {
106                    return Err(BasicJwtError::InvalidEcAlgorithm(pkey.algorithm.oid));
107                }
108
109                let Some(params) = pkey.algorithm.parameters.as_ref() else {
110                    return Err(BasicJwtError::MissingEcParameters);
111                };
112
113                match params
114                    .decode_as::<ObjectIdentifier>()
115                    .map_err(BasicJwtError::DecodeEcParamsAsOID)?
116                {
117                    const_oid::db::rfc5912::SECP_256_R_1 => Ok(Self::ES256 {
118                        r#priv: key.to_string(),
119                    }),
120                    const_oid::db::rfc5912::SECP_384_R_1 => Ok(Self::ES384 {
121                        r#priv: key.to_string(),
122                    }),
123                    oid => Err(BasicJwtError::UnsupportedEcParams(oid)),
124                }
125            }
126            Some(f) => Err(BasicJwtError::UnsupportedAlgorithmFamily(f)),
127        }
128    }
129
130    /// Generate a new elliptic curve 256 signing key
131    pub fn generate_ec256_signing_key() -> Res<Self> {
132        let signing_key = p256::ecdsa::SigningKey::random(&mut OsRng);
133        let priv_pem = signing_key
134            .to_pkcs8_der()
135            .map_err(BasicJwtError::EncodeKeyToPKCS8)?
136            .to_pem("PRIVATE KEY", LineEnding::LF)
137            .map_err(BasicJwtError::EnvelopeEncodedKeyInPem)?
138            .to_string();
139
140        Ok(Self::ES256 { r#priv: priv_pem })
141    }
142
143    /// Generate a new ES384 signing key
144    pub fn generate_ec384_signing_key() -> Res<Self> {
145        let signing_key = p384::ecdsa::SigningKey::random(&mut OsRng);
146        let priv_pem = signing_key
147            .to_pkcs8_der()
148            .map_err(BasicJwtError::EncodeKeyToPKCS8)?
149            .to_pem("PRIVATE KEY", LineEnding::LF)
150            .map_err(BasicJwtError::EnvelopeEncodedKeyInPem)?
151            .to_string();
152
153        Ok(Self::ES384 { r#priv: priv_pem })
154    }
155
156    /// Get associated public key
157    pub fn to_public_key(&self) -> Res<JWTPublicKey> {
158        match self {
159            JWTPrivateKey::ES256 { r#priv } => {
160                let signing_key = p256::ecdsa::SigningKey::from_str(r#priv)
161                    .map_err(BasicJwtError::ParseSigningKey)?;
162
163                let pub_key = p256::ecdsa::VerifyingKey::from(signing_key);
164                let pub_pem = pub_key
165                    .to_public_key_pem(LineEnding::LF)
166                    .map_err(BasicJwtError::EncodePublicKeyToPem)?;
167
168                Ok(JWTPublicKey::ES256 { public: pub_pem })
169            }
170            JWTPrivateKey::ES384 { r#priv } => {
171                let signing_key = p384::ecdsa::SigningKey::from_str(r#priv)
172                    .map_err(BasicJwtError::ParseSigningKey)?;
173
174                let pub_key = p384::ecdsa::VerifyingKey::from(signing_key);
175                let pub_pem = pub_key
176                    .to_public_key_pem(LineEnding::LF)
177                    .map_err(BasicJwtError::EncodePublicKeyToPem)?;
178
179                Ok(JWTPublicKey::ES384 { public: pub_pem })
180            }
181        }
182    }
183
184    /// Get the decoding key & algorithm associated with a private key
185    pub fn get_encoding_key(&self) -> Res<(Zeroizing<EncodingKey>, Algorithm)> {
186        Ok(match self {
187            JWTPrivateKey::ES256 { r#priv } => (
188                Zeroizing::new(
189                    EncodingKey::from_ec_pem(r#priv.as_bytes())
190                        .map_err(BasicJwtError::ParseEncodingKey)?,
191                ),
192                Algorithm::ES256,
193            ),
194            JWTPrivateKey::ES384 { r#priv } => (
195                Zeroizing::new(
196                    EncodingKey::from_ec_pem(r#priv.as_bytes())
197                        .map_err(BasicJwtError::ParseEncodingKey)?,
198                ),
199                Algorithm::ES384,
200            ),
201        })
202    }
203
204    /// Sign a JWT
205    pub fn sign_jwt<C: Serialize>(&self, claims: &C) -> Res<String> {
206        let (encoding_key, algorithm) = self.get_encoding_key()?;
207
208        jsonwebtoken::encode(
209            &jsonwebtoken::Header::new(algorithm),
210            &claims,
211            &encoding_key,
212        )
213        .map_err(BasicJwtError::EncodeJWT)
214    }
215}
216
217impl JWTPublicKey {
218    /// Get the decoding key & algorithm associated with a public key
219    pub fn get_decoding_key(&self) -> Res<(DecodingKey, Algorithm)> {
220        Ok(match self {
221            JWTPublicKey::ES256 { public } => (
222                DecodingKey::from_ec_pem(public.as_bytes())
223                    .map_err(BasicJwtError::ParseDecodingKey)?,
224                Algorithm::ES256,
225            ),
226            JWTPublicKey::ES384 { public } => (
227                DecodingKey::from_ec_pem(public.as_bytes())
228                    .map_err(BasicJwtError::ParseDecodingKey)?,
229                Algorithm::ES384,
230            ),
231        })
232    }
233
234    /// Validate a given JWT
235    pub fn validate_jwt<E: DeserializeOwned + Clone>(&self, jwt: &str) -> Res<E> {
236        let (decoding_key, algorithm) = self.get_decoding_key()?;
237
238        let validation = Validation::new(algorithm);
239        Ok(jsonwebtoken::decode::<E>(jwt, &decoding_key, &validation)
240            .map_err(BasicJwtError::DecodeJWT)?
241            .claims)
242    }
243}
244
245#[cfg(test)]
246mod test {
247    use std::time::{SystemTime, UNIX_EPOCH};
248
249    use crate::JWTPrivateKey;
250    use serde::{Deserialize, Serialize};
251
252    fn time() -> u64 {
253        SystemTime::now()
254            .duration_since(UNIX_EPOCH)
255            .unwrap()
256            .as_secs()
257    }
258
259    #[derive(Debug, Serialize, Deserialize, Eq, PartialEq, Clone)]
260    pub struct Claims {
261        sub: String,
262        exp: u64,
263    }
264
265    impl Default for Claims {
266        fn default() -> Self {
267            Self {
268                sub: "my-sub".to_string(),
269                exp: time() + 100,
270            }
271        }
272    }
273
274    #[test]
275    fn jwt_encode_sign_verify_valid_p256() {
276        let priv_key = JWTPrivateKey::generate_ec256_signing_key().unwrap();
277        let pub_key = priv_key.to_public_key().unwrap();
278
279        let claims = Claims::default();
280        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
281        let claims_out = pub_key
282            .validate_jwt::<Claims>(&jwt)
283            .expect("Failed to validate JWT!");
284
285        assert_eq!(claims, claims_out)
286    }
287
288    #[test]
289    fn jwt_encode_sign_verify_valid_p384() {
290        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
291        let pub_key = priv_key.to_public_key().unwrap();
292
293        let claims = Claims::default();
294        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
295        let claims_out = pub_key
296            .validate_jwt::<Claims>(&jwt)
297            .expect("Failed to validate JWT!");
298
299        assert_eq!(claims, claims_out)
300    }
301
302    #[test]
303    fn parse_keys() {
304        const PEM_ONE: &str = r"-----BEGIN PRIVATE KEY-----
305MIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQgktVspUQwcArgXoLx
306hhGfw2yY6BB3/Cx8K2fIAck2KD2hRANCAAQ6PWADCFG5Ih6JMruZnuWKGmXEQtAf
3070ii/D5ubYKRW+iqx63h/OcW7DAABuh9Go0mRUf85A3PdoGAbkTAZAQoT
308-----END PRIVATE KEY-----
309";
310        let key_one = JWTPrivateKey::parse_key_pem(PEM_ONE).unwrap();
311
312        assert!(matches!(key_one, JWTPrivateKey::ES256 { .. }));
313
314        const PEM_TWO: &str = r"-----BEGIN PRIVATE KEY-----
315MIG2AgEAMBAGByqGSM49AgEGBSuBBAAiBIGeMIGbAgEBBDDMAHDNXDQfIZD0n52g
316ag4AxAJPK25TaPNK9TIbjx67Zs9U0JmvCbbbqFs6EiS08EyhZANiAAQZh4e/2BDk
317pECHm6hsokTKIn9EAgOQ0RtrWh02CZkTBJKvHC58KdwNB1eWSUHUKPKsrE2+h3cW
318Apd2mjEkBPSlCoTunNVAq+niutY+9LgcGZ3iFDTiI3GPDepQDtX8b6A=
319-----END PRIVATE KEY-----
320";
321        let key_one = JWTPrivateKey::parse_key_pem(PEM_TWO).unwrap();
322
323        assert!(matches!(key_one, JWTPrivateKey::ES384 { .. }));
324    }
325
326    #[test]
327    fn jwt_encode_sign_verify_invalid_key() {
328        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
329        let pub_key_2 = JWTPrivateKey::generate_ec384_signing_key()
330            .unwrap()
331            .to_public_key()
332            .unwrap();
333
334        let claims = Claims::default();
335        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
336        pub_key_2
337            .validate_jwt::<Claims>(&jwt)
338            .expect_err("JWT should not have validated!");
339    }
340
341    #[test]
342    fn jwt_verify_random_string() {
343        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
344        let pub_key = priv_key.to_public_key().unwrap();
345
346        pub_key
347            .validate_jwt::<Claims>("random_string")
348            .expect_err("JWT should not have validated!");
349    }
350
351    #[test]
352    fn jwt_expired() {
353        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
354        let pub_key = priv_key.to_public_key().unwrap();
355
356        let claims = Claims {
357            exp: time() - 100,
358            ..Default::default()
359        };
360        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
361        pub_key
362            .validate_jwt::<Claims>(&jwt)
363            .expect_err("JWT should not have validated!");
364    }
365
366    #[test]
367    fn jwt_invalid_signature() {
368        let priv_key = JWTPrivateKey::generate_ec384_signing_key().unwrap();
369        let pub_key = priv_key.to_public_key().unwrap();
370
371        let claims = Claims::default();
372        let jwt = priv_key.sign_jwt(&claims).expect("Failed to sign JWT!");
373        pub_key
374            .validate_jwt::<Claims>(&format!("{jwt}bad"))
375            .expect_err("JWT should not have validated!");
376    }
377}