use super::portable::{Schedule, BLOCK_LEN};
pub const LANES: usize = 4;
pub const GROUP: usize = LANES * BLOCK_LEN;
type Planes = [u64; 8];
const SQUARE_TERMS: [&[usize]; 8] = [
&[0, 4, 6],
&[4, 6, 7],
&[1, 5],
&[4, 5, 6, 7],
&[2, 4, 7],
&[5, 6],
&[3, 5],
&[6, 7],
];
#[inline(always)]
fn transpose8(mut x: u64) -> u64 {
x = (x & 0xAA55_AA55_AA55_AA55)
| ((x & 0x00AA_00AA_00AA_00AA) << 7)
| ((x >> 7) & 0x00AA_00AA_00AA_00AA);
x = (x & 0xCCCC_3333_CCCC_3333)
| ((x & 0x0000_CCCC_0000_CCCC) << 14)
| ((x >> 14) & 0x0000_CCCC_0000_CCCC);
x = (x & 0xF0F0_F0F0_0F0F_0F0F)
| ((x & 0x0000_0000_F0F0_F0F0) << 28)
| ((x >> 28) & 0x0000_0000_F0F0_F0F0);
x
}
fn transpose_in(bytes: &[u8]) -> Planes {
let mut p = [0u64; 8];
for (w, chunk) in bytes.chunks_exact(8).enumerate() {
let mut word = [0u8; 8];
word.copy_from_slice(chunk);
let t = transpose8(u64::from_le_bytes(word));
for (i, plane) in p.iter_mut().enumerate() {
*plane |= ((t >> (8 * i)) & 0xff) << (8 * w);
}
}
p
}
fn transpose_out(p: &Planes, out: &mut [u8]) {
for (w, chunk) in out.chunks_exact_mut(8).enumerate() {
let mut t = 0u64;
for (i, plane) in p.iter().enumerate() {
t |= ((plane >> (8 * w)) & 0xff) << (8 * i);
}
chunk.copy_from_slice(&transpose8(t).to_le_bytes());
}
}
fn square(a: &Planes) -> Planes {
let mut out = [0u64; 8];
for (i, slot) in out.iter_mut().enumerate() {
let mut v = 0u64;
for &j in SQUARE_TERMS[i] {
v ^= a[j];
}
*slot = v;
}
out
}
fn mul(a: &Planes, b: &Planes) -> Planes {
let mut t = [0u64; 15];
for i in 0..8 {
for j in 0..8 {
t[i + j] ^= a[i] & b[j];
}
}
let mut k = 14;
while k >= 8 {
let v = t[k];
t[k - 4] ^= v;
t[k - 5] ^= v;
t[k - 7] ^= v;
t[k - 8] ^= v;
k -= 1;
}
let mut out = [0u64; 8];
out.copy_from_slice(&t[..8]);
out
}
fn inv(a: &Planes) -> Planes {
let mut r = *a;
let mut bit = 6i32;
while bit >= 0 {
r = square(&r);
if bit > 0 {
r = mul(&r, a);
}
bit -= 1;
}
r
}
fn sbox(a: &Planes) -> Planes {
let y = inv(a);
let mut out = [0u64; 8];
for (i, slot) in out.iter_mut().enumerate() {
*slot = y[i] ^ y[(i + 7) % 8] ^ y[(i + 6) % 8] ^ y[(i + 5) % 8] ^ y[(i + 4) % 8];
}
for i in [0, 1, 5, 6] {
out[i] = !out[i];
}
out
}
fn xtime(a: &Planes) -> Planes {
[
a[7],
a[0] ^ a[7],
a[1],
a[2] ^ a[7],
a[3] ^ a[7],
a[4],
a[5],
a[6],
]
}
const NIBBLE_LOW: u64 = 0x1111_1111_1111_1111;
fn rotate_column(v: u64) -> u64 {
((v >> 1) & 0x7777_7777_7777_7777) | ((v & NIBBLE_LOW) << 3)
}
fn mix_columns(a: &Planes) -> Planes {
let r1 = a.map(rotate_column);
let r2 = r1.map(rotate_column);
let r3 = r2.map(rotate_column);
let xa = xtime(a);
let xr1 = xtime(&r1);
let mut out = [0u64; 8];
for (i, slot) in out.iter_mut().enumerate() {
*slot = xa[i] ^ xr1[i] ^ r1[i] ^ r2[i] ^ r3[i];
}
out
}
fn shift_rows(a: &Planes) -> Planes {
let mut out = [0u64; 8];
for (slot, &v) in out.iter_mut().zip(a.iter()) {
let mut acc = v & NIBBLE_LOW;
for r in 1..4u32 {
let row = v & (NIBBLE_LOW << r);
let s = 4 * r;
let m = (1u64 << s) - 1;
let low_mask = m | (m << 16) | (m << 32) | (m << 48);
let lo = row & low_mask;
let hi = row & !low_mask;
acc |= (hi >> s) | (lo << (16 - s));
}
*slot = acc;
}
out
}
pub struct RoundKeys {
planes: [Planes; 15],
rounds: usize,
}
impl RoundKeys {
pub fn new(sched: &Schedule) -> Self {
let mut planes = [[0u64; 8]; 15];
for (r, slot) in planes.iter_mut().enumerate().take(sched.rounds + 1) {
let rk = sched.round_key(r);
let mut wide = [0u8; GROUP];
for lane in 0..LANES {
wide[lane * BLOCK_LEN..(lane + 1) * BLOCK_LEN].copy_from_slice(rk);
}
*slot = transpose_in(&wide);
}
Self {
planes,
rounds: sched.rounds,
}
}
}
pub fn encrypt_group(keys: &RoundKeys, data: &mut [u8]) {
debug_assert_eq!(data.len(), GROUP);
let mut s = transpose_in(data);
for (slot, k) in s.iter_mut().zip(keys.planes[0].iter()) {
*slot ^= k;
}
for r in 1..keys.rounds {
s = sbox(&s);
s = shift_rows(&s);
s = mix_columns(&s);
for (slot, k) in s.iter_mut().zip(keys.planes[r].iter()) {
*slot ^= k;
}
}
s = sbox(&s);
s = shift_rows(&s);
for (slot, k) in s.iter_mut().zip(keys.planes[keys.rounds].iter()) {
*slot ^= k;
}
transpose_out(&s, data);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::gf;
fn through<F: Fn(&Planes) -> Planes>(vals: &[u8; GROUP], f: F) -> [u8; GROUP] {
let out_planes = f(&transpose_in(vals));
let mut out = [0u8; GROUP];
transpose_out(&out_planes, &mut out);
out
}
#[test]
fn the_byte_transpose_matches_a_naive_one() {
fn naive(x: u64) -> u64 {
let mut r = 0u64;
for j in 0..8 {
for i in 0..8 {
if (x >> (8 * j + i)) & 1 == 1 {
r |= 1 << (8 * i + j);
}
}
}
r
}
for b in 0..64 {
let v = 1u64 << b;
assert_eq!(transpose8(v), naive(v), "single bit {b}");
}
let mut state = 0x9e37_79b9_7f4a_7c15u64;
for _ in 0..20_000 {
state ^= state >> 12;
state ^= state << 25;
state ^= state >> 27;
let v = state.wrapping_mul(0x2545_f491_4f6c_dd1d);
assert_eq!(transpose8(v), naive(v), "word {v:#018x}");
}
}
#[test]
fn transposing_round_trips() {
let mut v = [0u8; GROUP];
for (i, slot) in v.iter_mut().enumerate() {
*slot = (i as u8).wrapping_mul(7).wrapping_add(3);
}
assert_eq!(through(&v, |p| *p), v);
for byte in 0..GROUP {
for bit in 0..8 {
let mut one = [0u8; GROUP];
one[byte] = 1 << bit;
assert_eq!(through(&one, |p| *p), one, "byte {byte} bit {bit}");
}
}
}
#[test]
fn squaring_matches_the_byte_at_a_time_path() {
for base in (0..=255u16).step_by(GROUP) {
let mut vals = [0u8; GROUP];
for (i, slot) in vals.iter_mut().enumerate() {
*slot = (base as usize + i).min(255) as u8;
}
let got = through(&vals, square);
for (i, &v) in vals.iter().enumerate() {
assert_eq!(got[i], gf::mul(v, v), "square({v:#04x})");
}
}
}
#[test]
fn field_multiply_matches_the_byte_at_a_time_one() {
for a in 0..=255u8 {
let a_vals = [a; GROUP];
let a_planes = transpose_in(&a_vals);
for chunk in 0..(256 / GROUP) {
let mut b_vals = [0u8; GROUP];
for (i, slot) in b_vals.iter_mut().enumerate() {
*slot = (chunk * GROUP + i) as u8;
}
let planes = mul(&a_planes, &transpose_in(&b_vals));
let mut got = [0u8; GROUP];
transpose_out(&planes, &mut got);
for (i, &b) in b_vals.iter().enumerate() {
assert_eq!(got[i], gf::mul(a, b), "{a:#04x} * {b:#04x}");
}
}
}
}
#[test]
fn sbox_matches_the_byte_at_a_time_path() {
for chunk in 0..(256 / GROUP) {
let mut vals = [0u8; GROUP];
for (i, slot) in vals.iter_mut().enumerate() {
*slot = (chunk * GROUP + i) as u8;
}
let got = through(&vals, sbox);
for (i, &v) in vals.iter().enumerate() {
assert_eq!(got[i], gf::sbox(v), "sbox({v:#04x})");
}
}
}
#[test]
fn four_blocks_match_the_byte_at_a_time_path() {
for key_len in [16usize, 24, 32] {
let key: Vec<u8> = (0..key_len).map(|i| (i * 11 + 5) as u8).collect();
let sched = Schedule::expand(&key).unwrap();
let keys = RoundKeys::new(&sched);
for case in 0..64u32 {
let mut data = [0u8; GROUP];
for (i, slot) in data.iter_mut().enumerate() {
*slot = (i as u32)
.wrapping_mul(case.wrapping_add(1))
.wrapping_add(case) as u8;
}
let mut want = data;
for block in want.chunks_exact_mut(BLOCK_LEN) {
super::super::portable::encrypt_block(&sched, block).unwrap();
}
let mut got = data;
encrypt_group(&keys, &mut got);
assert_eq!(got, want, "key_len {key_len}, case {case}");
}
}
}
#[test]
fn lanes_do_not_leak_into_each_other() {
let key = [0x42u8; 32];
let sched = Schedule::expand(&key).unwrap();
let keys = RoundKeys::new(&sched);
for lane in 0..LANES {
let mut data = [0u8; GROUP];
for k in 0..BLOCK_LEN {
data[lane * BLOCK_LEN + k] = (k as u8).wrapping_mul(37).wrapping_add(1);
}
let mut want = data;
for block in want.chunks_exact_mut(BLOCK_LEN) {
super::super::portable::encrypt_block(&sched, block).unwrap();
}
let mut got = data;
encrypt_group(&keys, &mut got);
assert_eq!(got, want, "lane {lane}");
}
}
}