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