use aes_gcm::{Key};
use hkdf::Hkdf;
use rand::{RngCore, rngs::OsRng};
use aes_gcm::{
aead::{Aead, Payload},
Aes128Gcm, Aes256Gcm, KeyInit, Nonce,
};
use sha2::Sha256;
use crate::error::{CasError, CasResult};
use super::cas_symmetric_encryption::{
CASAES128AadEncryption, CASAES128Encryption, CASAES256AadEncryption, CASAES256Encryption,
};
const AES_NONCE_LEN: usize = 12;
const AES128_KEY_LEN: usize = 16;
const AES256_KEY_LEN: usize = 32;
pub struct CASAES128;
pub struct CASAES256;
impl CASAES256Encryption for CASAES256 {
fn key_from_vec(key_slice: Vec<u8>) -> CasResult<Vec<u8>> {
if key_slice.len() != AES256_KEY_LEN {
return Err(CasError::InvalidKey);
}
let key = Key::<Aes256Gcm>::from_slice(key_slice.as_slice());
Ok(key.to_vec())
}
fn generate_key() -> Vec<u8> {
let mut os_rng = OsRng;
return Aes256Gcm::generate_key(&mut os_rng).to_vec();
}
fn encrypt_plaintext(aes_key: Vec<u8>, nonce: Vec<u8>, plaintext: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES256_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes256Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes256Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.encrypt(nonce, plaintext.as_ref())
.map_err(|_| CasError::EncryptionFailed)
}
fn decrypt_ciphertext(aes_key: Vec<u8>, nonce: Vec<u8>, ciphertext: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES256_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes256Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes256Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.decrypt(nonce, ciphertext.as_ref())
.map_err(|_| CasError::DecryptionFailed)
}
fn key_from_x25519_shared_secret(shared_secret: Vec<u8>) -> CasResult<Vec<u8>> {
let hk = Hkdf::<Sha256>::new(None, &shared_secret);
let mut aes_key: Box<[u8; 32]> = Box::new([0u8; 32]);
hk.expand(b"aes key", &mut *aes_key)
.map_err(|_| CasError::KeyGenerationFailed)?;
Ok(aes_key.to_vec())
}
fn generate_nonce() -> Vec<u8> {
let mut os_rng = OsRng;
let mut nonce = [0u8; AES_NONCE_LEN];
os_rng.fill_bytes(&mut nonce);
nonce.to_vec()
}
}
impl CASAES256AadEncryption for CASAES256 {
fn encrypt_plaintext_with_aad(aes_key: Vec<u8>, nonce: Vec<u8>, plaintext: Vec<u8>, aad: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES256_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes256Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes256Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.encrypt(nonce, Payload { msg: plaintext.as_slice(), aad: aad.as_slice() })
.map_err(|_| CasError::EncryptionFailed)
}
fn decrypt_ciphertext_with_aad(aes_key: Vec<u8>, nonce: Vec<u8>, ciphertext: Vec<u8>, aad: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES256_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes256Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes256Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.decrypt(nonce, Payload { msg: ciphertext.as_slice(), aad: aad.as_slice() })
.map_err(|_| CasError::DecryptionFailed)
}
}
impl CASAES128AadEncryption for CASAES128 {
fn encrypt_plaintext_with_aad(aes_key: Vec<u8>, nonce: Vec<u8>, plaintext: Vec<u8>, aad: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES128_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes128Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes128Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.encrypt(nonce, Payload { msg: plaintext.as_slice(), aad: aad.as_slice() })
.map_err(|_| CasError::EncryptionFailed)
}
fn decrypt_ciphertext_with_aad(aes_key: Vec<u8>, nonce: Vec<u8>, ciphertext: Vec<u8>, aad: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES128_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes128Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes128Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.decrypt(nonce, Payload { msg: ciphertext.as_slice(), aad: aad.as_slice() })
.map_err(|_| CasError::DecryptionFailed)
}
}
impl CASAES128Encryption for CASAES128 {
fn key_from_vec(key_slice: Vec<u8>) -> CasResult<Vec<u8>> {
if key_slice.len() != AES128_KEY_LEN {
return Err(CasError::InvalidKey);
}
let key = Key::<Aes128Gcm>::from_slice(key_slice.as_slice());
Ok(key.to_vec())
}
fn generate_key() -> Vec<u8> {
let mut os_rng = OsRng;
return Aes128Gcm::generate_key(&mut os_rng).to_vec();
}
fn encrypt_plaintext(aes_key: Vec<u8>, nonce: Vec<u8>, plaintext: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES128_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes128Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes128Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.encrypt(nonce, plaintext.as_ref())
.map_err(|_| CasError::EncryptionFailed)
}
fn decrypt_ciphertext(aes_key: Vec<u8>, nonce: Vec<u8>, ciphertext: Vec<u8>) -> CasResult<Vec<u8>> {
if aes_key.len() != AES128_KEY_LEN {
return Err(CasError::InvalidKey);
}
if nonce.len() != AES_NONCE_LEN {
return Err(CasError::InvalidNonce);
}
let key = Key::<Aes128Gcm>::from_slice(aes_key.as_slice());
let cipher = Aes128Gcm::new(key);
let nonce = Nonce::from_slice(nonce.as_slice());
cipher
.decrypt(nonce, ciphertext.as_ref())
.map_err(|_| CasError::DecryptionFailed)
}
fn key_from_x25519_shared_secret(shared_secret: Vec<u8>) -> CasResult<Vec<u8>> {
let hk = Hkdf::<Sha256>::new(None, &shared_secret);
let mut aes_key = Box::new([0u8; 16]);
hk.expand(b"aes key", &mut *aes_key)
.map_err(|_| CasError::KeyGenerationFailed)?;
Ok(aes_key.to_vec())
}
fn generate_nonce() -> Vec<u8> {
let mut os_rng = OsRng;
let mut nonce = [0u8; AES_NONCE_LEN];
os_rng.fill_bytes(&mut nonce);
nonce.to_vec()
}
}