use std::{fmt, sync::Arc};
use aes_gcm::Aes256Gcm;
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
use chacha20poly1305::{
ChaCha20Poly1305, Nonce,
aead::{Aead, AeadCore, KeyInit, OsRng, Payload},
};
use rand::RngCore as _;
use serde::{Deserialize, Serialize};
use zeroize::{Zeroize as _, Zeroizing};
use crate::{AuthError, error::Result};
#[derive(Debug, Clone)]
pub struct EncryptedState {
pub ciphertext: Vec<u8>,
pub nonce: [u8; 12],
}
impl EncryptedState {
#[must_use]
pub const fn new(ciphertext: Vec<u8>, nonce: [u8; 12]) -> Self {
Self { ciphertext, nonce }
}
#[must_use]
pub fn to_bytes(&self) -> Vec<u8> {
let mut bytes = Vec::with_capacity(12 + self.ciphertext.len());
bytes.extend_from_slice(&self.nonce);
bytes.extend_from_slice(&self.ciphertext);
bytes
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
if bytes.len() < 12 {
return Err(AuthError::InvalidState);
}
let mut nonce = [0u8; 12];
nonce.copy_from_slice(&bytes[0..12]);
let ciphertext = bytes[12..].to_vec();
Ok(Self::new(ciphertext, nonce))
}
}
pub struct StateEncryption {
cipher: ChaCha20Poly1305,
}
impl StateEncryption {
pub fn new(key_bytes: &[u8; 32]) -> Result<Self> {
let cipher =
ChaCha20Poly1305::new_from_slice(key_bytes).map_err(|_| AuthError::ConfigError {
message: "Invalid state encryption key".to_string(),
})?;
Ok(Self { cipher })
}
pub fn encrypt(&self, state: &str) -> Result<EncryptedState> {
let mut nonce_bytes = [0u8; 12];
rand::rng().fill_bytes(&mut nonce_bytes);
let nonce = Nonce::from(nonce_bytes);
let ciphertext =
self.cipher.encrypt(&nonce, Payload::from(state.as_bytes())).map_err(|_| {
AuthError::Internal {
message: "State encryption failed".to_string(),
}
})?;
Ok(EncryptedState::new(ciphertext, nonce_bytes))
}
pub fn decrypt(&self, encrypted: &EncryptedState) -> Result<String> {
let nonce = Nonce::from(encrypted.nonce);
let plaintext = self
.cipher
.decrypt(&nonce, Payload::from(encrypted.ciphertext.as_slice()))
.map_err(|_| AuthError::InvalidState)?;
String::from_utf8(plaintext).map_err(|_| AuthError::InvalidState)
}
pub fn encrypt_to_bytes(&self, state: &str) -> Result<Vec<u8>> {
let encrypted = self.encrypt(state)?;
Ok(encrypted.to_bytes())
}
pub fn decrypt_from_bytes(&self, bytes: &[u8]) -> Result<String> {
let encrypted = EncryptedState::from_bytes(bytes)?;
self.decrypt(&encrypted)
}
}
#[must_use]
pub fn generate_state_encryption_key() -> Zeroizing<[u8; 32]> {
let mut key = [0u8; 32];
rand::rng().fill_bytes(&mut key);
Zeroizing::new(key)
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum DecryptionError {
#[error("authentication failed — ciphertext may be tampered or key is wrong")]
AuthenticationFailed,
#[error("invalid input: {0}")]
InvalidInput(String),
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum KeyError {
#[error("hex key must be 64 chars (32 bytes); got {0} chars")]
WrongLength(usize),
#[error("invalid hex character in key")]
InvalidHex,
}
#[derive(Debug, Clone, Default, Deserialize, Serialize, PartialEq, Eq)]
#[non_exhaustive]
pub enum EncryptionAlgorithm {
#[default]
#[serde(rename = "chacha20-poly1305")]
Chacha20Poly1305,
#[serde(rename = "aes-256-gcm")]
Aes256Gcm,
}
impl fmt::Display for EncryptionAlgorithm {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Chacha20Poly1305 => f.write_str("chacha20-poly1305"),
Self::Aes256Gcm => f.write_str("aes-256-gcm"),
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
pub struct StateEncryptionConfig {
pub enabled: bool,
pub algorithm: EncryptionAlgorithm,
pub key_env: Option<String>,
}
impl Default for StateEncryptionConfig {
fn default() -> Self {
Self {
enabled: false,
algorithm: EncryptionAlgorithm::default(),
key_env: Some("STATE_ENCRYPTION_KEY".to_string()),
}
}
}
pub struct StateEncryptionService {
algorithm: EncryptionAlgorithm,
key: [u8; 32],
}
impl fmt::Debug for StateEncryptionService {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StateEncryptionService")
.field("algorithm", &self.algorithm)
.field("key", &"[REDACTED]")
.finish()
}
}
impl Drop for StateEncryptionService {
fn drop(&mut self) {
self.key.zeroize();
}
}
impl StateEncryptionService {
#[must_use]
pub const fn from_raw_key(key: &[u8; 32], algorithm: EncryptionAlgorithm) -> Self {
Self {
algorithm,
key: *key,
}
}
pub fn from_hex_key(
hex: &str,
algorithm: EncryptionAlgorithm,
) -> std::result::Result<Self, KeyError> {
if hex.len() != 64 {
return Err(KeyError::WrongLength(hex.len()));
}
let bytes = hex::decode(hex).map_err(|_| KeyError::InvalidHex)?;
let mut key = [0u8; 32];
key.copy_from_slice(&bytes);
Ok(Self { algorithm, key })
}
pub fn new_from_env(
var: &str,
algorithm: EncryptionAlgorithm,
) -> std::result::Result<Self, anyhow::Error> {
let hex = std::env::var(var).map_err(|_| anyhow::anyhow!("env var {var} not set"))?;
Ok(Self::from_hex_key(&hex, algorithm)?)
}
pub fn from_compiled_schema(
security_json: &serde_json::Value,
) -> std::result::Result<Option<Arc<Self>>, anyhow::Error> {
let cfg: StateEncryptionConfig = match security_json.get("state_encryption") {
None | Some(serde_json::Value::Null) => return Ok(None),
Some(v) => serde_json::from_value(v.clone())
.map_err(|e| anyhow::anyhow!("invalid state_encryption config: {e}"))?,
};
if !cfg.enabled {
return Ok(None);
}
let key_env = cfg.key_env.as_deref().unwrap_or("STATE_ENCRYPTION_KEY");
Self::new_from_env(key_env, cfg.algorithm)
.map(|svc| Some(Arc::new(svc)))
.map_err(|e| {
anyhow::anyhow!(
"state_encryption enabled but key env var '{}' failed: {e}",
key_env
)
})
}
pub fn encrypt(&self, plaintext: &[u8]) -> std::result::Result<String, anyhow::Error> {
let combined = match self.algorithm {
EncryptionAlgorithm::Chacha20Poly1305 => {
let cipher = ChaCha20Poly1305::new_from_slice(&self.key)
.map_err(|_| anyhow::anyhow!("invalid key for ChaCha20-Poly1305"))?;
let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng);
let ct = cipher
.encrypt(&nonce, plaintext)
.map_err(|_| anyhow::anyhow!("ChaCha20-Poly1305 encryption failed"))?;
let mut out = nonce.to_vec();
out.extend_from_slice(&ct);
out
},
EncryptionAlgorithm::Aes256Gcm => {
let cipher = Aes256Gcm::new_from_slice(&self.key)
.map_err(|_| anyhow::anyhow!("invalid key for AES-256-GCM"))?;
let nonce = Aes256Gcm::generate_nonce(&mut OsRng);
let ct = cipher
.encrypt(&nonce, plaintext)
.map_err(|_| anyhow::anyhow!("AES-256-GCM encryption failed"))?;
let mut out = nonce.to_vec();
out.extend_from_slice(&ct);
out
},
};
Ok(URL_SAFE_NO_PAD.encode(&combined))
}
pub fn decrypt(&self, encoded: &str) -> std::result::Result<Vec<u8>, DecryptionError> {
const NONCE_SIZE: usize = 12;
if encoded.is_empty() {
return Err(DecryptionError::InvalidInput("empty input".into()));
}
let combined = URL_SAFE_NO_PAD
.decode(encoded)
.map_err(|_| DecryptionError::InvalidInput("invalid base64".into()))?;
if combined.len() < NONCE_SIZE {
return Err(DecryptionError::InvalidInput(format!(
"too short: {} bytes (minimum {NONCE_SIZE})",
combined.len()
)));
}
let (nonce_bytes, ct) = combined.split_at(NONCE_SIZE);
match self.algorithm {
EncryptionAlgorithm::Chacha20Poly1305 => {
let cipher = ChaCha20Poly1305::new_from_slice(&self.key)
.map_err(|_| DecryptionError::InvalidInput("invalid key".into()))?;
let nonce = chacha20poly1305::Nonce::from_slice(nonce_bytes);
cipher.decrypt(nonce, ct).map_err(|_| DecryptionError::AuthenticationFailed)
},
EncryptionAlgorithm::Aes256Gcm => {
let cipher = Aes256Gcm::new_from_slice(&self.key)
.map_err(|_| DecryptionError::InvalidInput("invalid key".into()))?;
let nonce = aes_gcm::Nonce::from_slice(nonce_bytes);
cipher.decrypt(nonce, ct).map_err(|_| DecryptionError::AuthenticationFailed)
},
}
}
}