1use alloc::{
37 borrow::Cow,
38 string::{String, ToString},
39 vec::Vec,
40};
41use core::fmt;
42
43use serde::{Deserialize, Deserializer, Serialize, Serializer};
44use sha2::digest::{Digest, Output};
45
46use crate::alg::SecretBytes;
47
48#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
50#[non_exhaustive]
51pub enum KeyType {
52 Rsa,
54 EllipticCurve,
57 Symmetric,
59 KeyPair,
61}
62
63impl fmt::Display for KeyType {
64 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
65 formatter.write_str(match self {
66 Self::Rsa => "RSA",
67 Self::EllipticCurve => "EC",
68 Self::Symmetric => "oct",
69 Self::KeyPair => "OKP",
70 })
71 }
72}
73
74#[derive(Debug)]
77#[non_exhaustive]
78pub enum JwkError {
79 NoField(String),
81 UnexpectedKeyType {
83 expected: KeyType,
85 actual: KeyType,
87 },
88 UnexpectedValue {
90 field: String,
92 expected: String,
94 actual: String,
96 },
97 UnexpectedLen {
99 field: String,
101 expected: usize,
103 actual: usize,
105 },
106 MismatchedKeys,
108 Custom(anyhow::Error),
110}
111
112impl fmt::Display for JwkError {
113 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
114 match self {
115 Self::UnexpectedKeyType { expected, actual } => {
116 write!(
117 formatter,
118 "unexpected key type: {actual} (expected {expected})"
119 )
120 }
121 Self::NoField(field) => write!(formatter, "field `{field}` is absent from JWK"),
122 Self::UnexpectedValue {
123 field,
124 expected,
125 actual,
126 } => {
127 write!(
128 formatter,
129 "field `{field}` has unexpected value (expected: {expected}, got: {actual})"
130 )
131 }
132 Self::UnexpectedLen {
133 field,
134 expected,
135 actual,
136 } => {
137 write!(
138 formatter,
139 "field `{field}` has unexpected length (expected: {expected}, got: {actual})"
140 )
141 }
142 Self::MismatchedKeys => {
143 formatter.write_str("private and public keys encoded in JWK do not match")
144 }
145 Self::Custom(err) => fmt::Display::fmt(err, formatter),
146 }
147 }
148}
149
150impl core::error::Error for JwkError {
151 fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
152 match self {
153 Self::Custom(err) => Some(err.as_ref()),
154 _ => None,
155 }
156 }
157}
158
159impl JwkError {
160 pub fn custom(err: impl Into<anyhow::Error>) -> Self {
162 Self::Custom(err.into())
163 }
164
165 pub(crate) fn key_type(jwk: &JsonWebKey<'_>, expected: KeyType) -> Self {
166 let actual = jwk.key_type();
167 debug_assert_ne!(actual, expected);
168 Self::UnexpectedKeyType { actual, expected }
169 }
170}
171
172impl Serialize for SecretBytes<'_> {
173 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
174 base64url::serialize(self.as_ref(), serializer)
175 }
176}
177
178impl<'de> Deserialize<'de> for SecretBytes<'_> {
179 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
180 base64url::deserialize(deserializer).map(SecretBytes::new)
181 }
182}
183
184#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
224#[serde(tag = "kty")]
225#[non_exhaustive]
226pub enum JsonWebKey<'a> {
227 #[serde(rename = "RSA")]
229 Rsa {
230 #[serde(rename = "n", with = "base64url")]
232 modulus: Cow<'a, [u8]>,
233 #[serde(rename = "e", with = "base64url")]
235 public_exponent: Cow<'a, [u8]>,
236 #[serde(flatten)]
238 private_parts: Option<RsaPrivateParts<'a>>,
239 },
240 #[serde(rename = "EC")]
242 EllipticCurve {
243 #[serde(rename = "crv")]
245 curve: Cow<'a, str>,
246 #[serde(with = "base64url")]
248 x: Cow<'a, [u8]>,
249 #[serde(with = "base64url")]
251 y: Cow<'a, [u8]>,
252 #[serde(rename = "d", default, skip_serializing_if = "Option::is_none")]
254 secret: Option<SecretBytes<'a>>,
255 },
256 #[serde(rename = "oct")]
258 Symmetric {
259 #[serde(rename = "k")]
261 secret: SecretBytes<'a>,
262 },
263 #[serde(rename = "OKP")]
265 KeyPair {
266 #[serde(rename = "crv")]
268 curve: Cow<'a, str>,
269 #[serde(with = "base64url")]
272 x: Cow<'a, [u8]>,
273 #[serde(rename = "d", default, skip_serializing_if = "Option::is_none")]
275 secret: Option<SecretBytes<'a>>,
276 },
277}
278
279impl JsonWebKey<'_> {
280 pub fn key_type(&self) -> KeyType {
282 match self {
283 Self::Rsa { .. } => KeyType::Rsa,
284 Self::EllipticCurve { .. } => KeyType::EllipticCurve,
285 Self::Symmetric { .. } => KeyType::Symmetric,
286 Self::KeyPair { .. } => KeyType::KeyPair,
287 }
288 }
289
290 pub fn is_signing_key(&self) -> bool {
292 match self {
293 Self::Rsa { private_parts, .. } => private_parts.is_some(),
294 Self::EllipticCurve { secret, .. } | Self::KeyPair { secret, .. } => secret.is_some(),
295 Self::Symmetric { .. } => true,
296 }
297 }
298
299 #[must_use]
301 pub fn to_verifying_key(&self) -> Self {
302 match self {
303 Self::Rsa {
304 modulus,
305 public_exponent,
306 ..
307 } => Self::Rsa {
308 modulus: modulus.clone(),
309 public_exponent: public_exponent.clone(),
310 private_parts: None,
311 },
312
313 Self::EllipticCurve { curve, x, y, .. } => Self::EllipticCurve {
314 curve: curve.clone(),
315 x: x.clone(),
316 y: y.clone(),
317 secret: None,
318 },
319
320 Self::Symmetric { secret } => Self::Symmetric {
321 secret: secret.clone(),
322 },
323
324 Self::KeyPair { curve, x, .. } => Self::KeyPair {
325 curve: curve.clone(),
326 x: x.clone(),
327 secret: None,
328 },
329 }
330 }
331
332 pub fn thumbprint<D: Digest>(&self) -> Output<D> {
337 let hashed_key = if self.is_signing_key() {
338 Cow::Owned(self.to_verifying_key())
339 } else {
340 Cow::Borrowed(self)
341 };
342 D::digest(hashed_key.to_string().as_bytes())
343 }
344}
345
346impl fmt::Display for JsonWebKey<'_> {
347 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
349 let json_value = serde_json::to_value(self).expect("Cannot convert JsonWebKey to JSON");
350 let json_value = json_value.as_object().unwrap();
351 let mut json_entries: Vec<_> = json_value.iter().collect();
354 json_entries.sort_unstable_by_key(|(x, _)| *x);
355
356 formatter.write_str("{")?;
357 let field_count = json_entries.len();
358 for (i, (name, value)) in json_entries.into_iter().enumerate() {
359 write!(formatter, "\"{name}\":{value}")?;
360 if i + 1 < field_count {
361 formatter.write_str(",")?;
362 }
363 }
364 formatter.write_str("}")
365 }
366}
367
368#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
376pub struct RsaPrivateParts<'a> {
377 #[serde(rename = "d")]
379 pub private_exponent: SecretBytes<'a>,
380 #[serde(rename = "p")]
382 pub prime_factor_p: SecretBytes<'a>,
383 #[serde(rename = "q")]
385 pub prime_factor_q: SecretBytes<'a>,
386 #[serde(rename = "dp", default, skip_serializing_if = "Option::is_none")]
388 pub p_crt_exponent: Option<SecretBytes<'a>>,
389 #[serde(rename = "dq", default, skip_serializing_if = "Option::is_none")]
391 pub q_crt_exponent: Option<SecretBytes<'a>>,
392 #[serde(rename = "qi", default, skip_serializing_if = "Option::is_none")]
394 pub q_crt_coefficient: Option<SecretBytes<'a>>,
395 #[serde(rename = "oth", default, skip_serializing_if = "Vec::is_empty")]
397 pub other_prime_factors: Vec<RsaPrimeFactor<'a>>,
398}
399
400#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
408pub struct RsaPrimeFactor<'a> {
409 #[serde(rename = "r")]
411 pub factor: SecretBytes<'a>,
412 #[serde(rename = "d", default, skip_serializing_if = "Option::is_none")]
414 pub crt_exponent: Option<SecretBytes<'a>>,
415 #[serde(rename = "t", default, skip_serializing_if = "Option::is_none")]
417 pub crt_coefficient: Option<SecretBytes<'a>>,
418}
419
420#[cfg(any(
421 feature = "es256k",
422 feature = "k256",
423 feature = "p256",
424 feature = "exonum-crypto",
425 feature = "ed25519-dalek",
426 feature = "ed25519-compact"
427))]
428mod helpers {
429 use super::{JsonWebKey, JwkError};
430 use crate::{Algorithm, alg::SigningKey};
431
432 impl JsonWebKey<'_> {
433 pub(crate) fn ensure_curve(curve: &str, expected: &str) -> Result<(), JwkError> {
434 if curve == expected {
435 Ok(())
436 } else {
437 Err(JwkError::UnexpectedValue {
438 field: "crv".into(),
439 expected: expected.into(),
440 actual: curve.into(),
441 })
442 }
443 }
444
445 pub(crate) fn ensure_len(
446 field: &str,
447 bytes: &[u8],
448 expected_len: usize,
449 ) -> Result<(), JwkError> {
450 if bytes.len() == expected_len {
451 Ok(())
452 } else {
453 Err(JwkError::UnexpectedLen {
454 field: field.into(),
455 expected: expected_len,
456 actual: bytes.len(),
457 })
458 }
459 }
460
461 pub(crate) fn ensure_key_match<Alg, K>(&self, signing_key: K) -> Result<K, JwkError>
464 where
465 Alg: Algorithm<SigningKey = K>,
466 K: SigningKey<Alg>,
467 Alg::VerifyingKey: for<'jwk> TryFrom<&'jwk Self, Error = JwkError> + PartialEq,
468 {
469 let verifying_key = <Alg::VerifyingKey>::try_from(self)?;
470 if verifying_key == signing_key.to_verifying_key() {
471 Ok(signing_key)
472 } else {
473 Err(JwkError::MismatchedKeys)
474 }
475 }
476 }
477}
478
479mod base64url {
480 use alloc::{borrow::Cow, vec::Vec};
481 use core::fmt;
482
483 use base64ct::{Base64UrlUnpadded, Encoding};
484 use serde::{
485 Deserializer, Serializer,
486 de::{Error as DeError, Unexpected, Visitor},
487 };
488
489 pub fn serialize<S>(value: &[u8], serializer: S) -> Result<S::Ok, S::Error>
490 where
491 S: Serializer,
492 {
493 if serializer.is_human_readable() {
494 serializer.serialize_str(&Base64UrlUnpadded::encode_string(value))
495 } else {
496 serializer.serialize_bytes(value)
497 }
498 }
499
500 pub fn deserialize<'de, D>(deserializer: D) -> Result<Cow<'static, [u8]>, D::Error>
501 where
502 D: Deserializer<'de>,
503 {
504 struct Base64Visitor;
505
506 impl Visitor<'_> for Base64Visitor {
507 type Value = Vec<u8>;
508
509 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
510 formatter.write_str("base64url-encoded data")
511 }
512
513 fn visit_str<E: DeError>(self, value: &str) -> Result<Self::Value, E> {
514 Base64UrlUnpadded::decode_vec(value)
515 .map_err(|_| E::invalid_value(Unexpected::Str(value), &self))
516 }
517
518 fn visit_bytes<E: DeError>(self, value: &[u8]) -> Result<Self::Value, E> {
519 Ok(value.to_vec())
520 }
521
522 fn visit_byte_buf<E: DeError>(self, value: Vec<u8>) -> Result<Self::Value, E> {
523 Ok(value)
524 }
525 }
526
527 struct BytesVisitor;
528
529 impl Visitor<'_> for BytesVisitor {
530 type Value = Vec<u8>;
531
532 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
533 formatter.write_str("byte buffer")
534 }
535
536 fn visit_bytes<E: DeError>(self, value: &[u8]) -> Result<Self::Value, E> {
537 Ok(value.to_vec())
538 }
539
540 fn visit_byte_buf<E: DeError>(self, value: Vec<u8>) -> Result<Self::Value, E> {
541 Ok(value)
542 }
543 }
544
545 let maybe_bytes = if deserializer.is_human_readable() {
546 deserializer.deserialize_str(Base64Visitor)
547 } else {
548 deserializer.deserialize_bytes(BytesVisitor)
549 };
550 maybe_bytes.map(Cow::Owned)
551 }
552}
553
554#[cfg(test)]
555mod tests {
556 use assert_matches::assert_matches;
557
558 use super::*;
559 use crate::alg::Hs256Key;
560
561 fn create_jwk() -> JsonWebKey<'static> {
562 JsonWebKey::KeyPair {
563 curve: Cow::Borrowed("Ed25519"),
564 x: Cow::Borrowed(b"test"),
565 secret: None,
566 }
567 }
568
569 #[test]
570 fn serializing_jwk() {
571 let jwk = create_jwk();
572
573 let json = serde_json::to_value(&jwk).unwrap();
574 assert_eq!(
575 json,
576 serde_json::json!({ "crv": "Ed25519", "kty": "OKP", "x": "dGVzdA" })
577 );
578
579 let restored: JsonWebKey<'_> = serde_json::from_value(json).unwrap();
580 assert_eq!(restored, jwk);
581 }
582
583 #[test]
584 fn jwk_deserialization_errors() {
585 let missing_field_json = r#"{"crv":"Ed25519"}"#;
586 let missing_field_err = serde_json::from_str::<JsonWebKey<'_>>(missing_field_json)
587 .unwrap_err()
588 .to_string();
589 assert!(
590 missing_field_err.contains("missing field `kty`"),
591 "{missing_field_err}"
592 );
593
594 let base64_json = r#"{"crv":"Ed25519","kty":"OKP","x":"??"}"#;
595 let base64_err = serde_json::from_str::<JsonWebKey<'_>>(base64_json)
596 .unwrap_err()
597 .to_string();
598 assert!(
599 base64_err.contains("invalid value: string \"??\""),
600 "{base64_err}"
601 );
602 assert!(
603 base64_err.contains("base64url-encoded data"),
604 "{base64_err}"
605 );
606 }
607
608 #[test]
609 fn extra_jwk_fields() {
610 #[derive(Debug, Serialize, Deserialize)]
611 struct ExtendedJsonWebKey<'a, T> {
612 #[serde(flatten)]
613 base: JsonWebKey<'a>,
614 #[serde(flatten)]
615 extra: T,
616 }
617
618 #[derive(Debug, Deserialize)]
619 struct Extra {
620 #[serde(rename = "kid")]
621 key_id: String,
622 #[serde(rename = "use")]
623 key_use: KeyUse,
624 }
625
626 #[derive(Debug, Deserialize, PartialEq)]
627 enum KeyUse {
628 #[serde(rename = "sig")]
629 Signature,
630 #[serde(rename = "enc")]
631 Encryption,
632 }
633
634 let json_str = r#"
635 { "kty": "oct", "kid": "my-unique-key", "k": "dGVzdA", "use": "sig" }
636 "#;
637 let jwk: ExtendedJsonWebKey<'_, Extra> = serde_json::from_str(json_str).unwrap();
638
639 assert_matches!(&jwk.base, JsonWebKey::Symmetric { secret } if secret.as_ref() == b"test");
640 assert_eq!(jwk.extra.key_id, "my-unique-key");
641 assert_eq!(jwk.extra.key_use, KeyUse::Signature);
642
643 let key = Hs256Key::try_from(&jwk.base).unwrap();
644 let jwk_from_key = JsonWebKey::from(&key);
645
646 assert_matches!(
647 jwk_from_key,
648 JsonWebKey::Symmetric { secret } if secret.as_ref() == b"test"
649 );
650 }
651
652 #[test]
653 #[cfg(feature = "ciborium")]
654 fn jwk_with_cbor() {
655 let key = JsonWebKey::KeyPair {
656 curve: Cow::Borrowed("Ed25519"),
657 x: Cow::Borrowed(b"public"),
658 secret: Some(SecretBytes::borrowed(b"private")),
659 };
660 let mut bytes = Vec::new();
661 ciborium::into_writer(&key, &mut bytes).unwrap();
662 assert!(bytes.windows(6).any(|window| window == b"public"));
663 assert!(bytes.windows(7).any(|window| window == b"private"));
664
665 let restored: JsonWebKey<'_> = ciborium::from_reader(&bytes[..]).unwrap();
666 assert_eq!(restored, key);
667 }
668}