#[cfg(hax)]
use hax_lib::int::*;
use crate::{
generic_keccak::KeccakState,
traits::{Absorb, KeccakItem, Squeeze},
};
#[cfg(hax)]
use crate::proof_utils::valid_rate;
#[hax_lib::fstar::before(
r#"
#push-options "--split_queries always --z3rlimit 300"
"#
)]
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,
squeeze_buf: [u8; RATE],
squeeze_pos: usize,
}
#[inline(always)]
#[hax_lib::fstar::replace(
"let buf_to_slices
(v_PARALLEL_LANES v_RATE: usize)
(buf: t_Array (t_Array u8 v_RATE) v_PARALLEL_LANES)
: t_Array (t_Slice u8) v_PARALLEL_LANES =
Core_models.Array.from_fn #(t_Slice u8)
v_PARALLEL_LANES
#(usize -> t_Slice u8)
(fun i -> Core_models.Array.impl_23__as_slice #u8 v_RATE (buf.[ i ]))
"
)]
fn buf_to_slices<const PARALLEL_LANES: usize, const RATE: usize>(
buf: &[[u8; RATE]; PARALLEL_LANES],
) -> [&[u8]; PARALLEL_LANES] {
core::array::from_fn(|i| buf[i].as_slice())
}
#[hax_lib::attributes]
#[hax_lib::fstar::options("--split_queries always --z3rlimit 300")] 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]
}
#[hax_lib::requires(
PARALLEL_LANES == 1 &&
valid_rate(RATE)
)]
#[hax_lib::ensures(|result|
result.state_inv() &&
result.squeeze_pos == RATE
)]
pub(crate) fn new() -> Self {
Self {
inner: KeccakState::new(),
buf: [Self::zero_block(); PARALLEL_LANES],
buf_len: 0,
sponge: false,
squeeze_buf: Self::zero_block(),
squeeze_pos: RATE,
}
}
#[hax_lib::requires(
PARALLEL_LANES == 1 &&
self.state_inv()
)]
#[hax_lib::ensures(|consumed|
future(self).state_inv() &&
if consumed == 0 {
future(self).buf_len == self.buf_len && (
self.buf_len == 0 ||
self.buf_len == RATE ||
self.buf_len.to_int() + inputs[0].len().to_int() < RATE.to_int()
)
} else {
future(self).buf_len == RATE &&
consumed.to_int() == RATE.to_int() - self.buf_len.to_int()
}
)]
pub(crate) fn fill_buffer(&mut self, inputs: &[&[u8]; PARALLEL_LANES]) -> usize {
let input_len = inputs[0].len();
if self.buf_len != 0 && input_len >= RATE - self.buf_len {
let consumed = RATE - self.buf_len;
#[cfg(hax)]
let self_buf_len = self.buf_len;
#[cfg(hax)]
let self_squeeze_pos = self.squeeze_pos;
#[allow(clippy::needless_range_loop)]
for i in 0..PARALLEL_LANES {
hax_lib::loop_invariant!(|_: usize| {
self.buf_len == self_buf_len && self.squeeze_pos == self_squeeze_pos
});
self.buf[i][self.buf_len..].copy_from_slice(&inputs[i][..consumed]);
}
self.buf_len = RATE;
consumed
} else {
0
}
}
#[hax_lib::requires(
PARALLEL_LANES == 1 &&
self.state_inv()
)]
#[hax_lib::ensures(|remainder|
future(self).state_inv() &&
future(self).buf_len.to_int() + remainder.to_int() <= RATE.to_int()
)]
pub(crate) fn absorb_full(&mut self, inputs: &[&[u8]; PARALLEL_LANES]) -> usize
where
KeccakState<PARALLEL_LANES, STATE>: Absorb<PARALLEL_LANES>,
{
debug_assert!(PARALLEL_LANES > 0);
debug_assert!(self.buf_len <= RATE);
#[cfg(all(debug_assertions, not(eurydice), not(hax)))]
{
for block in inputs {
debug_assert!(block.len() == inputs[0].len());
}
}
let consumed = self.fill_buffer(inputs);
if self.buf_len == RATE {
let borrowed = buf_to_slices(&self.buf);
self.inner.load_block::<RATE>(&borrowed, 0);
self.inner.keccakf1600();
self.buf_len = 0;
}
let input_to_consume = inputs[0].len() - consumed;
let num_blocks = input_to_consume / RATE;
let remainder = input_to_consume % RATE;
#[cfg(hax)]
let (self_buf_len, self_squeeze_pos, end) = {
let end = consumed + num_blocks * RATE;
hax_lib::assert!(end <= inputs[0].len());
(self.buf_len, self.squeeze_pos, end)
};
for i in 0..num_blocks {
hax_lib::loop_invariant!(
|_: usize| self.buf_len == self_buf_len && self.squeeze_pos == self_squeeze_pos
);
#[cfg(hax)]
crate::proof_utils::lemma_mul_succ_le(i, num_blocks, RATE);
let start = i * RATE + consumed;
#[cfg(hax)]
hax_lib::assert!(start + RATE <= end);
self.inner.load_block::<RATE>(inputs, start);
self.inner.keccakf1600();
}
remainder
}
#[inline(always)]
#[hax_lib::requires(
PARALLEL_LANES == 1 &&
self.state_inv()
)]
#[hax_lib::ensures(|_| future(self).state_inv())]
pub(crate) fn absorb(&mut self, inputs: &[&[u8]; PARALLEL_LANES])
where
KeccakState<PARALLEL_LANES, STATE>: Absorb<PARALLEL_LANES>,
{
let remainder = self.absorb_full(inputs);
if remainder > 0 {
#[cfg(not(eurydice))]
debug_assert!(
self.buf_len == 0 || self.buf_len + remainder <= RATE
);
#[cfg(hax)]
hax_lib::assert!(remainder.to_int() + self.buf_len.to_int() <= RATE.to_int());
let input_len = inputs[0].len();
#[cfg(hax)]
let self_buf_len = self.buf_len;
#[cfg(hax)]
let self_squeeze_pos = self.squeeze_pos;
#[allow(clippy::needless_range_loop)]
for i in 0..PARALLEL_LANES {
hax_lib::loop_invariant!(
|_: usize| self.buf_len == self_buf_len && self.squeeze_pos == self_squeeze_pos
);
self.buf[i][self.buf_len..self.buf_len + remainder]
.copy_from_slice(&inputs[i][input_len - remainder..input_len]);
}
self.buf_len += remainder;
}
}
#[inline(always)]
#[hax_lib::requires(
PARALLEL_LANES == 1 &&
self.state_inv()
)]
#[hax_lib::ensures(|_| future(self).state_inv())]
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 borrowed = buf_to_slices(&self.buf);
self.inner
.load_last::<RATE, DELIMITER>(&borrowed, 0, self.buf_len);
self.inner.keccakf1600();
}
}
#[hax_lib::attributes]
impl<const RATE: usize, STATE: KeccakItem<1>> KeccakXofState<1, RATE, STATE> {
#[inline(always)]
#[hax_lib::requires(
self.state_inv()
)]
#[hax_lib::ensures(|_|
future(self).state_inv() &&
future(out).len() == out.len()
)]
pub(crate) fn squeeze(&mut self, out: &mut [u8])
where
KeccakState<1, STATE>: Squeeze<STATE>,
{
let out_len = out.len();
if out_len == 0 {
return;
}
let mut out_offset = 0;
if self.squeeze_pos < RATE {
let avail = RATE - self.squeeze_pos;
let take = if avail < out_len { avail } else { out_len };
out[..take]
.copy_from_slice(&self.squeeze_buf[self.squeeze_pos..self.squeeze_pos + take]);
self.squeeze_pos += take;
#[cfg(hax)]
hax_lib::assert!(self.squeeze_pos <= RATE);
out_offset = take;
}
if out_offset == out_len {
return;
}
#[cfg(hax)]
hax_lib::assert!(self.squeeze_pos == RATE);
if self.sponge {
self.inner.keccakf1600();
}
self.sponge = true;
let remaining = out_len - out_offset;
let blocks = remaining / RATE;
let last_full = out_offset + blocks * RATE;
if blocks == 0 {
self.inner.squeeze::<RATE>(&mut self.squeeze_buf, 0, RATE);
out[out_offset..out_len].copy_from_slice(&self.squeeze_buf[..remaining]);
self.squeeze_pos = remaining;
} else {
self.inner.squeeze::<RATE>(out, out_offset, RATE);
#[cfg(hax)]
let self_buf_len = self.buf_len;
#[cfg(hax)]
let self_squeeze_pos = self.squeeze_pos;
for i in 1..blocks {
hax_lib::loop_invariant!(|_: usize| out.len() == out_len
&& self_buf_len == self.buf_len
&& self_squeeze_pos == self.squeeze_pos);
#[cfg(hax)]
hax_lib::assert!(
out_offset.to_int() + i.to_int() * RATE.to_int() <= out.len().to_int()
);
self.inner.keccakf1600();
self.inner.squeeze::<RATE>(out, out_offset + i * RATE, RATE);
}
let trailing = out_len - last_full;
if trailing > 0 {
self.inner.keccakf1600();
self.inner.squeeze::<RATE>(&mut self.squeeze_buf, 0, RATE);
out[last_full..out_len].copy_from_slice(&self.squeeze_buf[..trailing]);
self.squeeze_pos = trailing;
}
}
}
}
#[cfg(hax)]
mod proof_utils {
impl<
const PARALLEL_LANES: usize,
const RATE: usize,
State: super::KeccakItem<PARALLEL_LANES>,
> super::KeccakXofState<PARALLEL_LANES, RATE, State>
{
pub(crate) fn state_inv(&self) -> bool {
crate::proof_utils::valid_rate(RATE) && self.buf_len <= RATE && self.squeeze_pos <= RATE
}
}
}