authkestra_engine/token/jwk.rs
1use crate::auth::error::AuthError;
2use jsonwebtoken::DecodingKey;
3use serde::{Deserialize, Serialize};
4
5/// A JSON Web Key, as published at `/jwks.json`.
6///
7/// This struct is widened (not an enum) so that every existing call site
8/// that builds a `Jwk` with a plain struct literal — inside this crate and
9/// downstream — keeps compiling: it only needs two more fields (`crv`, `x`),
10/// both `None` for the RSA shape it already builds. See the `to_decoding_key`
11/// doc comment for why an enum/`#[serde(untagged)]` representation was
12/// rejected in favor of this.
13///
14/// Two shapes are represented today:
15/// - RSA (`kty: "RSA"`): `n`, `e` are populated; `crv`, `x` are `None`.
16/// - OKP/Ed25519 (`kty: "OKP"`): `crv` (always `"Ed25519"`), `x` are
17/// populated; `n`, `e` are `None`.
18///
19/// `None` fields are omitted from the serialized JSON (`skip_serializing_if`)
20/// so each shape's wire format matches its RFC exactly: RFC 7517 §6.3.1 for
21/// RSA (`kty`, `n`, `e`), RFC 8037 §2 for OKP (`kty`, `crv`, `x`). Neither
22/// shape ever emits the other's fields, and neither emits a stray `"n":null`
23/// / `"x":null`.
24#[derive(Debug, Clone, Serialize, Deserialize)]
25pub struct Jwk {
26 #[serde(skip_serializing_if = "Option::is_none")]
27 pub kid: Option<String>,
28 pub kty: String,
29 #[serde(skip_serializing_if = "Option::is_none")]
30 pub alg: Option<String>,
31 /// RSA modulus (base64url, unpadded). `None` for OKP keys.
32 #[serde(skip_serializing_if = "Option::is_none")]
33 pub n: Option<String>,
34 /// RSA public exponent (base64url, unpadded). `None` for OKP keys.
35 #[serde(skip_serializing_if = "Option::is_none")]
36 pub e: Option<String>,
37 /// OKP subtype curve name, e.g. `"Ed25519"` (RFC 8037 §2). `None` for
38 /// RSA keys.
39 #[serde(skip_serializing_if = "Option::is_none")]
40 pub crv: Option<String>,
41 /// OKP public key (base64url, unpadded, RFC 8037 §2). `None` for RSA
42 /// keys.
43 #[serde(skip_serializing_if = "Option::is_none")]
44 pub x: Option<String>,
45}
46
47impl Jwk {
48 /// Derives a `DecodingKey` from this JWK, dispatching on `kty`.
49 ///
50 /// Supports `"RSA"` (unchanged from before this key gained the OKP
51 /// shape) and `"OKP"` with `crv: "Ed25519"` (RFC 8037). Any other `kty`,
52 /// or an OKP key advertising an unsupported curve, is rejected.
53 pub fn to_decoding_key(&self) -> Result<DecodingKey, AuthError> {
54 match self.kty.as_str() {
55 "RSA" => {
56 let n = self
57 .n
58 .as_ref()
59 .ok_or_else(|| AuthError::Token("Missing 'n' component in JWK".to_string()))?;
60 let e = self
61 .e
62 .as_ref()
63 .ok_or_else(|| AuthError::Token("Missing 'e' component in JWK".to_string()))?;
64
65 DecodingKey::from_rsa_components(n, e).map_err(|e| AuthError::Token(e.to_string()))
66 }
67 "OKP" => {
68 match self.crv.as_deref() {
69 Some("Ed25519") => {}
70 Some(other) => {
71 return Err(AuthError::Token(format!(
72 "Unsupported OKP curve '{}' in JWK",
73 other
74 )));
75 }
76 None => {
77 return Err(AuthError::Token(
78 "Missing 'crv' component in OKP JWK".to_string(),
79 ));
80 }
81 }
82
83 let x_str = self
84 .x
85 .as_ref()
86 .ok_or_else(|| AuthError::Token("Missing 'x' component in JWK".to_string()))?;
87
88 authkestra_crypto_util::parse_ed25519_verifying_key_strict(x_str)
89 .map_err(|e| AuthError::Token(e.to_string()))?;
90
91 DecodingKey::from_ed_components(x_str).map_err(|e| AuthError::Token(e.to_string()))
92 }
93 other => Err(AuthError::Token(format!(
94 "Unsupported JWK 'kty' '{}' — only RSA and OKP are supported",
95 other
96 ))),
97 }
98 }
99}
100
101#[cfg(test)]
102mod tests {
103 use super::*;
104
105 #[test]
106 fn rejects_low_order_ed25519_key() {
107 // The identity point: the canonical universal low-order vector.
108 let identity_b64 = "AQAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA";
109 let jwk = Jwk {
110 kid: None,
111 kty: "OKP".to_string(),
112 alg: None,
113 n: None,
114 e: None,
115 crv: Some("Ed25519".to_string()),
116 x: Some(identity_b64.to_string()),
117 };
118
119 let err = jwk
120 .to_decoding_key()
121 .expect_err("should reject low order point");
122 assert!(
123 err.to_string().contains("low-order"),
124 "expected low-order point rejection, got: {}",
125 err
126 );
127 }
128}