Skip to main content

fabric_protocol/
lib.rs

1//! ForgeWire wire protocol primitives.
2//!
3//! Stage C.1: ed25519 sign + verify over canonical JSON envelopes, byte-compatible
4//! with the existing Python implementation in `scripts/remote/runner/identity.py`
5//! and `scripts/remote/runner/runner_capabilities.py::canonical_payload`.
6//!
7//! The canonical form is `serde_json` with sorted keys and the compact separators
8//! `(',', ':')`, matching Python's
9//! `json.dumps(payload, sort_keys=True, separators=(",", ":"))`.
10
11use ed25519_dalek::{Signature, SigningKey, Verifier, VerifyingKey, SECRET_KEY_LENGTH};
12use serde_json::{Map, Value};
13use thiserror::Error;
14
15/// Errors surfaced by the protocol layer.
16///
17/// Verification errors are deliberately collapsed to a single variant so callers
18/// cannot use the discriminant to fingerprint why a given attempt failed; the
19/// Python side returns a bare `False` for the same reason.
20#[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
41/// Canonical-JSON encoding of a `serde_json::Value`.
42///
43/// Object keys are emitted in sorted order; separators are compact (`,` and `:`).
44/// This matches `json.dumps(payload, sort_keys=True, separators=(",", ":"))`.
45pub 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        // Numbers/bools/null: serde_json's default encoding matches Python's
80        // json.dumps for these scalar shapes (e.g. integers as bare digits,
81        // booleans as `true`/`false`, null as `null`).
82        other => {
83            serde_json::to_writer(&mut *out, other)?;
84        }
85    }
86    Ok(())
87}
88
89/// Encode a string the way `json.dumps(..., ensure_ascii=True)` does:
90/// - quoted with `"`
91/// - escapes for `"`, `\\`, and the standard short escapes (`\b\f\n\r\t`)
92/// - control characters `< 0x20` as `\u00XX`
93/// - any non-ASCII character as `\uXXXX` (with UTF-16 surrogate pairs above U+FFFF)
94fn 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                // Printable ASCII (excluding the controls / quote / backslash already handled).
110                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                    // Encode as UTF-16 surrogate pair to match Python's
118                    // ensure_ascii output for non-BMP code points.
119                    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        // `& 0xf` above guarantees n <= 15, so this cast can never truncate --
136        // `try_from` would only add an unreachable error path.
137        #[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
164/// Verify an ed25519 signature over `payload` using a hex-encoded public key.
165///
166/// Returns `Ok(true)` only if the signature is valid; `Ok(false)` for any
167/// recoverable mismatch (invalid signature). Returns `Err` only for input
168/// shape problems (bad hex, wrong length).
169pub 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
188/// Sign `payload` with a hex-encoded 32-byte ed25519 secret key.
189///
190/// Returns the 64-byte signature as a lowercase hex string, matching the Python
191/// runner identity's `sign(payload).hex()` output.
192pub 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
205/// Convenience: verify a signed envelope (object) using a hex public key and
206/// hex signature; canonicalizes the envelope before verifying.
207pub 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
216/// Convenience: sign an envelope (object) with a hex secret key.
217pub 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        // Deterministic: ed25519 from a fixed 32-byte seed.
270        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        // Tamper with the envelope and ensure verification fails.
280        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    // M2.9.0: cross-language signed-command-canonical fixture.
293    // Verifies that `canonicalize` produces byte-identical output to Python's
294    // `canonical_payload` for each envelope case in the fixture file. The
295    // agent-kind case is the additive-only proof (loom fields absent → unchanged
296    // canonical).
297    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}