1use 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
60pub 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}