use core::fmt;
use aes_gcm::aead::Aead;
use aes_gcm::{Aes256Gcm, KeyInit, Nonce};
use crate::crypto::fill_random;
use crate::crypto::zeroize::Zeroizing;
use crate::encoding::{Base64DecodeError, base64url_decode, base64url_encode};
const NONCE_LEN: usize = 12;
const PREFIX: &str = "enc:";
const PREFIX_V2: &str = "enc2:";
pub struct SecretBox {
cipher: Aes256Gcm,
}
impl SecretBox {
#[must_use]
pub fn from_key(key: &[u8; 32]) -> Self {
Self {
cipher: Aes256Gcm::new(key.into()),
}
}
#[must_use]
pub fn is_encrypted(encoded: &str) -> bool {
encoded.starts_with(PREFIX) || encoded.starts_with(PREFIX_V2)
}
pub fn encrypt(&self, plaintext: &[u8]) -> Result<String, SecretBoxError> {
self.seal(plaintext, b"", PREFIX)
}
pub fn encrypt_with_context(
&self,
plaintext: &[u8],
context: &[u8],
) -> Result<String, SecretBoxError> {
self.seal(plaintext, context, PREFIX_V2)
}
fn seal(&self, plaintext: &[u8], aad: &[u8], prefix: &str) -> Result<String, SecretBoxError> {
let mut nonce_bytes = [0u8; NONCE_LEN];
fill_random(&mut nonce_bytes)
.map_err(|_| SecretBoxError::new(SecretBoxErrorKind::Random))?;
let nonce = Nonce::from_slice(&nonce_bytes);
let ciphertext = self
.cipher
.encrypt(
nonce,
aes_gcm::aead::Payload {
msg: plaintext,
aad,
},
)
.map_err(|_| SecretBoxError::new(SecretBoxErrorKind::Aead))?;
let nonce_b64 = base64url_encode(&nonce_bytes);
let ct_b64 = base64url_encode(&ciphertext);
Ok(format!("{prefix}{nonce_b64}:{ct_b64}"))
}
pub fn decrypt(&self, encoded: &str) -> Result<Zeroizing<Vec<u8>>, SecretBoxError> {
self.decrypt_with_context(encoded, b"")
}
pub fn decrypt_with_context(
&self,
encoded: &str,
context: &[u8],
) -> Result<Zeroizing<Vec<u8>>, SecretBoxError> {
let aad: &[u8] = if encoded.starts_with(PREFIX_V2) {
context
} else if encoded.starts_with(PREFIX) {
b""
} else {
return Err(SecretBoxError::new(SecretBoxErrorKind::InvalidFormat));
};
let parts: Vec<&str> = encoded.split(':').collect();
if parts.len() != 3 {
return Err(SecretBoxError::new(SecretBoxErrorKind::InvalidFormat));
}
let nonce_bytes = base64url_decode(parts[1]).map_err(SecretBoxError::base64)?;
let ciphertext = base64url_decode(parts[2]).map_err(SecretBoxError::base64)?;
if nonce_bytes.len() != NONCE_LEN {
return Err(SecretBoxError::new(SecretBoxErrorKind::BadNonce));
}
let nonce = Nonce::from_slice(&nonce_bytes);
self.cipher
.decrypt(
nonce,
aes_gcm::aead::Payload {
msg: ciphertext.as_ref(),
aad,
},
)
.map(Zeroizing::new)
.map_err(|_| SecretBoxError::new(SecretBoxErrorKind::Aead))
}
}
impl fmt::Debug for SecretBox {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SecretBox")
.field("cipher", &"[REDACTED]")
.finish()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum SecretBoxErrorKind {
InvalidFormat,
Base64(Base64DecodeError),
Aead,
BadNonce,
Random,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SecretBoxError {
kind: SecretBoxErrorKind,
}
impl SecretBoxError {
const fn new(kind: SecretBoxErrorKind) -> Self {
Self { kind }
}
fn base64(err: Base64DecodeError) -> Self {
Self {
kind: SecretBoxErrorKind::Base64(err),
}
}
}
impl fmt::Display for SecretBoxError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.kind {
SecretBoxErrorKind::InvalidFormat => {
write!(f, "secret_box: invalid encrypted format")
}
SecretBoxErrorKind::Base64(err) => {
write!(f, "secret_box: base64 decode failed: {err}")
}
SecretBoxErrorKind::Aead => {
write!(f, "secret_box: aead error")
}
SecretBoxErrorKind::BadNonce => {
write!(f, "secret_box: nonce wrong length")
}
SecretBoxErrorKind::Random => {
write!(f, "secret_box: csprng unavailable")
}
}
}
}
impl std::error::Error for SecretBoxError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match &self.kind {
SecretBoxErrorKind::Base64(err) => Some(err),
_ => None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const TEST_KEY: [u8; 32] = [0u8; 32];
const OTHER_KEY: [u8; 32] = [1u8; 32];
fn sb() -> SecretBox {
SecretBox::from_key(&TEST_KEY)
}
#[test]
fn round_trip_recovers_plaintext() {
let sb = sb();
let encoded = sb.encrypt(b"hello world").unwrap();
let plain = sb.decrypt(&encoded).unwrap();
assert_eq!(plain.as_slice(), b"hello world");
}
#[test]
fn round_trip_empty_plaintext() {
let sb = sb();
let encoded = sb.encrypt(b"").unwrap();
let plain = sb.decrypt(&encoded).unwrap();
assert_eq!(plain.as_slice(), b"");
}
#[test]
fn round_trip_long_plaintext() {
let sb = sb();
let plaintext = vec![0xABu8; 4096];
let encoded = sb.encrypt(&plaintext).unwrap();
let plain = sb.decrypt(&encoded).unwrap();
assert_eq!(plain.as_slice(), plaintext.as_slice());
}
#[test]
fn nonce_randomness_yields_distinct_ciphertexts() {
let sb = sb();
let a = sb.encrypt(b"same plaintext").unwrap();
let b = sb.encrypt(b"same plaintext").unwrap();
assert_ne!(a, b, "fresh nonce per encryption should differ");
}
#[test]
fn encrypted_format_starts_with_prefix() {
let sb = sb();
let encoded = sb.encrypt(b"x").unwrap();
assert!(encoded.starts_with("enc:"));
let parts: Vec<&str> = encoded.split(':').collect();
assert_eq!(parts.len(), 3);
assert_eq!(parts[0], "enc");
}
#[test]
fn is_encrypted_recognises_prefix() {
assert!(SecretBox::is_encrypted("enc:foo:bar"));
assert!(SecretBox::is_encrypted("enc:"));
assert!(!SecretBox::is_encrypted("plaintext"));
assert!(!SecretBox::is_encrypted(""));
assert!(!SecretBox::is_encrypted("ENC:foo:bar"));
}
#[test]
fn decrypt_rejects_missing_prefix() {
let sb = sb();
assert_eq!(
sb.decrypt("aGVsbG8:d29ybGQ").unwrap_err(),
SecretBoxError::new(SecretBoxErrorKind::InvalidFormat),
);
}
#[test]
fn decrypt_rejects_too_few_fields() {
let sb = sb();
assert_eq!(
sb.decrypt("enc:onlyone").unwrap_err(),
SecretBoxError::new(SecretBoxErrorKind::InvalidFormat),
);
}
#[test]
fn decrypt_rejects_too_many_fields() {
let sb = sb();
assert_eq!(
sb.decrypt("enc:a:b:c").unwrap_err(),
SecretBoxError::new(SecretBoxErrorKind::InvalidFormat),
);
}
#[test]
fn decrypt_rejects_bad_base64_nonce() {
let sb = sb();
let err = sb.decrypt("enc:!!!:AAAA").unwrap_err();
assert!(
matches!(err.kind, SecretBoxErrorKind::Base64(_)),
"expected Base64 variant, got {err:?}"
);
}
#[test]
fn decrypt_rejects_bad_base64_ciphertext() {
let sb = sb();
let nonce = base64url_encode(&[0u8; NONCE_LEN]);
let err = sb.decrypt(&format!("enc:{nonce}:!!!")).unwrap_err();
assert!(
matches!(err.kind, SecretBoxErrorKind::Base64(_)),
"expected Base64 variant, got {err:?}"
);
}
#[test]
fn decrypt_rejects_wrong_nonce_length() {
let sb = sb();
let bad_nonce = base64url_encode(&[0u8; 8]);
let ct = base64url_encode(&[0u8; 16]);
assert_eq!(
sb.decrypt(&format!("enc:{bad_nonce}:{ct}")).unwrap_err(),
SecretBoxError::new(SecretBoxErrorKind::BadNonce),
);
}
#[test]
fn decrypt_with_wrong_key_fails_aead() {
let alice = SecretBox::from_key(&TEST_KEY);
let bob = SecretBox::from_key(&OTHER_KEY);
let encoded = alice.encrypt(b"top secret").unwrap();
assert_eq!(
bob.decrypt(&encoded).unwrap_err(),
SecretBoxError::new(SecretBoxErrorKind::Aead),
);
}
#[test]
fn decrypt_rejects_tampered_ciphertext() {
let sb = sb();
let encoded = sb.encrypt(b"protected").unwrap();
let parts: Vec<&str> = encoded.split(':').collect();
let mut ct = base64url_decode(parts[2]).unwrap();
ct[0] ^= 0x01;
let tampered = format!("enc:{}:{}", parts[1], base64url_encode(&ct));
assert_eq!(
sb.decrypt(&tampered).unwrap_err(),
SecretBoxError::new(SecretBoxErrorKind::Aead),
);
}
#[test]
fn error_display_contains_no_secrets() {
let errors = [
SecretBoxError::new(SecretBoxErrorKind::InvalidFormat),
SecretBoxError::new(SecretBoxErrorKind::Aead),
SecretBoxError::new(SecretBoxErrorKind::BadNonce),
];
for err in &errors {
let msg = err.to_string();
assert!(
msg.starts_with("secret_box:"),
"error should be prefixed: {msg}"
);
assert!(
!msg.contains("key") && !msg.contains("plaintext"),
"error must not leak material: {msg}",
);
}
}
#[test]
fn error_implements_std_error() {
let err: Box<dyn std::error::Error> =
Box::new(SecretBoxError::new(SecretBoxErrorKind::Aead));
let _ = err.to_string();
}
#[test]
fn debug_redacts_cipher() {
let sb = sb();
let dbg = format!("{sb:?}");
assert!(dbg.contains("[REDACTED]"), "debug should redact: {dbg}");
}
#[test]
fn context_bound_blob_only_decrypts_under_its_context() {
let sb = SecretBox::from_key(&[7u8; 32]);
let ctx = b"totp_secret:tenant-a:user-1";
let blob = sb.encrypt_with_context(b"seed", ctx).unwrap();
assert!(blob.starts_with("enc2:"));
assert!(SecretBox::is_encrypted(&blob));
assert_eq!(
sb.decrypt_with_context(&blob, ctx).unwrap().as_slice(),
b"seed"
);
assert!(
sb.decrypt_with_context(&blob, b"totp_secret:tenant-b:user-1")
.is_err()
);
assert!(sb.decrypt_with_context(&blob, b"").is_err());
assert!(sb.decrypt(&blob).is_err());
}
#[test]
fn legacy_blob_still_decrypts_under_any_context() {
let sb = SecretBox::from_key(&[9u8; 32]);
let legacy = sb.encrypt(b"old").unwrap();
assert!(legacy.starts_with("enc:"));
assert_eq!(
sb.decrypt_with_context(&legacy, b"whatever:1:2")
.unwrap()
.as_slice(),
b"old"
);
}
}