use crate::{
generic_keccak::KeccakState,
traits::{Absorb, KeccakItem, Squeeze1},
};
pub(crate) struct KeccakXofState<
const PARALLEL_LANES: usize,
const RATE: usize,
STATE: KeccakItem<PARALLEL_LANES>,
> {
inner: KeccakState<PARALLEL_LANES, STATE>,
buf: [[u8; RATE]; PARALLEL_LANES],
buf_len: usize,
sponge: bool,
}
#[hax_lib::attributes]
impl<const PARALLEL_LANES: usize, const RATE: usize, STATE: KeccakItem<PARALLEL_LANES>>
KeccakXofState<PARALLEL_LANES, RATE, STATE>
{
pub(crate) const fn zero_block() -> [u8; RATE] {
[0u8; RATE]
}
pub(crate) fn new() -> Self {
Self {
inner: KeccakState::new(),
buf: [Self::zero_block(); PARALLEL_LANES],
buf_len: 0,
sponge: false,
}
}
#[inline(always)]
pub(crate) fn absorb(&mut self, inputs: &[&[u8]; PARALLEL_LANES])
where
KeccakState<PARALLEL_LANES, STATE>: Absorb<PARALLEL_LANES>,
{
let input_remainder_len = self.absorb_full(inputs);
if input_remainder_len > 0 {
#[cfg(not(eurydice))]
debug_assert!(
self.buf_len == 0 || self.buf_len + input_remainder_len <= RATE
);
let input_len = inputs[0].len();
#[allow(clippy::needless_range_loop)]
for i in 0..PARALLEL_LANES {
self.buf[i][self.buf_len..self.buf_len + input_remainder_len]
.copy_from_slice(&inputs[i][input_len - input_remainder_len..]);
}
self.buf_len += input_remainder_len;
}
}
pub(crate) fn absorb_full(&mut self, inputs: &[&[u8]; PARALLEL_LANES]) -> usize
where
KeccakState<PARALLEL_LANES, STATE>: Absorb<PARALLEL_LANES>,
{
#[cfg(not(eurydice))]
debug_assert!(PARALLEL_LANES > 0);
#[cfg(not(eurydice))]
debug_assert!(self.buf_len < RATE);
#[cfg(all(debug_assertions, not(hax), not(eurydice)))]
{
for block in inputs {
debug_assert!(block.len() == inputs[0].len());
}
}
let input_consumed = self.fill_buffer(inputs);
if input_consumed > 0 {
let borrowed: [&[u8]; PARALLEL_LANES] =
core::array::from_fn(|i| self.buf[i].as_slice());
self.inner.load_block::<RATE>(&borrowed, 0);
self.inner.keccakf1600();
self.buf_len = 0;
}
let input_to_consume = inputs[0].len() - input_consumed;
let num_blocks = input_to_consume / RATE;
let remainder = input_to_consume % RATE;
for i in 0..num_blocks {
self.inner
.load_block::<RATE>(inputs, input_consumed + i * RATE);
self.inner.keccakf1600();
}
remainder
}
pub(crate) fn fill_buffer(&mut self, inputs: &[&[u8]; PARALLEL_LANES]) -> usize {
let input_len = inputs[0].len();
let mut consumed = 0;
if self.buf_len > 0 {
if self.buf_len + input_len >= RATE {
consumed = RATE - self.buf_len;
#[allow(clippy::needless_range_loop)]
for i in 0..PARALLEL_LANES {
self.buf[i][self.buf_len..].copy_from_slice(&inputs[i][..consumed]);
}
self.buf_len += consumed;
}
}
consumed
}
#[inline(always)]
pub(crate) fn absorb_final<const DELIMITER: u8>(&mut self, inputs: &[&[u8]; PARALLEL_LANES])
where
KeccakState<PARALLEL_LANES, STATE>: Absorb<PARALLEL_LANES>,
{
self.absorb(inputs);
let mut borrowed = [[0u8; RATE].as_slice(); PARALLEL_LANES];
#[allow(clippy::needless_range_loop)]
for i in 0..PARALLEL_LANES {
borrowed[i] = &self.buf[i];
}
self.inner
.load_last::<RATE, DELIMITER>(&borrowed, 0, self.buf_len);
self.inner.keccakf1600();
}
}
impl<const RATE: usize, STATE: KeccakItem<1>> KeccakXofState<1, RATE, STATE> {
#[inline(always)]
pub(crate) fn squeeze(&mut self, out: &mut [u8])
where
KeccakState<1, STATE>: Squeeze1<STATE>,
{
if self.sponge {
self.inner.keccakf1600();
}
let out_len = out.len();
if out_len > 0 {
if out_len <= RATE {
self.inner.squeeze::<RATE>(out, 0, out_len);
} else {
let blocks = out_len / RATE;
for i in 0..blocks {
self.inner.keccakf1600();
self.inner.squeeze::<RATE>(out, i * RATE, RATE);
}
let remaining = out_len % RATE;
if remaining > 0 {
self.inner.keccakf1600();
self.inner.squeeze::<RATE>(out, blocks * RATE, remaining);
}
}
self.sponge = true;
}
}
}