Skip to main content

headgate_crypto/
lib.rs

1//! Client-side AES-256-GCM payload encryption. Stores only ever see ciphertext.
2
3use std::{collections::BTreeMap, sync::Arc};
4
5use aes_gcm::{
6    Aes256Gcm, KeyInit, Nonce,
7    aead::{Aead, Payload},
8};
9use headgate::{CodecError, Envelope, JobCtx, JobError, Registry, Task};
10use rand::RngCore;
11
12const MAGIC: &[u8; 5] = b"HGEC\x01";
13const NONCE_LEN: usize = 12;
14
15#[derive(Debug)]
16pub struct CryptoError(String);
17impl std::fmt::Display for CryptoError {
18    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
19        f.write_str(&self.0)
20    }
21}
22impl std::error::Error for CryptoError {}
23
24pub trait KeyProvider: Send + Sync + 'static {
25    fn active_key(&self) -> Result<(String, [u8; 32]), CryptoError>;
26    fn key(&self, id: &str) -> Result<[u8; 32], CryptoError>;
27}
28
29#[derive(Clone)]
30pub struct StaticKeyring {
31    active: String,
32    keys: BTreeMap<String, [u8; 32]>,
33}
34
35impl StaticKeyring {
36    pub fn new(
37        active: impl Into<String>,
38        keys: BTreeMap<String, [u8; 32]>,
39    ) -> Result<Self, CryptoError> {
40        let active = active.into();
41        if !keys.contains_key(&active) {
42            return Err(CryptoError("active encryption key is missing".into()));
43        }
44        Ok(Self { active, keys })
45    }
46}
47
48impl KeyProvider for StaticKeyring {
49    fn active_key(&self) -> Result<(String, [u8; 32]), CryptoError> {
50        Ok((self.active.clone(), self.keys[&self.active]))
51    }
52    fn key(&self, id: &str) -> Result<[u8; 32], CryptoError> {
53        self.keys
54            .get(id)
55            .copied()
56            .ok_or_else(|| CryptoError(format!("encryption key `{id}` is unavailable")))
57    }
58}
59
60/// Encrypt an envelope in place. The poison-pill fingerprint is derived from plaintext
61/// before the randomized nonce is introduced, so identical bad jobs still quarantine
62/// together without revealing their payload to the store.
63pub fn encrypt_envelope(
64    keys: &dyn KeyProvider,
65    mut env: Envelope,
66) -> Result<Envelope, CryptoError> {
67    if env.fingerprint.is_empty() {
68        env.fingerprint = headgate::fingerprint(&env.kind, &env.payload);
69    }
70    let (key_id, key) = keys.active_key()?;
71    if key_id.is_empty() || key_id.len() > u16::MAX as usize {
72        return Err(CryptoError(
73            "encryption key id must be 1..65535 bytes".into(),
74        ));
75    }
76    let mut nonce = [0u8; NONCE_LEN];
77    rand::rng().fill_bytes(&mut nonce);
78    env.payload = seal(&key_id, &key, &nonce, &aad(&env), &env.payload)?;
79    Ok(env)
80}
81
82pub fn decrypt_envelope(keys: &dyn KeyProvider, env: &Envelope) -> Result<Vec<u8>, CryptoError> {
83    if env.payload.len() < MAGIC.len() + 2 + NONCE_LEN + 16 || &env.payload[..MAGIC.len()] != MAGIC
84    {
85        return Err(CryptoError(
86            "payload is not a headgate encrypted envelope".into(),
87        ));
88    }
89    let id_len = u16::from_be_bytes([env.payload[5], env.payload[6]]) as usize;
90    let id_start = 7;
91    let nonce_start = id_start + id_len;
92    if nonce_start + NONCE_LEN + 16 > env.payload.len() {
93        return Err(CryptoError("encrypted payload header is truncated".into()));
94    }
95    let key_id = std::str::from_utf8(&env.payload[id_start..nonce_start])
96        .map_err(|_| CryptoError("encryption key id is not UTF-8".into()))?;
97    let key = keys.key(key_id)?;
98    let cipher =
99        Aes256Gcm::new_from_slice(&key).map_err(|_| CryptoError("invalid AES key".into()))?;
100    cipher
101        .decrypt(
102            Nonce::from_slice(&env.payload[nonce_start..nonce_start + NONCE_LEN]),
103            Payload {
104                msg: &env.payload[nonce_start + NONCE_LEN..],
105                aad: &aad(env),
106            },
107        )
108        .map_err(|_| CryptoError("encrypted payload authentication failed".into()))
109}
110
111fn seal(
112    key_id: &str,
113    key: &[u8; 32],
114    nonce: &[u8; NONCE_LEN],
115    aad: &[u8],
116    plaintext: &[u8],
117) -> Result<Vec<u8>, CryptoError> {
118    let cipher =
119        Aes256Gcm::new_from_slice(key).map_err(|_| CryptoError("invalid AES key".into()))?;
120    let ciphertext = cipher
121        .encrypt(
122            Nonce::from_slice(nonce),
123            Payload {
124                msg: plaintext,
125                aad,
126            },
127        )
128        .map_err(|_| CryptoError("payload encryption failed".into()))?;
129    let mut out = Vec::with_capacity(MAGIC.len() + 2 + key_id.len() + NONCE_LEN + ciphertext.len());
130    out.extend_from_slice(MAGIC);
131    out.extend_from_slice(&(key_id.len() as u16).to_be_bytes());
132    out.extend_from_slice(key_id.as_bytes());
133    out.extend_from_slice(nonce);
134    out.extend_from_slice(&ciphertext);
135    Ok(out)
136}
137
138fn aad(env: &Envelope) -> Vec<u8> {
139    let mut out = Vec::with_capacity(env.id.len() + env.kind.len() + 12);
140    out.extend_from_slice(&(env.id.len() as u32).to_be_bytes());
141    out.extend_from_slice(env.id.as_bytes());
142    out.extend_from_slice(&(env.kind.len() as u32).to_be_bytes());
143    out.extend_from_slice(env.kind.as_bytes());
144    out.extend_from_slice(&env.schema_version.to_be_bytes());
145    out
146}
147
148pub fn register_encrypted<T, F, Fut>(
149    registry: &mut Registry,
150    keys: Arc<dyn KeyProvider>,
151    handler: F,
152) -> Result<(), String>
153where
154    T: Task,
155    F: Fn(JobCtx, T) -> Fut + Send + Sync + 'static,
156    Fut: std::future::Future<Output = Result<(), JobError>> + Send + 'static,
157{
158    registry.register_raw::<T, _, _>(move |ctx, mut env| {
159        let keys = keys.clone();
160        let plaintext =
161            decrypt_envelope(keys.as_ref(), &env).map_err(|e| CodecError::Malformed(e.to_string()));
162        let task = plaintext.and_then(|bytes| T::upcast(env.schema_version, &bytes));
163        env.payload.clear();
164        let future = task.map(|task| handler(ctx, task));
165        async move {
166            match future {
167                Ok(future) => future.await,
168                Err(error) => Err(Box::new(error) as JobError),
169            }
170        }
171    })
172}
173
174#[cfg(test)]
175mod tests {
176    use super::*;
177
178    fn ring() -> StaticKeyring {
179        StaticKeyring::new("k1", BTreeMap::from([("k1".into(), [7u8; 32])])).unwrap()
180    }
181
182    fn env() -> Envelope {
183        Envelope {
184            id: "job-1".into(),
185            kind: "secret:task".into(),
186            schema_version: 1,
187            payload: b"top secret".to_vec(),
188            queue: "default".into(),
189            ..Default::default()
190        }
191    }
192
193    #[test]
194    fn round_trip_binds_identity_and_preserves_plaintext_fingerprint() {
195        let a = encrypt_envelope(&ring(), env()).unwrap();
196        let b = encrypt_envelope(&ring(), env()).unwrap();
197        assert_ne!(a.payload, b.payload, "fresh nonce per enqueue");
198        assert_eq!(
199            a.fingerprint, b.fingerprint,
200            "fingerprint must not depend on nonce"
201        );
202        assert_eq!(decrypt_envelope(&ring(), &a).unwrap(), b"top secret");
203        let mut moved = a;
204        moved.id = "job-2".into();
205        assert!(
206            decrypt_envelope(&ring(), &moved).is_err(),
207            "job id is authenticated AAD"
208        );
209    }
210
211    #[test]
212    fn tampering_and_missing_keys_fail_authentication() {
213        let mut encrypted = encrypt_envelope(&ring(), env()).unwrap();
214        *encrypted.payload.last_mut().unwrap() ^= 1;
215        assert!(decrypt_envelope(&ring(), &encrypted).is_err());
216        let missing =
217            StaticKeyring::new("other", BTreeMap::from([("other".into(), [9u8; 32])])).unwrap();
218        assert!(decrypt_envelope(&missing, &encrypt_envelope(&ring(), env()).unwrap()).is_err());
219    }
220
221    #[test]
222    fn wire_vector_matches_go_byte_for_byte() {
223        let nonce = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11];
224        let env = env();
225        let wire = seal("k1", &[7u8; 32], &nonce, &aad(&env), b"top secret").unwrap();
226        let hex: String = wire.iter().map(|b| format!("{b:02x}")).collect();
227        assert_eq!(
228            hex,
229            "484745430100026b31000102030405060708090a0b6cee99506e6cba3b12c6527e0b794110389ff91129360bd1446d"
230        );
231    }
232}