use crate::CipherGeneration;
use crate::serialization::RubyMarshal;
use aes_gcm::{
Aes128Gcm,
aead::{Aead, KeyInit, Nonce},
};
use anyhow::anyhow;
use base64::{Engine as _, engine::general_purpose};
pub struct MessageEncryption {
message: Vec<u8>,
key: String,
aad: String,
}
impl MessageEncryption {
pub fn new(message: Vec<u8>, key: &str, aad: &str) -> Self {
MessageEncryption {
message,
key: key.to_string(),
aad: aad.to_string(),
}
}
pub fn decrypt(&self, iv: &str, tag: &str) -> anyhow::Result<String> {
if let (Ok(key), Ok(iv), Ok(message), Ok(tag)) = (
hex_to_bytes(&self.key),
general_purpose::STANDARD.decode(iv),
general_purpose::STANDARD.decode(&self.message),
general_purpose::STANDARD.decode(tag),
) && let (Ok(decipher), Ok(nonce)) = (
Aes128Gcm::new_from_slice(&key),
Nonce::<Aes128Gcm>::try_from(iv.as_slice()),
) {
let mut ciphertext = message;
ciphertext.extend_from_slice(&tag);
let payload = aes_gcm::aead::Payload {
msg: &ciphertext,
aad: self.aad.as_bytes(),
};
let plaintext = decipher.decrypt(&nonce, payload);
if let Ok(plaintext) = plaintext {
let content = RubyMarshal::deserialize(plaintext)?;
return Ok(String::from_utf8(content)?);
}
}
Err(anyhow!("Decryption not successful"))
}
pub fn encrypt(&self) -> anyhow::Result<String> {
if let Ok(key) = hex_to_bytes(&self.key) {
let random_iv = CipherGeneration::random_iv();
if let (Ok(cipher), Ok(nonce)) = (
Aes128Gcm::new_from_slice(&key),
Nonce::<Aes128Gcm>::try_from(random_iv.as_slice()),
) {
let serialized_message = RubyMarshal::serialize(std::str::from_utf8(&self.message)?)?;
let payload = aes_gcm::aead::Payload {
msg: &serialized_message,
aad: self.aad.as_bytes(),
};
let encrypted = cipher.encrypt(&nonce, payload);
if let Ok(encrypted) = encrypted {
let (ct, tag) = encrypted.split_at(encrypted.len() - 16);
let encryption_result = format!(
"{}--{}--{}",
general_purpose::STANDARD.encode(ct),
general_purpose::STANDARD.encode(&random_iv),
general_purpose::STANDARD.encode(tag)
);
return Ok(encryption_result);
}
}
}
Err(anyhow!("Encryption not successful"))
}
pub fn split_encrypted_contents(contents: &str) -> anyhow::Result<Vec<&str>> {
let contents = contents.split("--").fold(Vec::new(), |mut acc, content| {
acc.push(content);
acc
});
if contents.len() == 3 {
Ok(contents)
} else {
Err(anyhow!("Invalid encrypted contents"))
}
}
}
fn hex_to_bytes(raw_hex: &str) -> Result<Vec<u8>, hex::FromHexError> {
hex::decode(raw_hex)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_encrypt_decrypt_cycle() {
let key = "8872ebc11db3ea2ed08cc629d199b164";
let aad = "";
let plaintext_message = b"banana: true
apple: false
orange: false";
let encryptor = MessageEncryption::new(plaintext_message.to_vec(), key, aad);
let encrypted_result = match encryptor.encrypt() {
Ok(encrypted_contents) => encrypted_contents,
Err(..) => panic!("first encryption failed"),
};
let split_data = MessageEncryption::split_encrypted_contents(&encrypted_result).unwrap();
let new_message = split_data[0];
let new_iv = split_data[1];
let new_aad = split_data[2];
let decryptor = MessageEncryption::new(new_message.as_bytes().to_vec(), key, aad);
let decrypted_result = decryptor.decrypt(new_iv, new_aad);
let encryptor = match decrypted_result {
Ok(decrypted_contents) => {
MessageEncryption::new(decrypted_contents.as_bytes().to_vec(), key, aad)
}
Err(why) => panic!("first decryption failed {}", why),
};
let encrypted_result = match encryptor.encrypt() {
Ok(encrypted_contents) => encrypted_contents,
Err(why) => panic!("second encryption failed {}", why),
};
let split_data = MessageEncryption::split_encrypted_contents(&encrypted_result).unwrap();
let new_message = split_data[0];
let new_iv = split_data[1];
let new_aad = split_data[2];
let decryptor = MessageEncryption::new(new_message.as_bytes().to_vec(), key, aad);
let decrypted_result = decryptor.decrypt(new_iv, new_aad);
match decrypted_result {
Ok(decrypted_contents) => {
assert_eq!(decrypted_contents.as_bytes(), plaintext_message);
}
Err(_) => panic!("second decryption failed"),
};
}
#[test]
fn test_encryption_decryption_with_aad() {
let key = "8872ebc11db3ea2ed08cc629d199b164";
let aad = "some value";
let plaintext_message = "banana: true
apple: false
orange: false";
let encryptor = MessageEncryption::new(plaintext_message.as_bytes().to_vec(), key, aad);
let encrypted_result = match encryptor.encrypt() {
Ok(encrypted_contents) => encrypted_contents,
Err(..) => panic!("first encryption failed"),
};
let split_data = MessageEncryption::split_encrypted_contents(&encrypted_result).unwrap();
let new_message = split_data[0];
let new_iv = split_data[1];
let new_aad = split_data[2];
let decryptor = MessageEncryption::new(new_message.as_bytes().to_vec(), key, aad);
let result = decryptor.decrypt(new_iv, new_aad);
assert_eq!(plaintext_message, result.unwrap());
}
#[test]
fn test_decryption_fails_with_incorrect_iv() {
let key = "94b6b40cabf62ee59c9aa13a86f0e7d7";
let aad = "";
let encrypted_message = b"1alR88JGbSy1wz44cgVgZC3mH2Fg9HjRFtl6NwRoOfpqNzJ61Ub48O1YhJUqaszJgJ8=";
let decryptor = MessageEncryption::new(encrypted_message.to_vec(), key, aad);
let result = decryptor.decrypt("123456789012345", "pksKcg/so9Pq3UMHjfnVsg==");
assert!(result.is_err());
}
#[test]
fn test_encryption_fails_with_non_hex_key() {
let key = "8872ebc11db3ea2";
let aad = "";
let plaintext_message = b"banana: true
apple: false
orange: false";
let encryptor = MessageEncryption::new(plaintext_message.to_vec(), key, aad);
let result = encryptor.encrypt();
assert!(result.is_err());
}
#[test]
fn test_invalid_aad_for_decrypt() {
let invalid_aad = "66ag";
let aad = "";
let key = "8872ebc11db3ea2";
let plaintext_message = b"banana: true
apple: false
orange: false";
let decryptor = MessageEncryption::new(plaintext_message.to_vec(), key, aad);
let result = decryptor.decrypt("", invalid_aad);
assert!(result.is_err());
}
}