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#[allow(private_bounds)]
33pub struct KeyRegistry<K: IsKey> {
34 named_keys: HashMap<String, usize>,
38 unnamed_keys: HashSet<usize>,
40 key_to_ordinal: HashMap<K::Inner, (usize, Option<String>)>,
42 active_keys: BTreeSet<usize>,
44 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 pub fn private() -> Self {
63 Self::default()
64 }
65}
66
67impl KeyRegistry<PublicKey> {
68 pub fn public() -> Self {
70 Self::default()
71 }
72}
73
74impl KeyRegistry<Key> {
75 pub fn new() -> Self {
77 Self::default()
78 }
79}
80
81#[allow(private_bounds)]
82impl<K: IsKey> KeyRegistry<K> {
83 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 pub fn add_key(&mut self, key: K) {
100 self.remove_key(&key);
101
102 let (kid, inner) = key.into_inner();
103
104 if let Some(kid) = &kid {
106 if self.named_keys.contains_key(kid) {
107 self.remove_kid(kid);
108 }
109 }
110
111 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 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 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 pub fn len(&self) -> usize {
150 self.key_to_ordinal.len()
151 }
152
153 pub fn is_empty(&self) -> bool {
155 self.key_to_ordinal.is_empty()
156 }
157
158 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 }
170 }
171 Ok(added)
172 }
173
174 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 }
185 }
186 Ok(added)
187 }
188
189 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 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 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 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 !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 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 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}