Skip to main content

treeship_core/vi/
jws.rs

1//! base64url, SHA-256, P-256 keys and ES256 compact JWS, byte-for-byte the
2//! way the VI reference SDK does them.
3
4use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine};
5use p256::ecdsa::{
6    signature::Signer as _, signature::Verifier as _, Signature, SigningKey, VerifyingKey,
7};
8use serde::{Deserialize, Serialize};
9use serde_json::Value;
10use sha2::{Digest, Sha256};
11
12use super::ViError;
13
14/// base64url without padding, as every VI string field uses.
15pub fn b64u(bytes: &[u8]) -> String {
16    URL_SAFE_NO_PAD.encode(bytes)
17}
18
19/// Decode base64url, accepting the unpadded form (and padded, leniently).
20pub fn b64u_decode(s: &str) -> Result<Vec<u8>, ViError> {
21    let trimmed = s.trim_end_matches('=');
22    URL_SAFE_NO_PAD
23        .decode(trimmed)
24        .map_err(|e| ViError::Malformed(format!("base64url: {e}")))
25}
26
27/// `B64U(SHA-256(bytes))`: the spec's hash form for `sd_hash`, disclosure
28/// hashes and `checkout_hash`.
29pub fn sha256_b64u(bytes: &[u8]) -> String {
30    b64u(&Sha256::digest(bytes))
31}
32
33/// Compact JSON exactly as Python's `json.dumps(obj, separators=(",", ":"))`
34/// emits it with its default `ensure_ascii=True`: non-ASCII characters
35/// become `\uXXXX` escapes (surrogate pairs above the BMP). The reference
36/// verifier re-serializes a token's header and payload from the parsed
37/// dicts before checking the signature, so our signed bytes must survive
38/// that round trip unchanged.
39pub fn json_compact_ascii(v: &Value) -> String {
40    let raw = serde_json::to_string(v).expect("serde_json::Value serializes");
41    if raw.is_ascii() {
42        return raw;
43    }
44    let mut out = String::with_capacity(raw.len() + 16);
45    for ch in raw.chars() {
46        if ch.is_ascii() {
47            out.push(ch);
48        } else {
49            let mut buf = [0u16; 2];
50            for unit in ch.encode_utf16(&mut buf) {
51                use std::fmt::Write;
52                let _ = write!(out, "\\u{unit:04x}");
53            }
54        }
55    }
56    out
57}
58
59/// A P-256 public key in JWK form (`kty`, `crv`, `x`, `y`, optional `kid`).
60#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
61pub struct Jwk {
62    pub kty: String,
63    pub crv: String,
64    pub x: String,
65    pub y: String,
66    #[serde(default, skip_serializing_if = "Option::is_none")]
67    pub kid: Option<String>,
68}
69
70impl Jwk {
71    /// Parse a JWK from a JSON value, accepting extra members (`d`, `use`).
72    pub fn from_value(v: &Value) -> Result<Self, ViError> {
73        let get = |k: &str| -> Result<String, ViError> {
74            v.get(k)
75                .and_then(Value::as_str)
76                .map(str::to_string)
77                .ok_or_else(|| ViError::Key(format!("jwk missing '{k}'")))
78        };
79        let kty = get("kty")?;
80        let crv = get("crv")?;
81        if kty != "EC" || crv != "P-256" {
82            return Err(ViError::Key(format!(
83                "jwk must be EC/P-256, got {kty}/{crv}"
84            )));
85        }
86        Ok(Self {
87            kty,
88            crv,
89            x: get("x")?,
90            y: get("y")?,
91            kid: v.get("kid").and_then(Value::as_str).map(str::to_string),
92        })
93    }
94
95    pub fn to_value(&self) -> Value {
96        serde_json::to_value(self).expect("Jwk serializes")
97    }
98
99    /// The verifying key this JWK names.
100    pub fn verifying_key(&self) -> Result<VerifyingKey, ViError> {
101        let x = b64u_decode(&self.x)?;
102        let y = b64u_decode(&self.y)?;
103        if x.len() != 32 || y.len() != 32 {
104            return Err(ViError::Key("jwk x/y must be 32 bytes each".into()));
105        }
106        let mut sec1 = Vec::with_capacity(65);
107        sec1.push(0x04);
108        sec1.extend_from_slice(&x);
109        sec1.extend_from_slice(&y);
110        VerifyingKey::from_sec1_bytes(&sec1).map_err(|e| ViError::Key(format!("jwk point: {e}")))
111    }
112
113    /// Verify a raw `r || s` ES256 signature over `signing_input`.
114    pub fn verify(&self, signing_input: &[u8], sig: &[u8]) -> Result<(), ViError> {
115        let vk = self.verifying_key()?;
116        let sig = Signature::from_slice(sig)
117            .map_err(|e| ViError::Signature(format!("ES256 signature bytes: {e}")))?;
118        vk.verify(signing_input, &sig)
119            .map_err(|_| ViError::Signature("ES256 signature did not verify".into()))
120    }
121}
122
123/// A P-256 signing key: the agent's key that L2 binds under `cnf` and that
124/// signs L3a and L3b.
125#[derive(Clone)]
126pub struct AgentKey {
127    signing: SigningKey,
128    /// The key id the L2 mandate's `cnf.jwk.kid` carries and every L3
129    /// header repeats. Chosen at generation; stable for the key's life.
130    pub kid: String,
131}
132
133impl std::fmt::Debug for AgentKey {
134    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
135        f.debug_struct("AgentKey").field("kid", &self.kid).finish()
136    }
137}
138
139impl AgentKey {
140    /// Generate a fresh key. The kid is the first 16 hex characters of
141    /// SHA-256 over the SEC1 compressed point, prefixed `vik_`.
142    pub fn generate() -> Self {
143        let signing = SigningKey::random(&mut rand::rngs::OsRng);
144        let kid = kid_for(signing.verifying_key());
145        Self { signing, kid }
146    }
147
148    /// Rebuild from the 32-byte private scalar and an explicit kid.
149    pub fn from_secret(d: &[u8], kid: impl Into<String>) -> Result<Self, ViError> {
150        let signing =
151            SigningKey::from_slice(d).map_err(|e| ViError::Key(format!("P-256 scalar: {e}")))?;
152        Ok(Self {
153            signing,
154            kid: kid.into(),
155        })
156    }
157
158    /// Import a private JWK (`d` present). Keeps the JWK's `kid` when it has
159    /// one, else derives the same kid `generate` would.
160    pub fn from_private_jwk(v: &Value) -> Result<Self, ViError> {
161        let d = v
162            .get("d")
163            .and_then(Value::as_str)
164            .ok_or_else(|| ViError::Key("private jwk missing 'd'".into()))?;
165        let d = b64u_decode(d)?;
166        let signing =
167            SigningKey::from_slice(&d).map_err(|e| ViError::Key(format!("P-256 scalar: {e}")))?;
168        let kid = v
169            .get("kid")
170            .and_then(Value::as_str)
171            .map(str::to_string)
172            .unwrap_or_else(|| kid_for(signing.verifying_key()));
173        let key = Self { signing, kid };
174        // The public half in the JWK, if present, must be this key's.
175        if v.get("x").is_some() {
176            let claimed = Jwk::from_value(v)?;
177            let ours = key.public_jwk();
178            if claimed.x != ours.x || claimed.y != ours.y {
179                return Err(ViError::Key("private jwk x/y do not match d".into()));
180            }
181        }
182        Ok(key)
183    }
184
185    /// The 32-byte private scalar, for sealing at rest.
186    pub fn secret_bytes(&self) -> [u8; 32] {
187        let b = self.signing.to_bytes();
188        let mut out = [0u8; 32];
189        out.copy_from_slice(&b);
190        out
191    }
192
193    /// The public JWK with this key's kid, as a wallet binds it under `cnf`.
194    pub fn public_jwk(&self) -> Jwk {
195        let point = self.signing.verifying_key().to_encoded_point(false);
196        Jwk {
197            kty: "EC".into(),
198            crv: "P-256".into(),
199            x: b64u(point.x().expect("uncompressed point has x")),
200            y: b64u(point.y().expect("uncompressed point has y")),
201            kid: Some(self.kid.clone()),
202        }
203    }
204
205    /// The private JWK (`d` included). For export to a wallet or HSM only.
206    pub fn private_jwk(&self) -> Value {
207        let mut v = self.public_jwk().to_value();
208        v["d"] = Value::String(b64u(&self.secret_bytes()));
209        v
210    }
211
212    /// ES256: raw `r || s`, 64 bytes, deterministic (RFC 6979).
213    pub fn sign(&self, signing_input: &[u8]) -> Vec<u8> {
214        let sig: Signature = self.signing.sign(signing_input);
215        sig.to_bytes().to_vec()
216    }
217}
218
219fn kid_for(vk: &VerifyingKey) -> String {
220    let compressed = vk.to_encoded_point(true);
221    let h = Sha256::digest(compressed.as_bytes());
222    format!("vik_{}", hex::encode(&h[..8]))
223}
224
225/// A decoded compact JWS: the raw segments are kept so a parsed token
226/// re-serializes byte-identically (which `sd_hash` depends on).
227#[derive(Debug, Clone)]
228pub struct CompactJws {
229    pub header: Value,
230    pub payload: Value,
231    pub raw_header_b64: String,
232    pub raw_payload_b64: String,
233    pub signature: Vec<u8>,
234}
235
236impl CompactJws {
237    /// Sign `header` and `payload` with `key`, producing `h.p.s`.
238    pub fn sign(header: &Value, payload: &Value, key: &AgentKey) -> Self {
239        let h = b64u(json_compact_ascii(header).as_bytes());
240        let p = b64u(json_compact_ascii(payload).as_bytes());
241        let signing_input = format!("{h}.{p}");
242        let signature = key.sign(signing_input.as_bytes());
243        Self {
244            header: header.clone(),
245            payload: payload.clone(),
246            raw_header_b64: h,
247            raw_payload_b64: p,
248            signature,
249        }
250    }
251
252    pub fn parse(token: &str) -> Result<Self, ViError> {
253        let parts: Vec<&str> = token.split('.').collect();
254        if parts.len() != 3 {
255            return Err(ViError::Malformed(format!(
256                "jwt: expected 3 parts, got {}",
257                parts.len()
258            )));
259        }
260        let header: Value = serde_json::from_slice(&b64u_decode(parts[0])?)
261            .map_err(|e| ViError::Malformed(format!("jwt header json: {e}")))?;
262        let payload: Value = serde_json::from_slice(&b64u_decode(parts[1])?)
263            .map_err(|e| ViError::Malformed(format!("jwt payload json: {e}")))?;
264        Ok(Self {
265            header,
266            payload,
267            raw_header_b64: parts[0].to_string(),
268            raw_payload_b64: parts[1].to_string(),
269            signature: b64u_decode(parts[2])?,
270        })
271    }
272
273    pub fn serialize(&self) -> String {
274        format!(
275            "{}.{}.{}",
276            self.raw_header_b64,
277            self.raw_payload_b64,
278            b64u(&self.signature)
279        )
280    }
281
282    /// The bytes the signature covers.
283    pub fn signing_input(&self) -> Vec<u8> {
284        format!("{}.{}", self.raw_header_b64, self.raw_payload_b64).into_bytes()
285    }
286
287    /// Verify with a P-256 public key.
288    pub fn verify(&self, jwk: &Jwk) -> Result<(), ViError> {
289        jwk.verify(&self.signing_input(), &self.signature)
290    }
291}
292
293#[cfg(test)]
294mod tests {
295    use super::*;
296
297    #[test]
298    fn compact_ascii_matches_python_ensure_ascii() {
299        let v = serde_json::json!({"a": "caf\u{e9} \u{1F600}", "n": 1});
300        // Python: json.dumps({"a": "café 😀", "n": 1}, separators=(",",":"))
301        assert_eq!(
302            json_compact_ascii(&v),
303            r#"{"a":"caf\u00e9 \ud83d\ude00","n":1}"#
304        );
305    }
306
307    #[test]
308    fn sign_and_verify_round_trip() {
309        let key = AgentKey::generate();
310        let header = serde_json::json!({"alg":"ES256","typ":"kb-sd-jwt","kid":key.kid});
311        let payload = serde_json::json!({"nonce":"n","aud":"https://m.example","iat":1});
312        let jws = CompactJws::sign(&header, &payload, &key);
313        let parsed = CompactJws::parse(&jws.serialize()).unwrap();
314        assert_eq!(parsed.header["kid"], key.kid);
315        parsed.verify(&key.public_jwk()).unwrap();
316        let other = AgentKey::generate();
317        assert!(parsed.verify(&other.public_jwk()).is_err());
318    }
319
320    #[test]
321    fn private_jwk_round_trip_keeps_kid_and_checks_point() {
322        let key = AgentKey::generate();
323        let priv_jwk = key.private_jwk();
324        let back = AgentKey::from_private_jwk(&priv_jwk).unwrap();
325        assert_eq!(back.kid, key.kid);
326        assert_eq!(back.public_jwk(), key.public_jwk());
327        let mut bad = priv_jwk.clone();
328        bad["x"] = Value::String(b64u(&[7u8; 32]));
329        assert!(AgentKey::from_private_jwk(&bad).is_err());
330    }
331}