use crate::Error;
use cipher::{
Array, Block, BlockCipherDecrypt, BlockCipherEncrypt, BlockSizeUser, Iv, IvSizeUser, Key,
KeyInit, KeyIvInit, KeySizeUser,
typenum::{Sum, U16, U32, U64, U128, Unsigned},
};
use core::fmt;
use hybrid_array::ArraySize;
#[cfg(feature = "zeroize")]
use zeroize::{Zeroize, ZeroizeOnDrop};
pub trait EmePoly: ArraySize {
fn mult_by_two(val: &Array<u8, Self>) -> Array<u8, Self>;
}
macro_rules! impl_eme_poly {
($size:ty, $limbs:expr, $mod_bytes:expr) => {
impl EmePoly for $size {
#[inline]
fn mult_by_two(val: &Array<u8, Self>) -> Array<u8, Self> {
let mut res = Array::<u8, Self>::default();
let mut v = [0u64; $limbs];
for i in 0..$limbs {
let mut buf = [0u8; 8];
buf.copy_from_slice(&val[i * 8..(i + 1) * 8]);
v[i] = u64::from_le_bytes(buf);
}
let carry_out = v[$limbs - 1] >> 63;
for i in (1..$limbs).rev() {
v[i] = (v[i] << 1) | (v[i - 1] >> 63);
}
v[0] <<= 1;
let mask = 0u64.wrapping_sub(carry_out);
let mod_bytes: [u8; 8] = $mod_bytes;
let mod_limb = u64::from_le_bytes(mod_bytes) & mask;
v[0] ^= mod_limb;
for i in 0..$limbs {
res[i * 8..(i + 1) * 8].copy_from_slice(&v[i].to_le_bytes());
}
res
}
}
};
}
impl_eme_poly!(U16, 2, [0x87, 0, 0, 0, 0, 0, 0, 0]);
impl_eme_poly!(U32, 4, [0x25, 0x04, 0, 0, 0, 0, 0, 0]);
impl_eme_poly!(U64, 8, [0x25, 0x01, 0, 0, 0, 0, 0, 0]);
impl_eme_poly!(U128, 16, [0xa3, 0x03, 0, 0, 0, 0, 0, 0]);
#[inline]
fn xor_blocks<C: BlockSizeUser>(out: &mut Block<C>, a: &Block<C>, b: &Block<C>) {
for (o, (x, y)) in out.iter_mut().zip(a.iter().zip(b.iter())) {
*o = *x ^ *y;
}
}
#[inline]
fn xor_into<C: BlockSizeUser>(out: &mut Block<C>, b: &Block<C>) {
for (o, x) in out.iter_mut().zip(b.iter()) {
*o ^= *x;
}
}
#[inline]
fn block_from_slice<C: BlockSizeUser>(s: &[u8]) -> Block<C> {
Block::<C>::try_from(s).unwrap_or_else(|_| unreachable!())
}
#[derive(Clone)]
#[cfg_attr(feature = "zeroize", derive(ZeroizeOnDrop))]
pub struct Eme2<C>
where
C: BlockCipherEncrypt + BlockCipherDecrypt + BlockSizeUser + KeySizeUser,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
{
#[cfg_attr(feature = "zeroize", zeroize(skip))]
cipher: C,
key2: Block<C>,
key3: Block<C>,
tweak: Block<C>,
}
impl<C> BlockSizeUser for Eme2<C>
where
C: BlockCipherEncrypt + BlockCipherDecrypt + BlockSizeUser + KeySizeUser,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
{
type BlockSize = C::BlockSize;
}
impl<C> IvSizeUser for Eme2<C>
where
C: BlockCipherEncrypt + BlockCipherDecrypt + BlockSizeUser + KeySizeUser,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
{
type IvSize = C::BlockSize;
}
impl<C> KeySizeUser for Eme2<C>
where
C: BlockCipherEncrypt + BlockCipherDecrypt + BlockSizeUser + KeySizeUser,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
C::KeySize: core::ops::Add<Sum<C::BlockSize, C::BlockSize>>,
Sum<C::KeySize, Sum<C::BlockSize, C::BlockSize>>: ArraySize,
{
type KeySize = Sum<C::KeySize, Sum<C::BlockSize, C::BlockSize>>;
}
impl<C> KeyInit for Eme2<C>
where
C: BlockCipherEncrypt + BlockCipherDecrypt + BlockSizeUser + KeyInit,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
C::KeySize: core::ops::Add<Sum<C::BlockSize, C::BlockSize>>,
Sum<C::KeySize, Sum<C::BlockSize, C::BlockSize>>: ArraySize,
{
fn new(key: &Key<Self>) -> Self {
let key_bytes = key.as_slice();
let ks = C::KeySize::USIZE;
let bs = C::BlockSize::USIZE;
let key1 = &key_bytes[..ks];
let key2 = block_from_slice::<C>(&key_bytes[ks..ks + bs]);
let key3 = block_from_slice::<C>(&key_bytes[ks + bs..ks + 2 * bs]);
let cipher = C::new_from_slice(key1).unwrap_or_else(|_| unreachable!());
let mut mode = Self {
cipher,
key2,
key3,
tweak: Block::<C>::default(),
};
mode.tweak = mode.hash_ad(&[]);
mode
}
}
impl<C> KeyIvInit for Eme2<C>
where
C: BlockCipherEncrypt + BlockCipherDecrypt + BlockSizeUser + KeyInit,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
C::KeySize: core::ops::Add<Sum<C::BlockSize, C::BlockSize>>,
Sum<C::KeySize, Sum<C::BlockSize, C::BlockSize>>: ArraySize,
{
#[inline]
fn new(key: &Key<Self>, iv: &Iv<Self>) -> Self {
let mut mode = <Self as KeyInit>::new(key);
mode.tweak = mode.hash_ad(iv.as_slice());
mode
}
}
impl<C> cipher::AlgorithmName for Eme2<C>
where
C: BlockCipherEncrypt
+ BlockCipherDecrypt
+ BlockSizeUser
+ KeySizeUser
+ cipher::AlgorithmName,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
{
fn write_alg_name(f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("extended::Eme2<")?;
<C as cipher::AlgorithmName>::write_alg_name(f)?;
f.write_str(">")
}
}
impl<C> fmt::Debug for Eme2<C>
where
C: BlockCipherEncrypt
+ BlockCipherDecrypt
+ BlockSizeUser
+ KeySizeUser
+ cipher::AlgorithmName,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("extended::Eme2<")?;
<C as cipher::AlgorithmName>::write_alg_name(f)?;
f.write_str("> { ... }")
}
}
impl<C> Eme2<C>
where
C: BlockCipherEncrypt + BlockCipherDecrypt + BlockSizeUser + KeySizeUser,
C::BlockSize: EmePoly + core::ops::Add<C::BlockSize>,
Sum<C::BlockSize, C::BlockSize>: ArraySize,
{
pub fn hash_ad(&self, ad: &[u8]) -> Block<C> {
let bs = C::BlockSize::USIZE;
if ad.is_empty() {
let mut t_star = self.key3.clone();
self.cipher.encrypt_block(&mut t_star);
return t_star;
}
let mut current_key3 = C::BlockSize::mult_by_two(&self.key3);
let mut tt = Block::<C>::default();
let chunks = ad.chunks(bs);
let r = chunks.len();
for (i, chunk) in chunks.enumerate() {
let is_last = i == r - 1;
let mut block = Block::<C>::default();
if is_last {
if chunk.len() < bs {
block[..chunk.len()].copy_from_slice(chunk);
block[chunk.len()] = 0x80;
current_key3 = C::BlockSize::mult_by_two(¤t_key3);
} else {
block.copy_from_slice(chunk);
}
} else {
block.copy_from_slice(chunk);
}
xor_into::<C>(&mut block, ¤t_key3);
self.cipher.encrypt_block(&mut block);
xor_into::<C>(&mut block, ¤t_key3);
xor_into::<C>(&mut tt, &block);
if !is_last {
current_key3 = C::BlockSize::mult_by_two(¤t_key3);
}
}
#[cfg(feature = "zeroize")]
{
current_key3.zeroize();
}
tt
}
pub const fn t_star(&self) -> &Block<C> {
&self.tweak
}
pub fn set_t_star(&mut self, t_star: Block<C>) {
self.tweak = t_star;
}
pub fn encrypt(&self, data: &mut [u8]) -> Result<(), Error> {
let tweak = self.tweak.clone();
self.encrypt_core(&tweak, data)
}
pub fn decrypt(&self, data: &mut [u8]) -> Result<(), Error> {
let tweak = self.tweak.clone();
self.decrypt_core(&tweak, data)
}
pub fn encrypt_with_ad(&self, associated_data: &[u8], data: &mut [u8]) -> Result<(), Error> {
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))]
let mut t_star = self.hash_ad(associated_data);
let res = self.encrypt_core(&t_star, data);
#[cfg(feature = "zeroize")]
{
t_star.zeroize();
}
res
}
pub fn decrypt_with_ad(&self, associated_data: &[u8], data: &mut [u8]) -> Result<(), Error> {
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))]
let mut t_star = self.hash_ad(associated_data);
let res = self.decrypt_core(&t_star, data);
#[cfg(feature = "zeroize")]
{
t_star.zeroize();
}
res
}
fn encrypt_core(&self, tweak: &Block<C>, data: &mut [u8]) -> Result<(), Error> {
let bs = C::BlockSize::USIZE;
let len = data.len();
if len < bs {
return Err(Error::DataTooShort);
}
let m = len.div_ceil(bs);
let last_full = if len.is_multiple_of(bs) { m } else { m - 1 };
let (mask_cipher_first, mask_delta_first, ccc_m, c_m) =
self.encrypt_pass1_and_mix(tweak, data, len, bs, m, last_full);
self.encrypt_pass2_and_3(
tweak,
data,
len,
bs,
m,
last_full,
mask_cipher_first,
mask_delta_first,
ccc_m,
c_m,
);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn encrypt_pass1_and_mix(
&self,
tweak: &Block<C>,
data: &mut [u8],
len: usize,
bs: usize,
m: usize,
last_full: usize,
) -> (Block<C>, Block<C>, Block<C>, Block<C>) {
let mut l_current = self.key2.clone();
for i in 0..last_full {
let mut block = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
xor_into::<C>(&mut block, &l_current);
self.cipher.encrypt_block(&mut block);
data[i * bs..(i + 1) * bs].copy_from_slice(&block);
l_current = C::BlockSize::mult_by_two(&l_current);
}
let mut ppp_m = Block::<C>::default();
if last_full < m {
let rem = len % bs;
ppp_m[..rem].copy_from_slice(&data[last_full * bs..len]);
ppp_m[rem] = 0x80;
}
let mut sp = Block::<C>::default();
for i in 1..last_full {
let ppp_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
xor_into::<C>(&mut sp, &ppp_i);
}
if last_full < m {
xor_into::<C>(&mut sp, &ppp_m);
}
let ppp_0 = block_from_slice::<C>(&data[0..bs]);
let mut mask_plain_first = Block::<C>::default();
xor_blocks::<C>(&mut mask_plain_first, &ppp_0, &sp);
xor_into::<C>(&mut mask_plain_first, tweak);
let mut ccc_m = Block::<C>::default();
let mut c_m = Block::<C>::default();
let mut mm = Block::<C>::default();
let mask_cipher_first = if last_full < m {
mm.copy_from_slice(mask_plain_first.as_slice());
self.cipher.encrypt_block(&mut mm);
let mut mask_cipher_first = mm.clone();
self.cipher.encrypt_block(&mut mask_cipher_first);
let rem = len % bs;
for i in 0..rem {
c_m[i] = data[last_full * bs + i] ^ mm[i];
}
ccc_m[..rem].copy_from_slice(&c_m[..rem]);
ccc_m[rem] = 0x80;
mask_cipher_first
} else {
let mut mask_cipher_first = mask_plain_first.clone();
self.cipher.encrypt_block(&mut mask_cipher_first);
mask_cipher_first
};
let mut mask_delta_first = Block::<C>::default();
xor_blocks::<C>(&mut mask_delta_first, &mask_plain_first, &mask_cipher_first);
#[cfg(feature = "zeroize")]
{
l_current.zeroize();
ppp_m.zeroize();
sp.zeroize();
mask_plain_first.zeroize();
mm.zeroize();
}
(mask_cipher_first, mask_delta_first, ccc_m, c_m)
}
#[allow(clippy::too_many_arguments)]
fn encrypt_pass2_and_3(
&self,
tweak: &Block<C>,
data: &mut [u8],
len: usize,
bs: usize,
m: usize,
last_full: usize,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut mask_cipher_first: Block<C>,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut mask_delta_first: Block<C>,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut ccc_m: Block<C>,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut c_m: Block<C>,
) {
let mut current_m_j = mask_delta_first.clone();
let mut current_m_k = mask_delta_first.clone();
for i in 1..last_full {
let ppp_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
let mut ccc_i = Block::<C>::default();
let k = i & 127;
if k == 0 {
let mut mask_plain_block = Block::<C>::default();
xor_blocks::<C>(&mut mask_plain_block, &ppp_i, &mask_delta_first);
let mut mask_cipher_block = mask_plain_block.clone();
self.cipher.encrypt_block(&mut mask_cipher_block);
xor_blocks::<C>(&mut current_m_j, &mask_plain_block, &mask_cipher_block);
xor_blocks::<C>(&mut ccc_i, &mask_cipher_block, &mask_delta_first);
current_m_k = current_m_j.clone();
#[cfg(feature = "zeroize")]
{
mask_plain_block.zeroize();
mask_cipher_block.zeroize();
}
} else {
current_m_k = C::BlockSize::mult_by_two(¤t_m_k);
xor_blocks::<C>(&mut ccc_i, &ppp_i, ¤t_m_k);
}
data[i * bs..(i + 1) * bs].copy_from_slice(&ccc_i);
}
let mut sc = Block::<C>::default();
for i in 1..last_full {
let ccc_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
xor_into::<C>(&mut sc, &ccc_i);
}
if last_full < m {
xor_into::<C>(&mut sc, &ccc_m);
}
let mut ccc_0 = Block::<C>::default();
xor_blocks::<C>(&mut ccc_0, &mask_cipher_first, &sc);
xor_into::<C>(&mut ccc_0, tweak);
data[0..bs].copy_from_slice(&ccc_0);
let mut l_current = self.key2.clone();
for i in 0..last_full {
let mut cc_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
self.cipher.encrypt_block(&mut cc_i);
xor_into::<C>(&mut cc_i, &l_current);
data[i * bs..(i + 1) * bs].copy_from_slice(&cc_i);
l_current = C::BlockSize::mult_by_two(&l_current);
}
if last_full < m {
let rem = len % bs;
data[last_full * bs..len].copy_from_slice(&c_m[..rem]);
}
#[cfg(feature = "zeroize")]
{
mask_cipher_first.zeroize();
ccc_m.zeroize();
c_m.zeroize();
mask_delta_first.zeroize();
current_m_j.zeroize();
current_m_k.zeroize();
sc.zeroize();
ccc_0.zeroize();
l_current.zeroize();
}
}
fn decrypt_core(&self, tweak: &Block<C>, data: &mut [u8]) -> Result<(), Error> {
let bs = C::BlockSize::USIZE;
let len = data.len();
if len < bs {
return Err(Error::DataTooShort);
}
let m = len.div_ceil(bs);
let last_full = if len.is_multiple_of(bs) { m } else { m - 1 };
let (mask_plain_first, mask_delta_first, ppp_m, p_m) =
self.decrypt_pass1_and_mix(tweak, data, len, bs, m, last_full);
self.decrypt_pass2_and_3(
tweak,
data,
len,
bs,
m,
last_full,
mask_plain_first,
mask_delta_first,
ppp_m,
p_m,
);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn decrypt_pass1_and_mix(
&self,
tweak: &Block<C>,
data: &mut [u8],
len: usize,
bs: usize,
m: usize,
last_full: usize,
) -> (Block<C>, Block<C>, Block<C>, Block<C>) {
let mut l_current = self.key2.clone();
for i in 0..last_full {
let mut block = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
xor_into::<C>(&mut block, &l_current);
self.cipher.decrypt_block(&mut block);
data[i * bs..(i + 1) * bs].copy_from_slice(&block);
l_current = C::BlockSize::mult_by_two(&l_current);
}
let mut ccc_m = Block::<C>::default();
if last_full < m {
let rem = len % bs;
ccc_m[..rem].copy_from_slice(&data[last_full * bs..len]);
ccc_m[rem] = 0x80;
}
let mut sc = Block::<C>::default();
for i in 1..last_full {
let ccc_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
xor_into::<C>(&mut sc, &ccc_i);
}
if last_full < m {
xor_into::<C>(&mut sc, &ccc_m);
}
let ccc_0 = block_from_slice::<C>(&data[0..bs]);
let mut mask_cipher_first = Block::<C>::default();
xor_blocks::<C>(&mut mask_cipher_first, &ccc_0, &sc);
xor_into::<C>(&mut mask_cipher_first, tweak);
let mut ppp_m = Block::<C>::default();
let mut p_m = Block::<C>::default();
let mut mm = Block::<C>::default();
let mask_plain_first = if last_full < m {
mm.copy_from_slice(mask_cipher_first.as_slice());
self.cipher.decrypt_block(&mut mm);
let mut mask_plain_first = mm.clone();
self.cipher.decrypt_block(&mut mask_plain_first);
let rem = len % bs;
for i in 0..rem {
p_m[i] = data[last_full * bs + i] ^ mm[i];
}
ppp_m[..rem].copy_from_slice(&p_m[..rem]);
ppp_m[rem] = 0x80;
mask_plain_first
} else {
let mut mask_plain_first = mask_cipher_first.clone();
self.cipher.decrypt_block(&mut mask_plain_first);
mask_plain_first
};
let mut mask_delta_first = Block::<C>::default();
xor_blocks::<C>(&mut mask_delta_first, &mask_plain_first, &mask_cipher_first);
#[cfg(feature = "zeroize")]
{
l_current.zeroize();
ccc_m.zeroize();
sc.zeroize();
mask_cipher_first.zeroize();
mm.zeroize();
}
(mask_plain_first, mask_delta_first, ppp_m, p_m)
}
#[allow(clippy::too_many_arguments)]
fn decrypt_pass2_and_3(
&self,
tweak: &Block<C>,
data: &mut [u8],
len: usize,
bs: usize,
m: usize,
last_full: usize,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut mask_plain_first: Block<C>,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut mask_delta_first: Block<C>,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut ppp_m: Block<C>,
#[cfg_attr(not(feature = "zeroize"), allow(unused_mut))] mut p_m: Block<C>,
) {
let mut current_m_j = mask_delta_first.clone();
let mut current_m_k = mask_delta_first.clone();
for i in 1..last_full {
let ccc_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
let mut ppp_i = Block::<C>::default();
let k = i & 127;
if k == 0 {
let mut mask_cipher_block = Block::<C>::default();
xor_blocks::<C>(&mut mask_cipher_block, &ccc_i, &mask_delta_first);
let mut mask_plain_block = mask_cipher_block.clone();
self.cipher.decrypt_block(&mut mask_plain_block);
xor_blocks::<C>(&mut current_m_j, &mask_plain_block, &mask_cipher_block);
xor_blocks::<C>(&mut ppp_i, &mask_plain_block, &mask_delta_first);
current_m_k = current_m_j.clone();
#[cfg(feature = "zeroize")]
{
mask_cipher_block.zeroize();
mask_plain_block.zeroize();
}
} else {
current_m_k = C::BlockSize::mult_by_two(¤t_m_k);
xor_blocks::<C>(&mut ppp_i, &ccc_i, ¤t_m_k);
}
data[i * bs..(i + 1) * bs].copy_from_slice(&ppp_i);
}
let mut sp = Block::<C>::default();
for i in 1..last_full {
let ppp_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
xor_into::<C>(&mut sp, &ppp_i);
}
if last_full < m {
xor_into::<C>(&mut sp, &ppp_m);
}
let mut ppp_0 = Block::<C>::default();
xor_blocks::<C>(&mut ppp_0, &mask_plain_first, &sp);
xor_into::<C>(&mut ppp_0, tweak);
data[0..bs].copy_from_slice(&ppp_0);
let mut l_current = self.key2.clone();
for i in 0..last_full {
let mut pp_i = block_from_slice::<C>(&data[i * bs..(i + 1) * bs]);
self.cipher.decrypt_block(&mut pp_i);
xor_into::<C>(&mut pp_i, &l_current);
data[i * bs..(i + 1) * bs].copy_from_slice(&pp_i);
l_current = C::BlockSize::mult_by_two(&l_current);
}
if last_full < m {
let rem = len % bs;
data[last_full * bs..len].copy_from_slice(&p_m[..rem]);
}
#[cfg(feature = "zeroize")]
{
mask_plain_first.zeroize();
ppp_m.zeroize();
p_m.zeroize();
mask_delta_first.zeroize();
current_m_j.zeroize();
current_m_k.zeroize();
sp.zeroize();
ppp_0.zeroize();
l_current.zeroize();
}
}
}