shared-aes-enc 0.3.8

A shared AES encryption library providing secure encryption and decryption functionality
Documentation
use std::slice;

use aes::Aes256;
use aes::cipher::{
    BlockCipher, BlockDecrypt, NewBlockCipher,
    generic_array::{ArrayLength, GenericArray, typenum::Unsigned},
};
use base64::{Engine as _, engine::general_purpose}; // For encoding/decoding
use block_modes::block_padding::{Padding, Pkcs7};
use block_modes::{BlockMode, Ecb};
use rand::Rng;
use sha3::{Digest, Sha3_256};

type Aes256Ecb = Ecb<Aes256, Pkcs7>;

pub(crate) fn to_blocks<N>(data: &mut [u8]) -> &mut [GenericArray<u8, N>]
where
    N: ArrayLength<u8>,
{
    let n = N::to_usize();
    debug_assert!(data.len() % n == 0);

    #[allow(unsafe_code)]
    unsafe {
        slice::from_raw_parts_mut(data.as_ptr() as *mut GenericArray<u8, N>, data.len() / n)
    }
}

fn pad_derived_key(key1: &[u8], key2: &[u8]) -> (Vec<u8>, usize) {
    let mut pow = key1.len().next_power_of_two() as i32;
    if pow < 0 {
        pow = key1.len().next_power_of_two().next_power_of_two() as i32;
    }
    let padding = vec![0u8; pow as usize - key1.len() - 2];
    return ([key1, &padding, key2].concat(), pow as usize);
}

fn reverse_bits_in_u32(value: u32, n: u8) -> u32 {
    assert!(n <= 32, "Cannot reverse more than 32 bits");

    let mut reversed = 0;
    for i in 0..n {
        let bit = (value >> i) & 1;
        reversed |= bit << (n - 1 - i);
    }
    reversed
}

fn bit_reverse(data: &[u8], pow: usize) -> Vec<u8> {
    assert!(pow & (pow - 1) == 0, "pow must be a power of two");
    let pow_log2 = pow.trailing_zeros();
    let mut result = Vec::with_capacity(pow);
    for i in 0..pow {
        result.push(data[reverse_bits_in_u32(i as u32, pow_log2 as u8) as usize]);
    }
    result
}

fn derive_key(
    key1: &str,
    key2: &str,
) -> Result<GenericArray<u8, <Aes256 as NewBlockCipher>::KeySize>, Box<dyn std::error::Error>> {
    // pad the key to the next power of two for bit shuffling
    let (padded, pow) = pad_derived_key(key2.as_bytes(), key1.as_bytes());
    // use bit reverse to make sure the key is not predictable
    let reversed = bit_reverse(&padded, pow);
    let mut hasher = Sha3_256::default();
    hasher.update(&reversed);
    let key = hasher.finalize();
    Ok(*GenericArray::from_slice(&key))
}

pub fn shared_key_encrypt_bytes(
    key1: &str,
    key2: &str,
    data: &[u8],
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
    let key = derive_key(key1, key2)?;
    let cipher = Aes256Ecb::new_from_slices(&key, &[])?;

    let ciphertext = cipher.encrypt_vec(data);

    Ok(ciphertext)
}

pub fn shared_key_encrypt(
    key1: &str,
    key2: &str,
    data: &str,
) -> Result<String, Box<dyn std::error::Error>> {
    let res = shared_key_encrypt_bytes(key1, key2, data.as_bytes())?;
    Ok(general_purpose::STANDARD.encode(&res))
}

pub fn shared_key_decrypt(
    key1: &str,
    key2: &str,
    encrypted_b64: &str,
) -> Result<String, Box<dyn std::error::Error>> {
    let encrypted = general_purpose::STANDARD
        .decode(encrypted_b64)
        .map_err(|e| format!("base64 decode failed: {:?}", e))?;

    let key = derive_key(key1, key2)?;
    let cipher = Aes256::new(&key);

    let bs = <Aes256 as BlockCipher>::BlockSize::to_usize();
    if encrypted.len() % bs != 0 {
        return Err("encrypted length not multiple of block size".into());
    }
    let mut buf = encrypted.to_vec();
    let blocks = to_blocks(&mut buf);
    cipher.decrypt_blocks(blocks);

    if let Ok(unpadded) = Pkcs7::unpad(&buf) {
        let unpadded_len = unpadded.len();
        buf.truncate(unpadded_len);
        if let Ok(s) = String::from_utf8(buf.to_vec()) {
            Ok(s)
        } else {
            Ok(general_purpose::STANDARD.encode(&buf))
        }
    } else {
        Ok(general_purpose::STANDARD.encode(&buf))
    }
}

pub fn generate_password(length: usize) -> String {
    const CHARSET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ\
                             abcdefghijklmnopqrstuvwxyz\
                             0123456789";
    let mut rng = rand::rng();

    (0..length)
        .map(|_| {
            let idx = rng.random_range(0..CHARSET.len());
            CHARSET[idx] as char
        })
        .collect()
}