1use std::{fmt, sync::Arc};
12
13use aes_gcm::Aes256Gcm;
16use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
17use chacha20poly1305::{
18 ChaCha20Poly1305, Nonce,
19 aead::{Aead, AeadCore, KeyInit, OsRng, Payload},
20};
21use rand::RngCore as _;
22use serde::{Deserialize, Serialize};
23use zeroize::Zeroizing;
24
25use crate::{AuthError, error::Result};
26
27#[derive(Debug, Clone)]
29pub struct EncryptedState {
30 pub ciphertext: Vec<u8>,
32 pub nonce: [u8; 12],
34}
35
36impl EncryptedState {
37 #[must_use]
39 pub const fn new(ciphertext: Vec<u8>, nonce: [u8; 12]) -> Self {
40 Self { ciphertext, nonce }
41 }
42
43 #[must_use]
46 pub fn to_bytes(&self) -> Vec<u8> {
47 let mut bytes = Vec::with_capacity(12 + self.ciphertext.len());
48 bytes.extend_from_slice(&self.nonce);
49 bytes.extend_from_slice(&self.ciphertext);
50 bytes
51 }
52
53 pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
60 if bytes.len() < 12 {
61 return Err(AuthError::InvalidState);
62 }
63
64 let mut nonce = [0u8; 12];
65 nonce.copy_from_slice(&bytes[0..12]);
66 let ciphertext = bytes[12..].to_vec();
67
68 Ok(Self::new(ciphertext, nonce))
69 }
70}
71
72pub struct StateEncryption {
84 cipher: ChaCha20Poly1305,
85}
86
87impl StateEncryption {
88 pub fn new(key_bytes: &[u8; 32]) -> Result<Self> {
96 let cipher =
97 ChaCha20Poly1305::new_from_slice(key_bytes).map_err(|_| AuthError::ConfigError {
98 message: "Invalid state encryption key".to_string(),
99 })?;
100
101 Ok(Self { cipher })
102 }
103
104 pub fn encrypt(&self, state: &str) -> Result<EncryptedState> {
118 let mut nonce_bytes = [0u8; 12];
120 rand::rng().fill_bytes(&mut nonce_bytes);
121 let nonce = Nonce::from(nonce_bytes);
122
123 let ciphertext =
125 self.cipher.encrypt(&nonce, Payload::from(state.as_bytes())).map_err(|_| {
126 AuthError::Internal {
127 message: "State encryption failed".to_string(),
128 }
129 })?;
130
131 Ok(EncryptedState::new(ciphertext, nonce_bytes))
132 }
133
134 pub fn decrypt(&self, encrypted: &EncryptedState) -> Result<String> {
151 let nonce = Nonce::from(encrypted.nonce);
152
153 let plaintext = self
155 .cipher
156 .decrypt(&nonce, Payload::from(encrypted.ciphertext.as_slice()))
157 .map_err(|_| AuthError::InvalidState)?;
158
159 String::from_utf8(plaintext).map_err(|_| AuthError::InvalidState)
161 }
162
163 pub fn encrypt_to_bytes(&self, state: &str) -> Result<Vec<u8>> {
169 let encrypted = self.encrypt(state)?;
170 Ok(encrypted.to_bytes())
171 }
172
173 pub fn decrypt_from_bytes(&self, bytes: &[u8]) -> Result<String> {
181 let encrypted = EncryptedState::from_bytes(bytes)?;
182 self.decrypt(&encrypted)
183 }
184}
185
186#[must_use]
188pub fn generate_state_encryption_key() -> Zeroizing<[u8; 32]> {
189 let mut key = [0u8; 32];
190 rand::rng().fill_bytes(&mut key);
191 Zeroizing::new(key)
192}
193
194#[derive(Debug, thiserror::Error)]
208#[non_exhaustive]
209pub enum DecryptionError {
210 #[error("authentication failed — ciphertext may be tampered or key is wrong")]
212 AuthenticationFailed,
213 #[error("invalid input: {0}")]
215 InvalidInput(String),
216}
217
218#[derive(Debug, thiserror::Error)]
220#[non_exhaustive]
221pub enum KeyError {
222 #[error("hex key must be 64 chars (32 bytes); got {0} chars")]
224 WrongLength(usize),
225 #[error("invalid hex character in key")]
227 InvalidHex,
228}
229
230#[derive(Debug, Clone, Default, Deserialize, Serialize, PartialEq, Eq)]
232#[non_exhaustive]
233pub enum EncryptionAlgorithm {
234 #[default]
236 #[serde(rename = "chacha20-poly1305")]
237 Chacha20Poly1305,
238 #[serde(rename = "aes-256-gcm")]
240 Aes256Gcm,
241}
242
243impl fmt::Display for EncryptionAlgorithm {
244 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
245 match self {
246 Self::Chacha20Poly1305 => f.write_str("chacha20-poly1305"),
247 Self::Aes256Gcm => f.write_str("aes-256-gcm"),
248 }
249 }
250}
251
252#[derive(Debug, Clone, Deserialize, Serialize)]
254#[serde(default)]
255pub struct StateEncryptionConfig {
256 pub enabled: bool,
258 pub algorithm: EncryptionAlgorithm,
260 pub key_env: Option<String>,
262}
263
264impl Default for StateEncryptionConfig {
265 fn default() -> Self {
266 Self {
267 enabled: false,
268 algorithm: EncryptionAlgorithm::default(),
269 key_env: Some("STATE_ENCRYPTION_KEY".to_string()),
270 }
271 }
272}
273
274pub struct StateEncryptionService {
280 algorithm: EncryptionAlgorithm,
281 key: [u8; 32],
282}
283
284impl fmt::Debug for StateEncryptionService {
285 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
286 f.debug_struct("StateEncryptionService")
287 .field("algorithm", &self.algorithm)
288 .field("key", &"[REDACTED]")
289 .finish()
290 }
291}
292
293impl StateEncryptionService {
294 #[must_use]
296 pub const fn from_raw_key(key: &[u8; 32], algorithm: EncryptionAlgorithm) -> Self {
297 Self {
298 algorithm,
299 key: *key,
300 }
301 }
302
303 pub fn from_hex_key(
310 hex: &str,
311 algorithm: EncryptionAlgorithm,
312 ) -> std::result::Result<Self, KeyError> {
313 if hex.len() != 64 {
314 return Err(KeyError::WrongLength(hex.len()));
315 }
316 let bytes = hex::decode(hex).map_err(|_| KeyError::InvalidHex)?;
317 let mut key = [0u8; 32];
318 key.copy_from_slice(&bytes);
319 Ok(Self { algorithm, key })
320 }
321
322 pub fn new_from_env(
328 var: &str,
329 algorithm: EncryptionAlgorithm,
330 ) -> std::result::Result<Self, anyhow::Error> {
331 let hex = std::env::var(var).map_err(|_| anyhow::anyhow!("env var {var} not set"))?;
332 Ok(Self::from_hex_key(&hex, algorithm)?)
333 }
334
335 pub fn from_compiled_schema(
344 security_json: &serde_json::Value,
345 ) -> std::result::Result<Option<Arc<Self>>, anyhow::Error> {
346 let cfg: StateEncryptionConfig = match security_json.get("state_encryption") {
347 None | Some(serde_json::Value::Null) => return Ok(None),
348 Some(v) => serde_json::from_value(v.clone())
349 .map_err(|e| anyhow::anyhow!("invalid state_encryption config: {e}"))?,
350 };
351
352 if !cfg.enabled {
353 return Ok(None);
354 }
355
356 let key_env = cfg.key_env.as_deref().unwrap_or("STATE_ENCRYPTION_KEY");
357 Self::new_from_env(key_env, cfg.algorithm)
358 .map(|svc| Some(Arc::new(svc)))
359 .map_err(|e| {
360 anyhow::anyhow!(
361 "state_encryption enabled but key env var '{}' failed: {e}",
362 key_env
363 )
364 })
365 }
366
367 pub fn encrypt(&self, plaintext: &[u8]) -> std::result::Result<String, anyhow::Error> {
375 let combined = match self.algorithm {
376 EncryptionAlgorithm::Chacha20Poly1305 => {
377 let cipher = ChaCha20Poly1305::new_from_slice(&self.key)
378 .map_err(|_| anyhow::anyhow!("invalid key for ChaCha20-Poly1305"))?;
379 let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng);
380 let ct = cipher
381 .encrypt(&nonce, plaintext)
382 .map_err(|_| anyhow::anyhow!("ChaCha20-Poly1305 encryption failed"))?;
383 let mut out = nonce.to_vec();
384 out.extend_from_slice(&ct);
385 out
386 },
387 EncryptionAlgorithm::Aes256Gcm => {
388 let cipher = Aes256Gcm::new_from_slice(&self.key)
389 .map_err(|_| anyhow::anyhow!("invalid key for AES-256-GCM"))?;
390 let nonce = Aes256Gcm::generate_nonce(&mut OsRng);
391 let ct = cipher
392 .encrypt(&nonce, plaintext)
393 .map_err(|_| anyhow::anyhow!("AES-256-GCM encryption failed"))?;
394 let mut out = nonce.to_vec();
395 out.extend_from_slice(&ct);
396 out
397 },
398 };
399 Ok(URL_SAFE_NO_PAD.encode(&combined))
400 }
401
402 pub fn decrypt(&self, encoded: &str) -> std::result::Result<Vec<u8>, DecryptionError> {
409 const NONCE_SIZE: usize = 12;
410 if encoded.is_empty() {
411 return Err(DecryptionError::InvalidInput("empty input".into()));
412 }
413 let combined = URL_SAFE_NO_PAD
414 .decode(encoded)
415 .map_err(|_| DecryptionError::InvalidInput("invalid base64".into()))?;
416
417 if combined.len() < NONCE_SIZE {
418 return Err(DecryptionError::InvalidInput(format!(
419 "too short: {} bytes (minimum {NONCE_SIZE})",
420 combined.len()
421 )));
422 }
423 let (nonce_bytes, ct) = combined.split_at(NONCE_SIZE);
424
425 match self.algorithm {
426 EncryptionAlgorithm::Chacha20Poly1305 => {
427 let cipher = ChaCha20Poly1305::new_from_slice(&self.key)
428 .map_err(|_| DecryptionError::InvalidInput("invalid key".into()))?;
429 let nonce = chacha20poly1305::Nonce::from_slice(nonce_bytes);
430 cipher.decrypt(nonce, ct).map_err(|_| DecryptionError::AuthenticationFailed)
431 },
432 EncryptionAlgorithm::Aes256Gcm => {
433 let cipher = Aes256Gcm::new_from_slice(&self.key)
434 .map_err(|_| DecryptionError::InvalidInput("invalid key".into()))?;
435 let nonce = aes_gcm::Nonce::from_slice(nonce_bytes);
436 cipher.decrypt(nonce, ct).map_err(|_| DecryptionError::AuthenticationFailed)
437 },
438 }
439 }
440}