Skip to main content

vault_core/
jwks.rs

1//! The public half of a `signing_key` as a JSON Web Key Set (Phase 24.5).
2//!
3//! A service that signs tokens publishes its *public* keys at a `jwks_uri`, and a
4//! rotation needs the old and the new key listed together for as long as tokens
5//! signed by the old one are still in flight - which is why the entry has a
6//! first-class "previous key" slot. This turns what the vault holds (a PEM or a
7//! JWK) into the document a verifier fetches.
8//!
9//! **Only public material is ever produced.** A JWK that arrives with private
10//! members (`d`, `p`, `q`, `dp`, `dq`, `qi`, `k`) has them stripped, and a PEM that
11//! is a private key is refused rather than derived from: publishing the wrong
12//! half is the one mistake a key set cannot take back. Nothing is verified or
13//! generated here; keys are re-encoded, never operated on.
14//!
15//! PEM `PUBLIC KEY` (SPKI) is read for RSA, EC (P-256, P-384, P-521), Ed25519 and
16//! X25519; `RSA PUBLIC KEY` (PKCS#1) for RSA. A key without a `kid` gets its RFC
17//! 7638 thumbprint, which is what a verifier would compute for it anyway.
18
19use base64::Engine;
20use serde_json::{json, Map, Value};
21use sha2::{Digest, Sha256};
22
23fn b64u(b: &[u8]) -> String {
24    base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(b)
25}
26
27/// One DER element: `(tag, contents, rest)`.
28fn tlv(b: &[u8]) -> Result<(u8, &[u8], &[u8]), String> {
29    let (&tag, b) = b.split_first().ok_or("Truncated key (DER)")?;
30    let (&l0, b) = b.split_first().ok_or("Truncated key (DER)")?;
31    let (len, b) = if l0 < 0x80 {
32        (usize::from(l0), b)
33    } else {
34        let n = usize::from(l0 & 0x7f);
35        if n == 0 || n > 4 || b.len() < n {
36            return Err("Unsupported DER length".into());
37        }
38        let len = b[..n]
39            .iter()
40            .fold(0usize, |a, &x| (a << 8) | usize::from(x));
41        (len, &b[n..])
42    };
43    if b.len() < len {
44        return Err("Truncated key (DER)".into());
45    }
46    Ok((tag, &b[..len], &b[len..]))
47}
48
49fn pem_body(text: &str, label: &str) -> Option<Vec<u8>> {
50    let begin = format!("-----BEGIN {label}-----");
51    let end = format!("-----END {label}-----");
52    let i = text.find(&begin)? + begin.len();
53    let j = text[i..].find(&end)? + i;
54    let b64: String = text[i..j].lines().map(str::trim).collect();
55    base64::engine::general_purpose::STANDARD
56        .decode(b64.as_bytes())
57        .ok()
58}
59
60/// An unsigned big-endian integer without its DER sign byte.
61fn uint(contents: &[u8]) -> &[u8] {
62    match contents {
63        [0, rest @ ..] if !rest.is_empty() => rest,
64        c => c,
65    }
66}
67
68fn rsa_jwk(seq: &[u8]) -> Result<Value, String> {
69    let (_, n, rest) = tlv(seq)?;
70    let (_, e, _) = tlv(rest)?;
71    Ok(json!({ "kty": "RSA", "n": b64u(uint(n)), "e": b64u(uint(e)) }))
72}
73
74/// A JWK from a PEM public key.
75fn jwk_from_pem(text: &str) -> Result<Value, String> {
76    if text.contains("PRIVATE KEY-----") {
77        return Err(
78            "That is a private key. Put the public key in the entry: a key set publishes the public half only.".into(),
79        );
80    }
81    if let Some(der) = pem_body(text, "RSA PUBLIC KEY") {
82        let (tag, seq, _) = tlv(&der)?;
83        if tag != 0x30 {
84            return Err("Not an RSA public key".into());
85        }
86        return rsa_jwk(seq);
87    }
88    let der = pem_body(text, "PUBLIC KEY")
89        .ok_or("Expected a PEM `PUBLIC KEY` (or `RSA PUBLIC KEY`) block, or a JWK in JSON")?;
90    let (tag, spki, _) = tlv(&der)?;
91    if tag != 0x30 {
92        return Err("Not a SubjectPublicKeyInfo".into());
93    }
94    let (tag, alg, rest) = tlv(spki)?;
95    if tag != 0x30 {
96        return Err("Not a SubjectPublicKeyInfo".into());
97    }
98    let (tag, bits, _) = tlv(rest)?;
99    if tag != 0x03 || bits.first() != Some(&0) {
100        return Err("Unsupported key bit string".into());
101    }
102    let key = &bits[1..];
103    let (_, oid, params) = tlv(alg)?;
104    const RSA: &[u8] = &[0x2a, 0x86, 0x48, 0x86, 0xf7, 0x0d, 0x01, 0x01, 0x01];
105    const EC: &[u8] = &[0x2a, 0x86, 0x48, 0xce, 0x3d, 0x02, 0x01];
106    const ED25519: &[u8] = &[0x2b, 0x65, 0x70];
107    const X25519: &[u8] = &[0x2b, 0x65, 0x6e];
108    match oid {
109        RSA => {
110            let (tag, seq, _) = tlv(key)?;
111            if tag != 0x30 {
112                return Err("Not an RSA public key".into());
113            }
114            rsa_jwk(seq)
115        }
116        EC => {
117            let (_, curve, _) = tlv(params)?;
118            let (name, size): (&str, usize) = match curve {
119                [0x2a, 0x86, 0x48, 0xce, 0x3d, 0x03, 0x01, 0x07] => ("P-256", 32),
120                [0x2b, 0x81, 0x04, 0x00, 0x22] => ("P-384", 48),
121                [0x2b, 0x81, 0x04, 0x00, 0x23] => ("P-521", 66),
122                _ => return Err("Only the P-256, P-384 and P-521 curves are supported".into()),
123            };
124            if key.len() != 1 + 2 * size || key[0] != 4 {
125                return Err("Only uncompressed EC points are supported".into());
126            }
127            Ok(json!({
128                "kty": "EC", "crv": name,
129                "x": b64u(&key[1..=size]), "y": b64u(&key[1 + size..]),
130            }))
131        }
132        ED25519 | X25519 if key.len() == 32 => Ok(json!({
133            "kty": "OKP",
134            "crv": if oid == ED25519 { "Ed25519" } else { "X25519" },
135            "x": b64u(key),
136        })),
137        _ => Err("Unsupported key type (RSA, EC P-256/384/521, Ed25519 and X25519 are)".into()),
138    }
139}
140
141const PRIVATE_MEMBERS: [&str; 8] = ["d", "p", "q", "dp", "dq", "qi", "k", "oth"];
142
143/// A public JWK from either spelling the vault might hold, with no private members.
144pub fn public_jwk(text: &str) -> Result<Value, String> {
145    let t = text.trim();
146    if t.starts_with('{') {
147        let mut v: Value =
148            serde_json::from_str(t).map_err(|e| format!("The JWK is not JSON: {e}"))?;
149        let o = v.as_object_mut().ok_or("A JWK is a JSON object")?;
150        if !o.contains_key("kty") {
151            return Err("The JWK has no `kty`".into());
152        }
153        for m in PRIVATE_MEMBERS {
154            o.remove(m);
155        }
156        return Ok(v);
157    }
158    jwk_from_pem(t)
159}
160
161/// The RFC 7638 thumbprint of a public JWK, which is the `kid` a verifier would
162/// derive and the one used when the entry names none.
163pub fn thumbprint(jwk: &Value) -> Result<String, String> {
164    let members: &[&str] = match jwk.get("kty").and_then(Value::as_str) {
165        Some("RSA") => &["e", "kty", "n"],
166        Some("EC") => &["crv", "kty", "x", "y"],
167        Some("OKP") => &["crv", "kty", "x"],
168        _ => return Err("Only RSA, EC and OKP keys have a thumbprint here".into()),
169    };
170    let mut m = Map::new();
171    for k in members {
172        let v = jwk
173            .get(*k)
174            .and_then(Value::as_str)
175            .ok_or_else(|| format!("The JWK has no `{k}`"))?;
176        m.insert((*k).into(), json!(v));
177    }
178    // serde_json's map is ordered by key, which is the order RFC 7638 asks for.
179    let canonical = serde_json::to_string(&Value::Object(m)).map_err(|e| e.to_string())?;
180    Ok(b64u(&Sha256::digest(canonical.as_bytes())))
181}
182
183/// `use` for a `usage` word the entry holds.
184fn use_for(usage: &str) -> &'static str {
185    match usage.trim().to_ascii_lowercase().as_str() {
186        "encrypt" | "enc" | "decrypt" => "enc",
187        _ => "sig",
188    }
189}
190
191/// One entry of the key set: the key with its `kid`, `alg` and `use`.
192pub fn entry_jwk(key: &str, kid: &str, alg: &str, usage: &str) -> Result<Value, String> {
193    let mut jwk = public_jwk(key)?;
194    let o = jwk.as_object_mut().ok_or("A JWK is a JSON object")?;
195    if !kid.is_empty() {
196        o.insert("kid".into(), json!(kid));
197    } else if !o.contains_key("kid") {
198        let t = thumbprint(&Value::Object(o.clone()))?;
199        o.insert("kid".into(), json!(t));
200    }
201    if !alg.is_empty() {
202        o.insert("alg".into(), json!(alg));
203    }
204    o.entry("use").or_insert_with(|| json!(use_for(usage)));
205    Ok(jwk)
206}
207
208/// The key set for a `signing_key` entry: the current key and, while a rotation is
209/// open, the previous one. `None` for a part that is empty.
210pub fn key_set(
211    current: (&str, &str, &str),
212    previous: (&str, &str),
213    alg: &str,
214    usage: &str,
215) -> Result<Value, String> {
216    let mut keys = Vec::new();
217    if current.0.trim().is_empty() {
218        return Err("`public_key` is empty; fill it in before emitting".into());
219    }
220    keys.push(entry_jwk(current.0, current.1, alg, usage)?);
221    if !previous.0.trim().is_empty() {
222        let k = entry_jwk(previous.0, previous.1, alg, usage)?;
223        if k["kid"] == keys[0]["kid"] {
224            return Err(
225                "The previous key has the same key id as the current one; a verifier could not tell them apart"
226                    .into(),
227            );
228        }
229        keys.push(k);
230    }
231    Ok(json!({ "keys": keys }))
232}
233
234#[cfg(test)]
235mod tests {
236    use super::*;
237
238    /// Public keys made by Python's `cryptography` package, with the JWK and the
239    /// RFC 7638 thumbprint it computed for each: a second implementation, not this
240    /// one agreeing with itself.
241    fn vectors() -> Value {
242        serde_json::from_str(include_str!("../tests/fixtures/jwk-vectors.json")).unwrap()
243    }
244
245    #[test]
246    fn every_supported_key_type_comes_out_as_the_jwk_an_independent_implementation_made() {
247        let v = vectors();
248        for name in ["rsa", "p256", "p384", "p521", "ed25519", "x25519"] {
249            let pem = v[name]["pem"].as_str().unwrap();
250            let got = public_jwk(pem).unwrap();
251            assert_eq!(got, v[name]["jwk"], "{name}");
252            assert_eq!(
253                thumbprint(&got).unwrap(),
254                v[name]["kid"].as_str().unwrap(),
255                "{name}"
256            );
257        }
258        // PKCS#1 spelling of the RSA key.
259        assert_eq!(
260            public_jwk(v["rsa"]["pkcs1"].as_str().unwrap()).unwrap(),
261            v["rsa"]["jwk"]
262        );
263    }
264
265    #[test]
266    fn the_thumbprint_matches_the_rfc_8037_example() {
267        let k =
268            json!({"kty":"OKP","crv":"Ed25519","x":"11qYAYKxCrfVS_7TyWQHOg7hcvPapiMlrwIaaPcHURo"});
269        assert_eq!(
270            thumbprint(&k).unwrap(),
271            "kPrK_qmxVWaYVA9wwBF6Iuo3vVzz7TxHCTwXBygrS4k"
272        );
273    }
274
275    #[test]
276    fn private_material_never_reaches_the_set() {
277        // A JWK pasted with its private members: they are dropped.
278        let with_d =
279            json!({"kty":"OKP","crv":"Ed25519","x":"AAAA","d":"SECRETSECRET","k":"x"}).to_string();
280        let k = public_jwk(&with_d).unwrap();
281        assert!(k.get("d").is_none() && k.get("k").is_none());
282        assert!(!k.to_string().contains("SECRETSECRET"));
283        // A private PEM is refused, not derived from.
284        let err =
285            public_jwk("-----BEGIN PRIVATE KEY-----\nAAAA\n-----END PRIVATE KEY-----").unwrap_err();
286        assert!(err.contains("private key"), "{err}");
287        let err =
288            public_jwk("-----BEGIN RSA PRIVATE KEY-----\nAAAA\n-----END RSA PRIVATE KEY-----")
289                .unwrap_err();
290        assert!(err.contains("private key"), "{err}");
291    }
292
293    #[test]
294    fn a_rotation_lists_both_keys_and_refuses_two_with_one_id() {
295        let v = vectors();
296        let a = v["ed25519"]["pem"].as_str().unwrap();
297        let b = v["p256"]["pem"].as_str().unwrap();
298        let set = key_set((a, "", ""), (b, ""), "EdDSA", "sign").unwrap();
299        let keys = set["keys"].as_array().unwrap();
300        assert_eq!(keys.len(), 2);
301        assert_eq!(
302            keys[0]["kid"], v["ed25519"]["kid"],
303            "no kid given: the thumbprint"
304        );
305        assert_eq!(keys[0]["use"], "sig");
306        assert_eq!(keys[0]["alg"], "EdDSA");
307        assert_eq!(keys[1]["crv"], "P-256");
308        // Only the current key while no rotation is open.
309        assert_eq!(
310            key_set((a, "k1", ""), ("", ""), "", "").unwrap()["keys"][0]["kid"],
311            "k1"
312        );
313        assert_eq!(
314            key_set((a, "k1", ""), ("", ""), "", "").unwrap()["keys"]
315                .as_array()
316                .unwrap()
317                .len(),
318            1
319        );
320        // The same id twice is a key set a verifier cannot use.
321        assert!(key_set((a, "same", ""), (b, "same"), "", "").is_err());
322        assert!(key_set(("", "", ""), ("", ""), "", "").is_err());
323        assert_eq!(entry_jwk(a, "", "", "encrypt").unwrap()["use"], "enc");
324    }
325
326    #[test]
327    fn garbage_is_an_error_and_never_a_panic() {
328        for bad in [
329            "",
330            "hello",
331            "{",
332            "[]",
333            "{\"a\":1}",
334            "-----BEGIN PUBLIC KEY-----\n!!!\n-----END PUBLIC KEY-----",
335        ] {
336            assert!(public_jwk(bad).is_err(), "{bad:?}");
337        }
338        for n in 0..200u32 {
339            let der: Vec<u8> = (0..40u32)
340                .map(|i| (i.wrapping_mul(n + 7) & 0xff) as u8)
341                .collect();
342            let pem = format!(
343                "-----BEGIN PUBLIC KEY-----\n{}\n-----END PUBLIC KEY-----",
344                base64::engine::general_purpose::STANDARD.encode(der)
345            );
346            let _ = public_jwk(&pem);
347        }
348    }
349}