use crate::{
aligned_buffer::{AlignedBufHolder, AlignedBufRouter},
arch::Simd,
base::{aegis::Aegis, block::Block},
careful::Nonce256,
easy::{AuthTag, AuthTag128, AuthTag256, Key256},
utils::num_bits,
};
use zerocopy::{FromZeros, transmute};
pub struct Aegis256<S: Simd, StateBlock, const OUTPUT_RATE_BYTES: usize> {
s0: StateBlock,
s1: StateBlock,
s2: StateBlock,
s3: StateBlock,
s4: StateBlock,
s5: StateBlock,
simd: S,
}
impl<S, StateBlock, const OUTPUT_RATE_BYTES: usize>
Aegis256<S, StateBlock, OUTPUT_RATE_BYTES>
where
S: Simd,
StateBlock: Block<SelfArray = [u8; OUTPUT_RATE_BYTES], Simd = S>,
{
#[inline(always)]
fn update(&mut self, m: StateBlock) {
let orig_s5 = self.s5;
self.s5 = StateBlock::aes_encrypt_round(self.s4, self.s5);
self.s4 = StateBlock::aes_encrypt_round(self.s3, self.s4);
self.s3 = StateBlock::aes_encrypt_round(self.s2, self.s3);
self.s2 = StateBlock::aes_encrypt_round(self.s1, self.s2);
self.s1 = StateBlock::aes_encrypt_round(self.s0, self.s1);
self.s0 = StateBlock::aes_encrypt_round(orig_s5, self.s0);
self.s0 = self.s0 ^ m;
}
}
impl<S, StateBlock, const OUTPUT_RATE_BYTES: usize> Aegis<OUTPUT_RATE_BYTES>
for Aegis256<S, StateBlock, OUTPUT_RATE_BYTES>
where
S: Simd,
StateBlock: Block<SelfArray = [u8; OUTPUT_RATE_BYTES], Simd = S>,
AlignedBufHolder: AlignedBufRouter<OUTPUT_RATE_BYTES>,
{
type AlignedBuf =
<AlignedBufHolder as AlignedBufRouter<OUTPUT_RATE_BYTES>>::AlignedBuf;
type Simd = S;
type Key = Key256;
type Nonce = Nonce256;
#[inline(always)]
fn init(simd: S, key: &Key256, nonce: Nonce256) -> Self {
let (k0_raw, k1_raw) = transmute!(*key.expose_secret());
let (n0_raw, n1_raw) = transmute!(nonce.into_array());
let (k0, k1, n0, n1) = (
StateBlock::from_128_bits(simd, k0_raw),
StateBlock::from_128_bits(simd, k1_raw),
StateBlock::from_128_bits(simd, n0_raw),
StateBlock::from_128_bits(simd, n1_raw),
);
let k0_xor_n0 = k0 ^ n0;
let k1_xor_n1 = k1 ^ n1;
let mut state = Self {
s0: k0_xor_n0,
s1: k1_xor_n1,
s2: StateBlock::c1(simd),
s3: StateBlock::c0(simd),
s4: k0 ^ StateBlock::c0(simd),
s5: k1 ^ StateBlock::c1(simd),
simd,
};
for _ in 0..4 {
for x in [k0, k1, k0_xor_n0, k1_xor_n1] {
state.s3 = state.s3 ^ StateBlock::ctx(simd);
state.s5 = state.s5 ^ StateBlock::ctx(simd);
state.update(x);
}
}
state
}
#[inline(always)]
fn absorb(&mut self, block: &[u8; OUTPUT_RATE_BYTES]) {
self.update(StateBlock::new(self.simd, *block));
}
#[inline(always)]
fn enc(
&mut self,
plaintext_chunk: &[u8; OUTPUT_RATE_BYTES],
ciphertext_chunk: &mut [u8; OUTPUT_RATE_BYTES],
) {
let xi = StateBlock::new(self.simd, *plaintext_chunk);
let z = xi ^ self.s1 ^ self.s4 ^ self.s5 ^ (self.s2 & self.s3);
*ciphertext_chunk = z.into_bytes();
self.update(xi); }
#[inline(always)]
fn dec(
&mut self,
ciphertext_chunk: &[u8; OUTPUT_RATE_BYTES],
plaintext_chunk: &mut [u8; OUTPUT_RATE_BYTES],
) {
let z = StateBlock::new(self.simd, *ciphertext_chunk)
^ self.s1
^ self.s4
^ self.s5
^ (self.s2 & self.s3);
*plaintext_chunk = z.into_bytes();
self.update(z); }
#[inline(always)]
fn dec_partial(
&mut self,
ciphertext_chunk: &[u8; OUTPUT_RATE_BYTES],
plaintext_remainder: &mut [u8],
) {
let z = StateBlock::new(self.simd, *ciphertext_chunk)
^ self.s1
^ self.s4
^ self.s5
^ (self.s2 & self.s3);
let mut out = z.into_bytes();
out[plaintext_remainder.len()..].zero(); plaintext_remainder.copy_from_slice(&out[..plaintext_remainder.len()]);
self.update(StateBlock::new(self.simd, out));
}
#[inline(always)]
fn finalize<Tag: AuthTag>(
&mut self,
plaintext: &[u8],
associated_data: &[u8],
) -> Tag {
let lengths = StateBlock::from_ints(
self.simd,
num_bits(associated_data),
num_bits(plaintext),
);
let t = self.s3 ^ lengths;
for _ in 0..7 {
self.update(t);
}
let mut tag = Tag::default();
match size_of::<Tag>() {
AuthTag128::BYTES => {
let block = self.s0 ^ self.s1 ^ self.s2 ^ self.s3 ^ self.s4 ^ self.s5;
tag.as_mut().copy_from_slice(&block.xor_down());
},
AuthTag256::BYTES => {
let high = (self.s0 ^ self.s1 ^ self.s2).xor_down();
let low = (self.s3 ^ self.s4 ^ self.s5).xor_down();
tag.as_mut().copy_from_slice([high, low].as_flattened());
},
_ => unreachable!("Tag limited to 16 or 32 bytes at compile-time."),
}
tag
}
}