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}; 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>> {
let (padded, pow) = pad_derived_key(key2.as_bytes(), key1.as_bytes());
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()
}