1use 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
14pub fn b64u(bytes: &[u8]) -> String {
16 URL_SAFE_NO_PAD.encode(bytes)
17}
18
19pub 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
27pub fn sha256_b64u(bytes: &[u8]) -> String {
30 b64u(&Sha256::digest(bytes))
31}
32
33pub 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#[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 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 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 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#[derive(Clone)]
126pub struct AgentKey {
127 signing: SigningKey,
128 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 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 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 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 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 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 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 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 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#[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 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 pub fn signing_input(&self) -> Vec<u8> {
284 format!("{}.{}", self.raw_header_b64, self.raw_payload_b64).into_bytes()
285 }
286
287 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 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}