Skip to main content

gel_jwt/
registry.rs

1use crate::{
2    bare_key::{SerializedKey, SerializedKeys},
3    key::*,
4    Any, KeyError, OpaqueValidationFailureReason, SignatureError, SigningContext,
5    ValidationContext, ValidationError,
6};
7use std::{
8    collections::{BTreeSet, HashMap, HashSet},
9    fmt::Debug,
10};
11
12pub(crate) trait IsKey {
13    type Inner: std::hash::Hash + Eq + Debug + Clone;
14
15    fn inner(&self) -> &Self::Inner;
16    fn key_type(inner: &Self::Inner) -> KeyType;
17    fn from_inner(kid: Option<String>, inner: Self::Inner) -> Self;
18    fn into_inner(self) -> (Option<String>, Self::Inner);
19    fn get_serialized_key(key: SerializedKey) -> Option<Self>
20    where
21        Self: Sized;
22    fn to_serialized_key(kid: Option<&str>, inner: &Self::Inner) -> SerializedKey;
23    fn from_pem(pem: &str) -> Result<Vec<Result<Self, KeyError>>, KeyError>
24    where
25        Self: Sized;
26    fn to_pem(inner: &Self::Inner) -> String;
27    fn encoding_key(inner: &Self::Inner) -> Option<&jsonwebtoken::EncodingKey>;
28    fn decoding_key(inner: &Self::Inner) -> &jsonwebtoken::DecodingKey;
29}
30
31/// A collection of [`Key`] or [`PublicKey`] objects.
32#[allow(private_bounds)]
33pub struct KeyRegistry<K: IsKey> {
34    // TODO: this could probably be optimized, especially if we can
35    // generate key signatures.
36    /// Map from key identifier (kid) to ordinal
37    named_keys: HashMap<String, usize>,
38    /// Set of ordinals for unnamed keys
39    unnamed_keys: HashSet<usize>,
40    /// Map from key to ordinal/kid for quick lookup
41    key_to_ordinal: HashMap<K::Inner, (usize, Option<String>)>,
42    /// Set of active ordinals
43    active_keys: BTreeSet<usize>,
44    /// Next ordinal to use for a new key
45    next: usize,
46}
47
48impl<K: IsKey> Default for KeyRegistry<K> {
49    fn default() -> Self {
50        Self {
51            named_keys: HashMap::default(),
52            unnamed_keys: HashSet::default(),
53            key_to_ordinal: HashMap::default(),
54            active_keys: BTreeSet::default(),
55            next: 0,
56        }
57    }
58}
59
60impl KeyRegistry<PrivateKey> {
61    /// Create a new private key registry.
62    pub fn private() -> Self {
63        Self::default()
64    }
65}
66
67impl KeyRegistry<PublicKey> {
68    /// Create a new public key registry.
69    pub fn public() -> Self {
70        Self::default()
71    }
72}
73
74impl KeyRegistry<Key> {
75    /// Create a new key registry that may contain both private and public keys.
76    pub fn new() -> Self {
77        Self::default()
78    }
79}
80
81#[allow(private_bounds)]
82impl<K: IsKey> KeyRegistry<K> {
83    /// Clear the registry.
84    pub fn clear(&mut self) {
85        *self = Self::default();
86    }
87
88    pub fn into_keys(self) -> impl Iterator<Item = K> {
89        self.key_to_ordinal
90            .into_iter()
91            .map(|(key, (_, kid))| K::from_inner(kid, key))
92    }
93
94    /// Add a key to the registry. If the key already exists, it will be
95    /// replaced. If the key specifies the same kid as another key already
96    /// added, the new key will replace the old one.
97    ///
98    /// Adding a key, even if it already exists, will make it the active key.
99    pub fn add_key(&mut self, key: K) {
100        self.remove_key(&key);
101
102        let (kid, inner) = key.into_inner();
103
104        // If the kid still exists, we need to remove that key too
105        if let Some(kid) = &kid {
106            if self.named_keys.contains_key(kid) {
107                self.remove_kid(kid);
108            }
109        }
110
111        // Key is new, add it to the registry
112        let ordinal = self.next;
113        self.next += 1;
114        self.key_to_ordinal.insert(inner, (ordinal, kid.clone()));
115        self.active_keys.insert(ordinal);
116
117        if let Some(kid) = kid {
118            self.named_keys.insert(kid, ordinal);
119        } else {
120            self.unnamed_keys.insert(ordinal);
121        }
122    }
123
124    /// Remove a key from the registry by its key.
125    pub fn remove_key(&mut self, key: &K) {
126        let inner = key.inner();
127        if let Some((ordinal, kid)) = self.key_to_ordinal.remove(inner) {
128            if let Some(kid) = kid {
129                self.named_keys.remove(&kid);
130            } else {
131                self.unnamed_keys.remove(&ordinal);
132            }
133            self.active_keys.remove(&ordinal);
134        }
135    }
136
137    /// Remove a key from the registry by its kid. Note: O(N).
138    pub fn remove_kid(&mut self, kid: &str) -> bool {
139        if let Some(ordinal) = self.named_keys.remove(kid) {
140            self.active_keys.remove(&ordinal);
141            self.key_to_ordinal.retain(|_, &mut (v, _)| v != ordinal);
142            true
143        } else {
144            false
145        }
146    }
147
148    /// Get the number of keys in the registry.
149    pub fn len(&self) -> usize {
150        self.key_to_ordinal.len()
151    }
152
153    /// Check if the registry is empty.
154    pub fn is_empty(&self) -> bool {
155        self.key_to_ordinal.is_empty()
156    }
157
158    /// Add keys from a JWKSet.
159    pub fn add_from_jwkset(&mut self, jwkset: &str) -> Result<usize, KeyError> {
160        let loaded: SerializedKeys =
161            serde_json::from_str(jwkset).map_err(|_| KeyError::InvalidJson)?;
162        let mut added = 0;
163        for key in loaded.keys {
164            if let Some(key) = K::get_serialized_key(key) {
165                self.add_key(key);
166                added += 1;
167            } else {
168                // TODO: log unknown or invalid key
169            }
170        }
171        Ok(added)
172    }
173
174    /// Add keys from a PEM file.
175    pub fn add_from_pem(&mut self, pem: &str) -> Result<usize, KeyError> {
176        let keys = K::from_pem(pem)?;
177        let mut added = 0;
178        for key in keys {
179            if let Ok(key) = key {
180                self.add_key(key);
181                added += 1;
182            } else {
183                // TODO: log unknown or invalid key
184            }
185        }
186        Ok(added)
187    }
188
189    /// Add keys from a source string which can be either a JWK set or a PEM file with
190    /// 1 or more keys.
191    pub fn add_from_any(&mut self, source: &str) -> Result<usize, KeyError> {
192        let source = source.trim();
193        if source.is_empty() {
194            return Ok(0);
195        }
196
197        // Get the first non-whitespace character
198        let first_char = source.chars().next().unwrap_or_default();
199        if first_char == '{' {
200            self.add_from_jwkset(source)
201        } else if first_char == '-' {
202            self.add_from_pem(source)
203        } else {
204            Err(KeyError::UnsupportedKeyType(format!(
205                "Expected JWK set or PEM file, got {first_char}"
206            )))
207        }
208    }
209
210    pub fn to_pem(&self) -> String {
211        let mut pem = String::new();
212        for (k, (_, _)) in &self.key_to_ordinal {
213            pem.push_str(&K::to_pem(k));
214        }
215        pem
216    }
217
218    pub fn to_json(&self) -> Result<String, KeyError> {
219        serde_json::to_string(&SerializedKeys {
220            keys: self
221                .key_to_ordinal
222                .iter()
223                .map(|(k, (_, kid))| K::to_serialized_key(kid.as_deref(), k))
224                .collect(),
225        })
226        .map_err(|_| KeyError::EncodeError)
227    }
228
229    /// Get the active key and kid.
230    fn active_key(&self) -> Option<(Option<&str>, &K::Inner)> {
231        if let Some(&i) = self.active_keys.last() {
232            for (k, &(v, ref kid)) in &self.key_to_ordinal {
233                if v == i {
234                    if let Some(kid) = kid {
235                        return Some((Some(kid.as_str()), k));
236                    } else {
237                        return Some((None, k));
238                    }
239                }
240            }
241        }
242        None
243    }
244
245    /// Decode a token without validating the signature.
246    pub fn unsafely_decode_without_validation(
247        &self,
248        token: &str,
249    ) -> Result<HashMap<String, Any>, ValidationError> {
250        let mut validation = jsonwebtoken::Validation::new(jsonwebtoken::Algorithm::default());
251        validation.insecure_disable_signature_validation();
252        validation.required_spec_claims.clear();
253        validation.validate_exp = false;
254        validation.validate_nbf = false;
255        validation.validate_aud = false;
256        let decoding_key = jsonwebtoken::DecodingKey::from_secret(b"");
257
258        let decoded =
259            jsonwebtoken::decode::<HashMap<String, Any>>(token, &decoding_key, &validation)
260                .map_err(|e| {
261                    ValidationError::Invalid(OpaqueValidationFailureReason::Failure(format!(
262                        "{:?}",
263                        e.kind()
264                    )))
265                })?;
266
267        Ok(decoded.claims)
268    }
269
270    pub fn validate(
271        &self,
272        token: &str,
273        ctx: &ValidationContext,
274    ) -> Result<HashMap<String, Any>, ValidationError> {
275        // If we have a named key that matches, use that.
276        if !self.named_keys.is_empty() {
277            if let Ok(header) = jsonwebtoken::decode_header(token) {
278                if let Some(header_kid) = header.kid {
279                    for (key, (_, kid)) in &self.key_to_ordinal {
280                        if kid.as_deref() == Some(header_kid.as_str()) {
281                            return validate_token(
282                                K::key_type(key),
283                                K::decoding_key(key),
284                                None,
285                                token,
286                                ctx,
287                            );
288                        }
289                    }
290                }
291            }
292        }
293
294        let mut result = None;
295        for (key, _) in self.key_to_ordinal.iter() {
296            let last_result =
297                validate_token(K::key_type(key), K::decoding_key(key), None, token, ctx);
298            match last_result {
299                Ok(result) => return Ok(result),
300                Err(e) => result = Some(e),
301            }
302        }
303        Err(result.unwrap_or(OpaqueValidationFailureReason::NoAppropriateKey.into()))
304    }
305
306    pub fn sign(
307        &self,
308        claims: HashMap<String, Any>,
309        ctx: &SigningContext,
310    ) -> Result<String, SignatureError> {
311        let (kid, key) = self.active_key().ok_or(SignatureError::NoAppropriateKey)?;
312        let encoding_key = K::encoding_key(key).ok_or(SignatureError::NoAppropriateKey)?;
313        sign_token(K::key_type(key), encoding_key, kid, claims, ctx)
314    }
315}
316
317impl KeyRegistry<PrivateKey> {
318    pub fn can_sign(&self) -> bool {
319        self.has_private_keys() || self.has_symmetric_keys()
320    }
321
322    pub fn can_validate(&self) -> bool {
323        self.has_public_keys() || self.has_symmetric_keys()
324    }
325
326    pub fn has_private_keys(&self) -> bool {
327        !self.is_empty()
328    }
329
330    pub fn has_public_keys(&self) -> bool {
331        self.key_to_ordinal
332            .iter()
333            .any(|(k, _)| k.bare_key.key_type() != KeyType::HS256)
334    }
335
336    pub fn has_symmetric_keys(&self) -> bool {
337        self.key_to_ordinal
338            .iter()
339            .any(|(k, _)| k.bare_key.key_type() == KeyType::HS256)
340    }
341
342    #[cfg(feature = "keygen")]
343    pub fn generate_key(&mut self, kid: Option<String>, key_type: KeyType) -> Result<(), KeyError> {
344        let key = PrivateKey::generate(kid, key_type)?;
345        self.add_key(key);
346        Ok(())
347    }
348}
349
350impl KeyRegistry<PublicKey> {
351    pub fn can_sign(&self) -> bool {
352        self.has_private_keys() || self.has_symmetric_keys()
353    }
354
355    pub fn can_validate(&self) -> bool {
356        self.has_public_keys() || self.has_symmetric_keys()
357    }
358
359    pub fn has_public_keys(&self) -> bool {
360        !self.is_empty()
361    }
362
363    pub fn has_private_keys(&self) -> bool {
364        false
365    }
366
367    pub fn has_symmetric_keys(&self) -> bool {
368        false
369    }
370}
371
372impl KeyRegistry<Key> {
373    pub fn can_sign(&self) -> bool {
374        self.has_private_keys() || self.has_symmetric_keys()
375    }
376
377    pub fn can_validate(&self) -> bool {
378        self.has_public_keys() || self.has_symmetric_keys()
379    }
380
381    pub fn has_private_keys(&self) -> bool {
382        for k in self.key_to_ordinal.keys() {
383            if let KeyInner::Private(_) = k {
384                return true;
385            }
386        }
387        false
388    }
389
390    pub fn has_public_keys(&self) -> bool {
391        for k in self.key_to_ordinal.keys() {
392            if let KeyInner::Public(_) = k {
393                return true;
394            }
395            if let KeyInner::Private(k) = k {
396                if k.bare_key.key_type() != KeyType::HS256 {
397                    return true;
398                }
399            }
400        }
401        false
402    }
403
404    pub fn has_symmetric_keys(&self) -> bool {
405        for k in self.key_to_ordinal.keys() {
406            if let KeyInner::Private(k) = k {
407                if k.bare_key.key_type() == KeyType::HS256 {
408                    return true;
409                }
410            }
411        }
412        false
413    }
414
415    /// Export the registry as a PEM file containing only the public keys.
416    /// This will fail if the registry contains symmetric keys.
417    pub fn to_pem_public(&self) -> Result<String, KeyError> {
418        let mut pem = String::new();
419        for (k, (_, _)) in &self.key_to_ordinal {
420            match k {
421                KeyInner::Private(k) => {
422                    pem.push_str(&k.bare_key.to_pem_public()?);
423                }
424                KeyInner::Public(k) => {
425                    pem.push_str(&k.bare_key.to_pem());
426                }
427            }
428        }
429        Ok(pem)
430    }
431
432    /// Export the registry as a JSON object containing only the public keys.
433    /// This will fail if the registry contains symmetric keys.
434    pub fn to_json_public(&self) -> Result<String, KeyError> {
435        let mut keys = Vec::new();
436        for (k, (_, kid)) in &self.key_to_ordinal {
437            match k {
438                KeyInner::Private(k) => {
439                    keys.push(SerializedKey::Public(
440                        kid.clone(),
441                        k.bare_key.to_public()?.clone_key(),
442                    ));
443                }
444                KeyInner::Public(k) => {
445                    keys.push(SerializedKey::Public(kid.clone(), k.bare_key.clone_key()));
446                }
447            }
448        }
449        serde_json::to_string(&SerializedKeys { keys }).map_err(|_| KeyError::EncodeError)
450    }
451
452    #[cfg(feature = "keygen")]
453    pub fn generate_key(&mut self, kid: Option<String>, key_type: KeyType) -> Result<(), KeyError> {
454        let key = PrivateKey::generate(kid, key_type)?;
455        self.add_key(key.into());
456        Ok(())
457    }
458}