use alloc::vec::Vec;
use core::array;
use miden_core::field::PrimeCharacteristicRing;
pub const STATE_WIDTH: usize = 12;
pub const MAX_SBOXES_PER_ROW: usize = STATE_WIDTH + 1;
pub const NUM_CUBE_REGS: usize = MAX_SBOXES_PER_ROW;
fn sbox<E: PrimeCharacteristicRing>(x: E, reg: &E) -> (E, E) {
(reg.clone().square() * x.clone(), reg.clone() - x.cube())
}
pub fn apply_single_ext<E: PrimeCharacteristicRing>(
h: &[E; STATE_WIDTH],
ark: &[E; STATE_WIDTH],
cube_regs: &[E],
) -> ([E; STATE_WIDTH], Vec<E>) {
let with_rc: [E; STATE_WIDTH] = array::from_fn(|i| h[i].clone() + ark[i].clone());
let mut checks = Vec::new();
let with_sbox: [E; STATE_WIDTH] = array::from_fn(|i| {
let (out, check) = sbox(with_rc[i].clone(), &cube_regs[i]);
checks.push(check);
out
});
(apply_matmul_external(&with_sbox), checks)
}
pub fn apply_init_plus_ext<E: PrimeCharacteristicRing>(
h: &[E; STATE_WIDTH],
ark_ext: &[E; STATE_WIDTH],
cube_regs: &[E],
) -> ([E; STATE_WIDTH], Vec<E>) {
let pre = apply_matmul_external(h);
let with_rc: [E; STATE_WIDTH] = array::from_fn(|i| pre[i].clone() + ark_ext[i].clone());
let mut checks = Vec::new();
let with_sbox: [E; STATE_WIDTH] = array::from_fn(|i| {
let (out, check) = sbox(with_rc[i].clone(), &cube_regs[i]);
checks.push(check);
out
});
(apply_matmul_external(&with_sbox), checks)
}
pub fn apply_packed_internals<E: PrimeCharacteristicRing>(
h: &[E; STATE_WIDTH],
w: &[E; 3],
ark_int: &[E; 3],
mat_diag: &[E; STATE_WIDTH],
cube_regs: &[E],
) -> ([E; STATE_WIDTH], [E; 3], Vec<E>) {
let mut state = h.clone();
let mut witness_checks: [E; 3] = array::from_fn(|_| E::ZERO);
let mut cube_checks = Vec::new();
for k in 0..3 {
let sbox_input = state[0].clone() + ark_int[k].clone();
let (out, check) = sbox(sbox_input, &cube_regs[k]);
witness_checks[k] = w[k].clone() - out;
cube_checks.push(check);
state[0] = w[k].clone();
state = apply_matmul_internal(&state, mat_diag);
}
(state, witness_checks, cube_checks)
}
pub fn apply_internal_plus_ext<E: PrimeCharacteristicRing>(
h: &[E; STATE_WIDTH],
w0: &E,
ark_int_const: E,
ark_ext: &[E; STATE_WIDTH],
mat_diag: &[E; STATE_WIDTH],
cube_regs: &[E],
) -> ([E; STATE_WIDTH], E, Vec<E>) {
let mut cube_checks = Vec::new();
let sbox_input = h[0].clone() + ark_int_const;
let (int_out, int_check) = sbox(sbox_input, &cube_regs[STATE_WIDTH]);
let witness_check = w0.clone() - int_out;
cube_checks.push(int_check);
let mut int_state = h.clone();
int_state[0] = w0.clone();
let intermediate = apply_matmul_internal(&int_state, mat_diag);
let with_rc: [E; STATE_WIDTH] =
array::from_fn(|i| intermediate[i].clone() + ark_ext[i].clone());
let with_sbox: [E; STATE_WIDTH] = array::from_fn(|i| {
let (out, check) = sbox(with_rc[i].clone(), &cube_regs[i]);
cube_checks.push(check);
out
});
let next_state = apply_matmul_external(&with_sbox);
(next_state, witness_check, cube_checks)
}
pub fn apply_matmul_external<E: PrimeCharacteristicRing>(
state: &[E; STATE_WIDTH],
) -> [E; STATE_WIDTH] {
let b0 = matmul_m4(array::from_fn(|i| state[i].clone()));
let b1 = matmul_m4(array::from_fn(|i| state[4 + i].clone()));
let b2 = matmul_m4(array::from_fn(|i| state[8 + i].clone()));
let sums: [E; 4] = array::from_fn(|j| b0[j].clone() + b1[j].clone() + b2[j].clone());
array::from_fn(|i| {
let block = i / 4;
let lane = i % 4;
let b = match block {
0 => &b0,
1 => &b1,
_ => &b2,
};
b[lane].clone() + sums[lane].clone()
})
}
pub fn matmul_m4<E: PrimeCharacteristicRing>(input: [E; 4]) -> [E; 4] {
let [a, b, c, d] = input;
let t01 = a.clone() + b.clone();
let t23 = c.clone() + d.clone();
let t0123 = t01.clone() + t23.clone();
let t01123 = t0123.clone() + b;
let t01233 = t0123 + d;
let out0 = t01123.clone() + t01;
let out1 = t01123 + c.double();
let out2 = t01233.clone() + t23;
let out3 = t01233 + a.double();
[out0, out1, out2, out3]
}
pub fn apply_matmul_internal<E: PrimeCharacteristicRing>(
state: &[E; STATE_WIDTH],
mat_diag: &[E; STATE_WIDTH],
) -> [E; STATE_WIDTH] {
let sum = E::sum_array::<STATE_WIDTH>(state);
array::from_fn(|i| state[i].clone() * mat_diag[i].clone() + sum.clone())
}