use aes_gcm::aead::{Aead, KeyInit, OsRng};
use aes_gcm::{Aes256Gcm, Nonce};
use base64::Engine;
use rand::RngCore;
use crate::error::StorageError;
pub struct Encryptor {
cipher: Aes256Gcm,
}
impl Encryptor {
pub fn from_base64_key(key_base64: &str) -> Result<Self, StorageError> {
let key_bytes = base64::engine::general_purpose::STANDARD
.decode(key_base64)
.map_err(|e| StorageError::ConfigError(format!("Invalid base64 key: {}", e)))?;
if key_bytes.len() != 32 {
return Err(StorageError::ConfigError(format!(
"Key must be 32 bytes (256 bits), got {} bytes",
key_bytes.len()
)));
}
let key = aes_gcm::Key::<Aes256Gcm>::from_slice(&key_bytes);
let cipher = Aes256Gcm::new(key);
Ok(Encryptor { cipher })
}
pub fn generate_key() -> String {
let mut key_bytes = [0u8; 32];
OsRng.fill_bytes(&mut key_bytes);
base64::engine::general_purpose::STANDARD.encode(&key_bytes)
}
pub fn encrypt(&self, plaintext: &str) -> Result<String, StorageError> {
let mut nonce_bytes = [0u8; 12];
OsRng.fill_bytes(&mut nonce_bytes);
let nonce = Nonce::from_slice(&nonce_bytes);
let ciphertext = self
.cipher
.encrypt(nonce, plaintext.as_bytes())
.map_err(|e| StorageError::ConfigError(format!("Encryption failed: {}", e)))?;
let nonce_b64 = base64::engine::general_purpose::STANDARD.encode(&nonce_bytes);
let ciphertext_b64 = base64::engine::general_purpose::STANDARD.encode(&ciphertext);
Ok(format!("{}.{}", nonce_b64, ciphertext_b64))
}
pub fn decrypt(&self, encrypted: &str) -> Result<String, StorageError> {
let parts: Vec<&str> = encrypted.splitn(2, '.').collect();
if parts.len() != 2 {
return Err(StorageError::ConfigError(
"Invalid encrypted format".to_string(),
));
}
let nonce_bytes = base64::engine::general_purpose::STANDARD
.decode(parts[0])
.map_err(|e| StorageError::ConfigError(format!("Invalid nonce: {}", e)))?;
let ciphertext = base64::engine::general_purpose::STANDARD
.decode(parts[1])
.map_err(|e| StorageError::ConfigError(format!("Invalid ciphertext: {}", e)))?;
let nonce = Nonce::from_slice(&nonce_bytes);
let plaintext = self
.cipher
.decrypt(nonce, ciphertext.as_ref())
.map_err(|e| StorageError::ConfigError(format!("Decryption failed: {}", e)))?;
String::from_utf8(plaintext).map_err(|e| {
StorageError::ConfigError(format!("Invalid UTF-8 after decryption: {}", e))
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generate_key() {
let key = Encryptor::generate_key();
assert_eq!(key.len(), 44);
}
#[test]
fn test_encrypt_decrypt() {
let key = Encryptor::generate_key();
let encryptor = Encryptor::from_base64_key(&key).unwrap();
let plaintext = "Hello, 杰哥! This is secret data.";
let encrypted = encryptor.encrypt(plaintext).unwrap();
assert_ne!(encrypted, plaintext);
let decrypted = encryptor.decrypt(&encrypted).unwrap();
assert_eq!(decrypted, plaintext);
}
#[test]
fn test_encrypt_produces_different_output() {
let key = Encryptor::generate_key();
let encryptor = Encryptor::from_base64_key(&key).unwrap();
let plaintext = "test";
let enc1 = encryptor.encrypt(plaintext).unwrap();
let enc2 = encryptor.encrypt(plaintext).unwrap();
assert_ne!(enc1, enc2);
}
#[test]
fn test_invalid_key_length() {
let result = Encryptor::from_base64_key("dG9vLXNob3J0");
assert!(result.is_err());
}
}