use std::array::from_fn;
use crate::shortint::server_key::LookupTable;
use crate::shortint::{Ciphertext, ServerKey};
use rayon::prelude::*;
use super::key::AesFheRoundKeys;
use super::sbox::sbox;
pub(super) fn xor_assign(sks: &ServerKey, lhs: &mut [Ciphertext], rhs: &[Ciphertext]) {
for (s, k) in lhs.iter_mut().zip(rhs) {
sks.unchecked_add_assign(s, k);
}
}
pub(super) fn shift_rows(state: &mut [Ciphertext; 128]) {
let (bytes, rest) = state.as_chunks_mut::<8>();
debug_assert!(rest.is_empty());
bytes.swap(1, 5);
bytes.swap(5, 9);
bytes.swap(9, 13);
bytes.swap(2, 10);
bytes.swap(6, 14);
bytes.swap(15, 11);
bytes.swap(11, 7);
bytes.swap(7, 3);
}
fn xtime(
sks: &ServerKey,
flush_lut: &LookupTable<Vec<u64>>,
bits: &[Ciphertext; 8],
) -> [Ciphertext; 8] {
let msb = &bits[7];
let mut result = [
msb.clone(),
bits[0].clone(),
bits[1].clone(),
bits[2].clone(),
bits[3].clone(),
bits[4].clone(),
bits[5].clone(),
bits[6].clone(),
];
let [_, r1, _, r3, r4, _, _, _] = &mut result;
[r1, r3, r4].par_iter_mut().for_each(|r| {
sks.unchecked_add_assign(r, msb);
sks.apply_lookup_table_assign(r, flush_lut);
});
result
}
pub(super) fn mix_columns_op(
sks: &ServerKey,
flush_lut: &LookupTable<Vec<u64>>,
input: &[Ciphertext; 32],
) -> [Ciphertext; 32] {
let (chunks, rest) = input.as_chunks::<8>();
debug_assert!(rest.is_empty());
let [c0, c1, c2, c3]: &[[Ciphertext; 8]; 4] = chunks.try_into().unwrap();
let [mut b0, mut b1, mut b2, mut b3] = [c0, c1, c2, c3]
.par_iter()
.map(|c| xtime(sks, flush_lut, c))
.collect::<Vec<_>>()
.try_into()
.unwrap();
let b0_copy = b0.clone();
xor_assign(sks, &mut b0, &b1);
xor_assign(sks, &mut b0, c1);
xor_assign(sks, &mut b0, c2);
xor_assign(sks, &mut b0, c3);
xor_assign(sks, &mut b1, c0);
xor_assign(sks, &mut b1, &b2);
xor_assign(sks, &mut b1, c2);
xor_assign(sks, &mut b1, c3);
xor_assign(sks, &mut b2, c0);
xor_assign(sks, &mut b2, c1);
xor_assign(sks, &mut b2, &b3);
xor_assign(sks, &mut b2, c3);
xor_assign(sks, &mut b3, c0);
xor_assign(sks, &mut b3, c1);
xor_assign(sks, &mut b3, c2);
xor_assign(sks, &mut b3, &b0_copy);
let mut iter = b0.into_iter().chain(b1).chain(b2).chain(b3);
from_fn(|_| iter.next().unwrap())
}
fn mix_columns(sks: &ServerKey, flush_lut: &LookupTable<Vec<u64>>, state: &mut [Ciphertext]) {
let (chunks, _) = state.as_chunks_mut::<32>();
chunks
.par_iter_mut()
.for_each(|col: &mut [Ciphertext; 32]| {
let mix_col = mix_columns_op(sks, flush_lut, col);
col.clone_from_slice(&mix_col);
});
}
fn sub_bytes(sks: &ServerKey, state: &mut [Ciphertext; 128], flush_lut: &LookupTable<Vec<u64>>) {
state
.par_chunks_mut(8)
.for_each(|chunk| sbox(sks, flush_lut, chunk));
}
fn flush_state(sks: &ServerKey, state: &mut [Ciphertext], flush_lut: &LookupTable<Vec<u64>>) {
state.par_iter_mut().for_each(|b| {
sks.apply_lookup_table_assign(b, flush_lut);
});
}
pub(super) fn encrypt_block(
sks: &ServerKey,
counter_value: u128,
key: &AesFheRoundKeys,
) -> [Ciphertext; 128] {
let bytes = counter_value.to_be_bytes();
let mut state: [Ciphertext; 128] =
from_fn(|i| sks.create_trivial(((bytes[i / 8] >> (i % 8)) & 1) as u64));
let round_keys = key.round_keys();
let flush_lut = key.flush_lut();
xor_assign(sks, &mut state, &round_keys[0].key);
for erk in &round_keys[1..10] {
sub_bytes(sks, &mut state, flush_lut);
flush_state(sks, &mut state, flush_lut);
shift_rows(&mut state);
mix_columns(sks, flush_lut, &mut state);
flush_state(sks, &mut state, flush_lut);
xor_assign(sks, &mut state, &erk.key);
flush_state(sks, &mut state, flush_lut);
}
sub_bytes(sks, &mut state, flush_lut);
shift_rows(&mut state);
xor_assign(sks, &mut state, &round_keys[10].key);
flush_state(sks, &mut state, flush_lut);
state
}