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