use crate::IggyError;
use crate::text;
use aes_gcm::aead::{Aead, OsRng};
use aes_gcm::{AeadCore, Aes256Gcm, KeyInit};
use std::fmt::Debug;
#[derive(Debug, Clone)]
pub enum EncryptorKind {
Aes256Gcm(Aes256GcmEncryptor),
}
impl EncryptorKind {
pub fn encrypt(&self, data: &[u8]) -> Result<Vec<u8>, IggyError> {
match self {
EncryptorKind::Aes256Gcm(e) => e.encrypt(data),
}
}
pub fn decrypt(&self, data: &[u8]) -> Result<Vec<u8>, IggyError> {
match self {
EncryptorKind::Aes256Gcm(e) => e.decrypt(data),
}
}
}
pub trait Encryptor {
fn encrypt(&self, data: &[u8]) -> Result<Vec<u8>, IggyError>;
fn decrypt(&self, data: &[u8]) -> Result<Vec<u8>, IggyError>;
}
#[derive(Clone)]
pub struct Aes256GcmEncryptor {
cipher: Aes256Gcm,
}
impl Debug for Aes256GcmEncryptor {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Encryptor").finish()
}
}
impl Aes256GcmEncryptor {
pub fn new(key: &[u8]) -> Result<Self, IggyError> {
if key.len() != 32 {
return Err(IggyError::InvalidEncryptionKey);
}
Ok(Self {
cipher: Aes256Gcm::new(key.into()),
})
}
pub fn from_base64_key(key: &str) -> Result<Self, IggyError> {
Self::new(&text::from_base64_as_bytes(key)?)
}
}
impl Encryptor for Aes256GcmEncryptor {
fn encrypt(&self, data: &[u8]) -> Result<Vec<u8>, IggyError> {
let nonce = Aes256Gcm::generate_nonce(&mut OsRng);
let encrypted_data = self.cipher.encrypt(&nonce, data);
if encrypted_data.is_err() {
return Err(IggyError::CannotEncryptData);
}
let payload = [&nonce, encrypted_data.unwrap().as_slice()].concat();
Ok(payload)
}
fn decrypt(&self, data: &[u8]) -> Result<Vec<u8>, IggyError> {
let nonce = (&data[0..12]).into();
let payload = self.cipher.decrypt(nonce, &data[12..]);
if payload.is_err() {
return Err(IggyError::CannotDecryptData);
}
Ok(payload.unwrap())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{HeaderKey, HeaderValue, IggyMessage};
#[test]
fn given_the_same_key_data_should_be_encrypted_and_decrypted_correctly() {
let key = [1; 32];
let encryptor = Aes256GcmEncryptor::new(&key).unwrap();
let data = b"Hello World!";
let encrypted_data = encryptor.encrypt(data);
assert!(encrypted_data.is_ok());
let encrypted_data = encrypted_data.unwrap();
let decrypted_data = encryptor.decrypt(&encrypted_data);
assert!(decrypted_data.is_ok());
let decrypted_data = decrypted_data.unwrap();
assert_eq!(data, decrypted_data.as_slice());
}
#[test]
fn given_the_invalid_key_data_should_not_be_decrypted_correctly() {
let first_key = [1; 32];
let second_key = [2; 32];
let first_encryptor = Aes256GcmEncryptor::new(&first_key).unwrap();
let second_encryptor = Aes256GcmEncryptor::new(&second_key).unwrap();
let data = b"Hello World!";
let encrypted_data = first_encryptor.encrypt(data);
assert!(encrypted_data.is_ok());
let encrypted_data = encrypted_data.unwrap();
let decrypted_data = second_encryptor.decrypt(&encrypted_data);
assert!(decrypted_data.is_err());
let error = decrypted_data.err().unwrap();
assert_eq!(error.as_code(), IggyError::CannotDecryptData.as_code());
}
#[test]
fn message_payload_and_headers_should_encrypt_and_decrypt_correctly() {
use bytes::Bytes;
use std::collections::BTreeMap;
let key = [1; 32];
let encryptor = Aes256GcmEncryptor::new(&key).unwrap();
let mut headers = BTreeMap::new();
headers.insert(
HeaderKey::try_from("batch").unwrap(),
HeaderValue::from(1u64),
);
headers.insert(
HeaderKey::try_from("type").unwrap(),
HeaderValue::try_from("test-message").unwrap(),
);
let mut message = IggyMessage::builder()
.payload(Bytes::from("test payload data"))
.user_headers(headers)
.build()
.unwrap();
let original_payload = message.payload.clone();
let original_headers = message.user_headers.clone().unwrap();
message.payload = Bytes::from(encryptor.encrypt(&message.payload).unwrap());
message.header.payload_length = message.payload.len() as u32;
let encrypted_headers = encryptor.encrypt(&original_headers).unwrap();
message.header.user_headers_length = encrypted_headers.len() as u32;
message.user_headers = Some(Bytes::from(encrypted_headers));
assert_ne!(message.payload, original_payload);
assert_ne!(message.user_headers.as_ref().unwrap(), &original_headers);
let decrypted_payload = encryptor.decrypt(&message.payload).unwrap();
message.payload = Bytes::from(decrypted_payload);
message.header.payload_length = message.payload.len() as u32;
let decrypted_headers = encryptor
.decrypt(message.user_headers.as_ref().unwrap())
.unwrap();
message.header.user_headers_length = decrypted_headers.len() as u32;
message.user_headers = Some(Bytes::from(decrypted_headers));
assert_eq!(message.payload, original_payload);
assert_eq!(message.user_headers.as_ref().unwrap(), &original_headers);
let parsed = message.user_headers_map().unwrap().unwrap();
assert_eq!(
parsed
.get(&HeaderKey::try_from("batch").unwrap())
.unwrap()
.as_uint64()
.unwrap(),
1
);
assert_eq!(
parsed
.get(&HeaderKey::try_from("type").unwrap())
.unwrap()
.as_str()
.unwrap(),
"test-message"
);
}
}