use zeroize::Zeroizing;
use crate::{AesGcmCrypter, Crypter, CryptoError};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum EncryptionAlgo {
#[default]
Aes256Gcm,
ChaCha20Poly1305,
Sm4Gcm,
}
impl EncryptionAlgo {
pub fn as_str(&self) -> &'static str {
match self {
Self::Aes256Gcm => "AES-256-GCM",
Self::ChaCha20Poly1305 => "ChaCha20-Poly1305",
Self::Sm4Gcm => "SM4-GCM",
}
}
pub fn is_supported(&self) -> bool {
matches!(self, Self::Aes256Gcm)
}
}
pub struct DekBuffer {
dek: Zeroizing<Vec<u8>>,
}
impl DekBuffer {
pub fn new(plaintext: Vec<u8>) -> Self {
Self {
dek: Zeroizing::new(plaintext),
}
}
pub fn from_slice(slice: &[u8]) -> Self {
Self::new(slice.to_vec())
}
pub fn as_bytes(&self) -> &[u8] {
&self.dek
}
pub fn len(&self) -> usize {
self.dek.len()
}
pub fn is_empty(&self) -> bool {
self.dek.is_empty()
}
pub fn encrypt(&self, plaintext: &[u8], algo: EncryptionAlgo) -> Result<Vec<u8>, CryptoError> {
match algo {
EncryptionAlgo::Aes256Gcm => {
if self.len() != 32 {
return Err(CryptoError::InvalidKey(format!(
"AES-256-GCM 需要 32 字节 DEK,实际 {} 字节",
self.len()
)));
}
let mut key = [0u8; 32];
key.copy_from_slice(&self.dek[..32]);
let crypter = AesGcmCrypter::new(&key);
let result = crypter.encrypt(plaintext);
use zeroize::Zeroize;
key.zeroize();
result
}
EncryptionAlgo::ChaCha20Poly1305 | EncryptionAlgo::Sm4Gcm => {
Err(CryptoError::EncryptionFailed(format!(
"算法 {} 未实现,请使用 AES-256-GCM",
algo.as_str()
)))
}
}
}
pub fn decrypt(&self, ciphertext: &[u8], algo: EncryptionAlgo) -> Result<Vec<u8>, CryptoError> {
match algo {
EncryptionAlgo::Aes256Gcm => {
if self.len() != 32 {
return Err(CryptoError::InvalidKey(format!(
"AES-256-GCM 需要 32 字节 DEK,实际 {} 字节",
self.len()
)));
}
let mut key = [0u8; 32];
key.copy_from_slice(&self.dek[..32]);
let crypter = AesGcmCrypter::new(&key);
let result = crypter.decrypt(ciphertext);
use zeroize::Zeroize;
key.zeroize();
result
}
EncryptionAlgo::ChaCha20Poly1305 | EncryptionAlgo::Sm4Gcm => {
Err(CryptoError::DecryptionFailed(format!(
"算法 {} 未实现,请使用 AES-256-GCM",
algo.as_str()
)))
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dek_buffer_encrypt_decrypt_aes256gcm() {
let dek = DekBuffer::new(vec![0x42u8; 32]);
let plaintext = b"sensitive data";
let ciphertext = dek.encrypt(plaintext, EncryptionAlgo::Aes256Gcm).unwrap();
assert_ne!(&ciphertext[..], &plaintext[..]);
let decrypted = dek.decrypt(&ciphertext, EncryptionAlgo::Aes256Gcm).unwrap();
assert_eq!(decrypted, plaintext);
}
#[test]
fn dek_buffer_empty_plaintext() {
let dek = DekBuffer::new(vec![0x42u8; 32]);
let ciphertext = dek.encrypt(b"", EncryptionAlgo::Aes256Gcm).unwrap();
let decrypted = dek.decrypt(&ciphertext, EncryptionAlgo::Aes256Gcm).unwrap();
assert_eq!(decrypted, b"");
}
#[test]
fn dek_buffer_wrong_key_fails() {
let dek1 = DekBuffer::new(vec![0x42u8; 32]);
let dek2 = DekBuffer::new(vec![0x43u8; 32]);
let ciphertext = dek1.encrypt(b"secret", EncryptionAlgo::Aes256Gcm).unwrap();
assert!(dek2
.decrypt(&ciphertext, EncryptionAlgo::Aes256Gcm)
.is_err());
}
#[test]
fn dek_buffer_invalid_key_length() {
let dek = DekBuffer::new(vec![0x42u8; 16]);
assert!(dek.encrypt(b"data", EncryptionAlgo::Aes256Gcm).is_err());
}
#[test]
fn dek_buffer_chacha20_unsupported() {
let dek = DekBuffer::new(vec![0x42u8; 32]);
assert!(dek
.encrypt(b"data", EncryptionAlgo::ChaCha20Poly1305)
.is_err());
}
#[test]
fn dek_buffer_sm4_unsupported() {
let dek = DekBuffer::new(vec![0x42u8; 32]);
assert!(dek.encrypt(b"data", EncryptionAlgo::Sm4Gcm).is_err());
}
#[test]
fn dek_buffer_drop_zeroizes() {
let dek = DekBuffer::new(vec![0xABu8; 32]);
let ptr = dek.as_bytes().as_ptr();
assert_eq!(unsafe { *ptr }, 0xAB);
drop(dek);
}
#[test]
fn encryption_algo_default() {
assert_eq!(EncryptionAlgo::default(), EncryptionAlgo::Aes256Gcm);
}
#[test]
fn encryption_algo_is_supported() {
assert!(EncryptionAlgo::Aes256Gcm.is_supported());
assert!(!EncryptionAlgo::ChaCha20Poly1305.is_supported());
assert!(!EncryptionAlgo::Sm4Gcm.is_supported());
}
}