1use ed25519_dalek::{Signature, SigningKey, Verifier, VerifyingKey, SECRET_KEY_LENGTH};
12use serde_json::{Map, Value};
13use thiserror::Error;
14
15#[derive(Debug, Error)]
21pub enum ProtocolError {
22 #[error("invalid hex: {0}")]
23 Hex(#[from] hex::FromHexError),
24
25 #[error("invalid public key length: expected 32 bytes, got {0}")]
26 PublicKeyLength(usize),
27
28 #[error("invalid private key length: expected 32 bytes, got {0}")]
29 PrivateKeyLength(usize),
30
31 #[error("invalid signature length: expected 64 bytes, got {0}")]
32 SignatureLength(usize),
33
34 #[error("invalid public key bytes")]
35 PublicKey,
36
37 #[error("canonicalization failed: {0}")]
38 Canonicalization(#[from] serde_json::Error),
39}
40
41pub fn canonicalize(value: &Value) -> Result<Vec<u8>, ProtocolError> {
46 let mut out = Vec::with_capacity(64);
47 write_canonical(value, &mut out)?;
48 Ok(out)
49}
50
51fn write_canonical(value: &Value, out: &mut Vec<u8>) -> Result<(), ProtocolError> {
52 match value {
53 Value::Object(map) => {
54 let sorted = sort_object(map);
55 out.push(b'{');
56 for (i, (k, v)) in sorted.iter().enumerate() {
57 if i > 0 {
58 out.push(b',');
59 }
60 write_ascii_string(k, out);
61 out.push(b':');
62 write_canonical(v, out)?;
63 }
64 out.push(b'}');
65 }
66 Value::Array(arr) => {
67 out.push(b'[');
68 for (i, v) in arr.iter().enumerate() {
69 if i > 0 {
70 out.push(b',');
71 }
72 write_canonical(v, out)?;
73 }
74 out.push(b']');
75 }
76 Value::String(s) => {
77 write_ascii_string(s, out);
78 }
79 other => {
83 serde_json::to_writer(&mut *out, other)?;
84 }
85 }
86 Ok(())
87}
88
89fn write_ascii_string(s: &str, out: &mut Vec<u8>) {
95 out.push(b'"');
96 for ch in s.chars() {
97 match ch {
98 '"' => out.extend_from_slice(b"\\\""),
99 '\\' => out.extend_from_slice(b"\\\\"),
100 '\u{0008}' => out.extend_from_slice(b"\\b"),
101 '\u{0009}' => out.extend_from_slice(b"\\t"),
102 '\u{000a}' => out.extend_from_slice(b"\\n"),
103 '\u{000c}' => out.extend_from_slice(b"\\f"),
104 '\u{000d}' => out.extend_from_slice(b"\\r"),
105 c if (c as u32) < 0x20 => {
106 write_unicode_escape(c as u32, out);
107 }
108 c if (c as u32) < 0x7f => {
109 out.push(c as u8);
111 }
112 c => {
113 let cp = c as u32;
114 if cp <= 0xffff {
115 write_unicode_escape(cp, out);
116 } else {
117 let v = cp - 0x10000;
120 let high = 0xd800 + (v >> 10);
121 let low = 0xdc00 + (v & 0x3ff);
122 write_unicode_escape(high, out);
123 write_unicode_escape(low, out);
124 }
125 }
126 }
127 }
128 out.push(b'"');
129}
130
131fn write_unicode_escape(cp: u32, out: &mut Vec<u8>) {
132 out.extend_from_slice(b"\\u");
133 let nibbles = [(cp >> 12) & 0xf, (cp >> 8) & 0xf, (cp >> 4) & 0xf, cp & 0xf];
134 for n in nibbles {
135 #[allow(clippy::cast_possible_truncation)]
138 let n = n as u8;
139 let byte = if n < 10 { b'0' + n } else { b'a' + (n - 10) };
140 out.push(byte);
141 }
142}
143
144fn sort_object(map: &Map<String, Value>) -> Vec<(&String, &Value)> {
145 let mut entries: Vec<(&String, &Value)> = map.iter().collect();
146 entries.sort_by(|a, b| a.0.cmp(b.0));
147 entries
148}
149
150fn decode_hex_fixed<const N: usize>(s: &str) -> Result<[u8; N], ProtocolError> {
151 let raw = hex::decode(s)?;
152 if raw.len() != N {
153 return match N {
154 32 if s.len() == 64 => Err(ProtocolError::PublicKeyLength(raw.len())),
155 64 => Err(ProtocolError::SignatureLength(raw.len())),
156 _ => Err(ProtocolError::PublicKeyLength(raw.len())),
157 };
158 }
159 let mut out = [0u8; N];
160 out.copy_from_slice(&raw);
161 Ok(out)
162}
163
164pub fn verify_signature_hex(
170 public_key_hex: &str,
171 payload: &[u8],
172 signature_hex: &str,
173) -> Result<bool, ProtocolError> {
174 let pk_bytes = decode_hex_fixed::<32>(public_key_hex)
175 .map_err(|_| ProtocolError::PublicKeyLength(public_key_hex.len() / 2))?;
176 let sig_bytes_raw = hex::decode(signature_hex)?;
177 if sig_bytes_raw.len() != 64 {
178 return Err(ProtocolError::SignatureLength(sig_bytes_raw.len()));
179 }
180 let mut sig_bytes = [0u8; 64];
181 sig_bytes.copy_from_slice(&sig_bytes_raw);
182
183 let pk = VerifyingKey::from_bytes(&pk_bytes).map_err(|_| ProtocolError::PublicKey)?;
184 let sig = Signature::from_bytes(&sig_bytes);
185 Ok(pk.verify(payload, &sig).is_ok())
186}
187
188pub fn sign_payload_hex(secret_key_hex: &str, payload: &[u8]) -> Result<String, ProtocolError> {
193 let sk_raw = hex::decode(secret_key_hex)?;
194 if sk_raw.len() != SECRET_KEY_LENGTH {
195 return Err(ProtocolError::PrivateKeyLength(sk_raw.len()));
196 }
197 let mut sk_bytes = [0u8; SECRET_KEY_LENGTH];
198 sk_bytes.copy_from_slice(&sk_raw);
199
200 let signing = SigningKey::from_bytes(&sk_bytes);
201 let sig: Signature = ed25519_dalek::Signer::sign(&signing, payload);
202 Ok(hex::encode(sig.to_bytes()))
203}
204
205pub fn verify_envelope_hex(
208 public_key_hex: &str,
209 envelope: &Value,
210 signature_hex: &str,
211) -> Result<bool, ProtocolError> {
212 let canonical = canonicalize(envelope)?;
213 verify_signature_hex(public_key_hex, &canonical, signature_hex)
214}
215
216pub fn sign_envelope_hex(secret_key_hex: &str, envelope: &Value) -> Result<String, ProtocolError> {
218 let canonical = canonicalize(envelope)?;
219 sign_payload_hex(secret_key_hex, &canonical)
220}
221
222#[cfg(test)]
223mod tests {
224 use super::*;
225 use serde_json::json;
226
227 #[test]
228 fn canonicalize_sorts_keys() {
229 let v = json!({"b": 1, "a": 2, "c": [3, 4]});
230 let s = canonicalize(&v).unwrap();
231 assert_eq!(
232 std::str::from_utf8(&s).unwrap(),
233 r#"{"a":2,"b":1,"c":[3,4]}"#
234 );
235 }
236
237 #[test]
238 fn canonicalize_nested_objects() {
239 let v = json!({"outer": {"z": 1, "a": 2}, "x": 3});
240 let s = canonicalize(&v).unwrap();
241 assert_eq!(
242 std::str::from_utf8(&s).unwrap(),
243 r#"{"outer":{"a":2,"z":1},"x":3}"#
244 );
245 }
246
247 #[test]
248 fn canonicalize_ensure_ascii_escapes_non_bmp() {
249 let v = json!({"unicode": "héllo", "emoji": "🚀"});
250 let s = canonicalize(&v).unwrap();
251 assert_eq!(
252 std::str::from_utf8(&s).unwrap(),
253 r#"{"emoji":"\ud83d\ude80","unicode":"h\u00e9llo"}"#
254 );
255 }
256
257 #[test]
258 fn canonicalize_escapes_control_chars() {
259 let v = json!({"k": "a\nb\tc\u{0001}d\""});
260 let s = canonicalize(&v).unwrap();
261 assert_eq!(
262 std::str::from_utf8(&s).unwrap(),
263 r#"{"k":"a\nb\tc\u0001d\""}"#
264 );
265 }
266
267 #[test]
268 fn sign_then_verify_roundtrip() {
269 let sk_hex = "1".repeat(64);
271 let signing =
272 SigningKey::from_bytes(&<[u8; 32]>::try_from(hex::decode(&sk_hex).unwrap()).unwrap());
273 let pk_hex = hex::encode(signing.verifying_key().to_bytes());
274 let envelope = json!({"op": "register", "runner_id": "abc", "ts": 1234});
275
276 let sig = sign_envelope_hex(&sk_hex, &envelope).unwrap();
277 assert!(verify_envelope_hex(&pk_hex, &envelope, &sig).unwrap());
278
279 let tampered = json!({"op": "register", "runner_id": "abc", "ts": 9999});
281 assert!(!verify_envelope_hex(&pk_hex, &tampered, &sig).unwrap());
282 }
283
284 #[test]
285 fn bad_signature_length_is_err() {
286 let pk_hex = "0".repeat(64);
287 let payload = b"hello";
288 let result = verify_signature_hex(&pk_hex, payload, "ab");
289 assert!(matches!(result, Err(ProtocolError::SignatureLength(_))));
290 }
291
292 const CANONICAL_FIXTURE: &str =
298 include_str!("../../../tests/fixtures/phase_2_9/signed_command_canonical.json");
299
300 #[test]
301 fn canonical_matches_signed_command_fixture() {
302 let doc: serde_json::Value =
303 serde_json::from_str(CANONICAL_FIXTURE).expect("parse signed_command_canonical.json");
304 let cases = doc["cases"].as_array().expect("cases array");
305 assert!(!cases.is_empty(), "fixture must have at least one case");
306
307 for case in cases {
308 let name = case["name"].as_str().unwrap_or("?");
309 let envelope = &case["envelope"];
310 let expected = case["expected_canonical"]
311 .as_str()
312 .unwrap_or_else(|| panic!("case {name}: expected_canonical must be a string"));
313
314 let got = canonicalize(envelope)
315 .unwrap_or_else(|e| panic!("case {name}: canonicalize error: {e}"));
316 assert_eq!(
317 std::str::from_utf8(&got).unwrap(),
318 expected,
319 "case {name}: canonical mismatch"
320 );
321 }
322 }
323}