const SBOX: [u8; 256] = [
0x63, 0x7c, 0x77, 0x7b, 0xf2, 0x6b, 0x6f, 0xc5, 0x30, 0x01, 0x67, 0x2b, 0xfe, 0xd7, 0xab, 0x76,
0xca, 0x82, 0xc9, 0x7d, 0xfa, 0x59, 0x47, 0xf0, 0xad, 0xd4, 0xa2, 0xaf, 0x9c, 0xa4, 0x72, 0xc0,
0xb7, 0xfd, 0x93, 0x26, 0x36, 0x3f, 0xf7, 0xcc, 0x34, 0xa5, 0xe5, 0xf1, 0x71, 0xd8, 0x31, 0x15,
0x04, 0xc7, 0x23, 0xc3, 0x18, 0x96, 0x05, 0x9a, 0x07, 0x12, 0x80, 0xe2, 0xeb, 0x27, 0xb2, 0x75,
0x09, 0x83, 0x2c, 0x1a, 0x1b, 0x6e, 0x5a, 0xa0, 0x52, 0x3b, 0xd6, 0xb3, 0x29, 0xe3, 0x2f, 0x84,
0x53, 0xd1, 0x00, 0xed, 0x20, 0xfc, 0xb1, 0x5b, 0x6a, 0xcb, 0xbe, 0x39, 0x4a, 0x4c, 0x58, 0xcf,
0xd0, 0xef, 0xaa, 0xfb, 0x43, 0x4d, 0x33, 0x85, 0x45, 0xf9, 0x02, 0x7f, 0x50, 0x3c, 0x9f, 0xa8,
0x51, 0xa3, 0x40, 0x8f, 0x92, 0x9d, 0x38, 0xf5, 0xbc, 0xb6, 0xda, 0x21, 0x10, 0xff, 0xf3, 0xd2,
0xcd, 0x0c, 0x13, 0xec, 0x5f, 0x97, 0x44, 0x17, 0xc4, 0xa7, 0x7e, 0x3d, 0x64, 0x5d, 0x19, 0x73,
0x60, 0x81, 0x4f, 0xdc, 0x22, 0x2a, 0x90, 0x88, 0x46, 0xee, 0xb8, 0x14, 0xde, 0x5e, 0x0b, 0xdb,
0xe0, 0x32, 0x3a, 0x0a, 0x49, 0x06, 0x24, 0x5c, 0xc2, 0xd3, 0xac, 0x62, 0x91, 0x95, 0xe4, 0x79,
0xe7, 0xc8, 0x37, 0x6d, 0x8d, 0xd5, 0x4e, 0xa9, 0x6c, 0x56, 0xf4, 0xea, 0x65, 0x7a, 0xae, 0x08,
0xba, 0x78, 0x25, 0x2e, 0x1c, 0xa6, 0xb4, 0xc6, 0xe8, 0xdd, 0x74, 0x1f, 0x4b, 0xbd, 0x8b, 0x8a,
0x70, 0x3e, 0xb5, 0x66, 0x48, 0x03, 0xf6, 0x0e, 0x61, 0x35, 0x57, 0xb9, 0x86, 0xc1, 0x1d, 0x9e,
0xe1, 0xf8, 0x98, 0x11, 0x69, 0xd9, 0x8e, 0x94, 0x9b, 0x1e, 0x87, 0xe9, 0xce, 0x55, 0x28, 0xdf,
0x8c, 0xa1, 0x89, 0x0d, 0xbf, 0xe6, 0x42, 0x68, 0x41, 0x99, 0x2d, 0x0f, 0xb0, 0x54, 0xbb, 0x16,
];
const INV_SBOX: [u8; 256] = {
let mut inv = [0u8; 256];
let mut x = 0usize;
while x < 256 {
let mut i = 0usize;
loop {
if SBOX[i] as usize == x {
inv[x] = i as u8;
break;
}
i += 1;
}
x += 1;
}
inv
};
fn xtime(x: u8) -> u8 {
let hi = ((x >> 7) & 1).wrapping_neg();
(x << 1) ^ (hi & 0x1b)
}
fn gf_mul(mut a: u8, mut b: u8) -> u8 {
let mut p = 0u8;
for _ in 0..8 {
p ^= a & ((b & 1).wrapping_neg());
let hi = ((a >> 7) & 1).wrapping_neg();
a = (a << 1) ^ (hi & 0x1b);
b >>= 1;
}
p
}
#[inline]
fn ct_table_lookup(table: &[u8; 256], x: u8) -> u8 {
let mut acc = 0u8;
for (i, &entry) in table.iter().enumerate() {
let eq = (((i as u8) ^ x) == 0) as u8;
acc |= entry & eq.wrapping_neg();
}
acc
}
#[inline]
fn sbox(x: u8) -> u8 {
ct_table_lookup(&SBOX, x)
}
#[inline]
#[cfg(test)]
fn inv_sbox(x: u8) -> u8 {
ct_table_lookup(&INV_SBOX, x)
}
#[inline]
fn sub_bytes(s: &mut [u8; 16]) {
let mut acc = [0u8; 16];
for (i, &entry) in SBOX.iter().enumerate() {
let idx = i as u8;
for (a, &x) in acc.iter_mut().zip(s.iter()) {
let eq = ((x == idx) as u8).wrapping_neg();
*a |= entry & eq;
}
}
*s = acc;
}
#[inline]
fn inv_sub_bytes(s: &mut [u8; 16]) {
let mut acc = [0u8; 16];
for (i, &entry) in INV_SBOX.iter().enumerate() {
let idx = i as u8;
for (a, &x) in acc.iter_mut().zip(s.iter()) {
let eq = ((x == idx) as u8).wrapping_neg();
*a |= entry & eq;
}
}
*s = acc;
}
fn sub_word(w: [u8; 4]) -> [u8; 4] {
[sbox(w[0]), sbox(w[1]), sbox(w[2]), sbox(w[3])]
}
macro_rules! aes_impl {
($name:ident, $nk:expr, $nr:expr, $doc:expr) => {
#[doc = $doc]
#[derive(Clone)]
pub struct $name {
/// 轮密钥(Nr+1 × 16 字节,按 FIPS-197 列序展开)。
rk: Vec<u8>,
rk_planes: Vec<[Planes; 16]>,
}
impl $name {
pub const KEY_LEN: usize = $nk * 4;
const NR: usize = $nr;
pub fn new(key: &[u8; $nk * 4]) -> Self {
let total = 16 * ($nr + 1);
let mut rk = vec![0u8; total];
let nk_bytes = $nk * 4;
rk[..nk_bytes].copy_from_slice(key);
let mut rcon = 1u8;
let mut i = nk_bytes;
while i < total {
let mut t: [u8; 4] = rk[i - 4..i].try_into().unwrap();
if i % nk_bytes == 0 {
t = sub_word([t[1], t[2], t[3], t[0]]);
t[0] ^= rcon;
rcon = xtime(rcon);
} else if $nk > 6 && i % nk_bytes == 16 {
t = sub_word(t);
}
for j in 0..4 {
rk[i + j] = rk[i - nk_bytes + j] ^ t[j];
}
i += 4;
}
let mut rk_planes = Vec::with_capacity($nr + 1);
for round in 0..=$nr {
let mut planes = [[0u64; 8]; 16];
for g in 0..16 {
for (b, plane) in planes[g].iter_mut().enumerate() {
*plane = u64::from((rk[round * 16 + g] >> b) & 1).wrapping_neg();
}
}
rk_planes.push(planes);
}
Self { rk, rk_planes }
}
fn add_round_key(&self, state: &mut [u8; 16], round: usize) {
for j in 0..16 {
state[j] ^= self.rk[round * 16 + j];
}
}
pub fn encrypt_block(&self, block: &mut [u8; 16]) {
let mut s = *block;
self.add_round_key(&mut s, 0);
for round in 1..Self::NR {
sub_bytes(&mut s);
shift_rows(&mut s);
mix_columns(&mut s);
self.add_round_key(&mut s, round);
}
sub_bytes(&mut s);
shift_rows(&mut s);
self.add_round_key(&mut s, Self::NR);
*block = s;
}
pub fn decrypt_block(&self, block: &mut [u8; 16]) {
let mut s = *block;
self.add_round_key(&mut s, Self::NR);
for round in (1..Self::NR).rev() {
inv_shift_rows(&mut s);
inv_sub_bytes(&mut s);
self.add_round_key(&mut s, round);
inv_mix_columns(&mut s);
}
inv_shift_rows(&mut s);
inv_sub_bytes(&mut s);
self.add_round_key(&mut s, 0);
*block = s;
}
#[allow(dead_code)]
pub(crate) fn encrypt_ctr_batch(&self, base: [u8; 16], n: usize, out: &mut [u8]) {
debug_assert!(n > 0 && n <= CTR_BATCH_BLOCKS && out.len() >= n * 16);
#[cfg(feature = "simd")]
{
if n <= 64 {
encrypt_ctr_batch_p::<u64>(&self.rk_planes, base, n, out)
} else if n <= 128 {
encrypt_ctr_batch_p::<Simd<u64, 2>>(&self.rk_planes, base, n, out)
} else if n <= 256 {
encrypt_ctr_batch_p::<Simd<u64, 4>>(&self.rk_planes, base, n, out)
} else {
encrypt_ctr_batch_p::<Simd<u64, 8>>(&self.rk_planes, base, n, out)
}
}
#[cfg(not(feature = "simd"))]
encrypt_ctr_batch_p::<u64>(&self.rk_planes, base, n, out);
}
}
impl Drop for $name {
fn drop(&mut self) {
self.rk.fill(0);
self.rk_planes.fill([[0u64; 8]; 16]);
}
}
impl std::fmt::Debug for $name {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(stringify!($name))
}
}
};
}
use core::ops::{BitAnd, BitAndAssign, BitOr, BitOrAssign, BitXor, BitXorAssign, Not};
#[cfg(feature = "simd")]
use std::simd::Simd;
trait Plane:
Copy
+ BitAnd<Output = Self>
+ BitOr<Output = Self>
+ BitXor<Output = Self>
+ Not<Output = Self>
+ BitAndAssign
+ BitOrAssign
+ BitXorAssign
{
const LANES: usize;
fn broadcast(x: u64) -> Self;
fn from_lane_groups(f: impl FnMut(usize) -> u64) -> Self;
fn to_lane_group(self, e: usize) -> u64;
}
impl Plane for u64 {
const LANES: usize = 64;
#[inline]
fn broadcast(x: u64) -> Self {
x
}
#[inline]
fn from_lane_groups(mut f: impl FnMut(usize) -> u64) -> Self {
f(0)
}
#[inline]
fn to_lane_group(self, e: usize) -> u64 {
debug_assert_eq!(e, 0);
self
}
}
#[cfg(feature = "simd")]
impl<const L: usize> Plane for Simd<u64, L> {
const LANES: usize = 64 * L;
#[inline]
fn broadcast(x: u64) -> Self {
Simd::splat(x)
}
#[inline]
fn from_lane_groups(f: impl FnMut(usize) -> u64) -> Self {
Simd::from_array(core::array::from_fn(f))
}
#[inline]
fn to_lane_group(self, e: usize) -> u64 {
self.to_array()[e]
}
}
type Planes = [u64; 8];
#[cfg(feature = "simd")]
pub(crate) const CTR_BATCH_BLOCKS: usize = 512;
#[cfg(not(feature = "simd"))]
pub(crate) const CTR_BATCH_BLOCKS: usize = 64;
fn bs_mul<P: Plane>(a: &[P; 8], b: &[P; 8]) -> [P; 8] {
let mut t = [P::broadcast(0); 15];
for i in 0..8 {
for j in 0..8 {
t[i + j] ^= a[i] & b[j];
}
}
[
t[0] ^ t[8] ^ t[12] ^ t[13],
t[1] ^ t[8] ^ t[9] ^ t[12] ^ t[14],
t[2] ^ t[9] ^ t[10] ^ t[13],
t[3] ^ t[8] ^ t[10] ^ t[11] ^ t[12] ^ t[13] ^ t[14],
t[4] ^ t[8] ^ t[9] ^ t[11] ^ t[14],
t[5] ^ t[9] ^ t[10] ^ t[12],
t[6] ^ t[10] ^ t[11] ^ t[13],
t[7] ^ t[11] ^ t[12] ^ t[14],
]
}
fn bs_sq<P: Plane>(a: &[P; 8]) -> [P; 8] {
[
a[0] ^ a[4] ^ a[6],
a[4] ^ a[6] ^ a[7],
a[1] ^ a[5],
a[4] ^ a[5] ^ a[6] ^ a[7],
a[2] ^ a[4] ^ a[7],
a[5] ^ a[6],
a[3] ^ a[5],
a[6] ^ a[7],
]
}
fn bs_xtime<P: Plane>(a: &[P; 8]) -> [P; 8] {
[
a[7],
a[0] ^ a[7],
a[1],
a[2] ^ a[7],
a[3] ^ a[7],
a[4],
a[5],
a[6],
]
}
fn bs_sbox<P: Plane>(x: &mut [P; 8]) {
let a = *x;
let x2 = bs_sq(&a);
let x3 = bs_mul(&a, &x2); let x6 = bs_sq(&x3);
let x12 = bs_sq(&x6);
let x24 = bs_sq(&x12);
let x48 = bs_sq(&x24);
let x96 = bs_sq(&x48);
let x192 = bs_sq(&x96);
let x4 = bs_sq(&x2);
let x7 = bs_mul(&x3, &x4); let x14 = bs_sq(&x7);
let t = bs_mul(&x192, &x48);
let inv = bs_mul(&t, &x14); for i in 0..8 {
let mut s =
inv[i] ^ inv[(i + 4) % 8] ^ inv[(i + 5) % 8] ^ inv[(i + 6) % 8] ^ inv[(i + 7) % 8];
if (0x63 >> i) & 1 == 1 {
s = !s; }
x[i] = s;
}
}
fn bs_rounds<P: Plane>(st: &mut [[P; 8]; 16], rk_planes: &[[Planes; 16]]) {
for rk in &rk_planes[1..rk_planes.len() - 1] {
for group in st.iter_mut() {
bs_sbox(group);
}
bs_shift_rows(st);
bs_mix_columns(st);
for (sg, rg) in st.iter_mut().zip(rk.iter()) {
for (s, r) in sg.iter_mut().zip(rg.iter()) {
*s ^= P::broadcast(*r);
}
}
}
for group in st.iter_mut() {
bs_sbox(group);
}
bs_shift_rows(st);
let rk = &rk_planes[rk_planes.len() - 1];
for (sg, rg) in st.iter_mut().zip(rk.iter()) {
for (s, r) in sg.iter_mut().zip(rg.iter()) {
*s ^= P::broadcast(*r);
}
}
}
fn bs_shift_rows<P: Plane>(s: &mut [[P; 8]; 16]) {
let t = *s;
for row in 1..4 {
for col in 0..4 {
s[4 * col + row] = t[4 * ((col + row) % 4) + row];
}
}
}
fn bs_mix_columns<P: Plane>(s: &mut [[P; 8]; 16]) {
for c in 0..4 {
let o = 4 * c;
let (a0, a1, a2, a3) = (s[o], s[o + 1], s[o + 2], s[o + 3]);
let xt0 = bs_xtime(&a0);
let xt1 = bs_xtime(&a1);
let xt2 = bs_xtime(&a2);
let xt3 = bs_xtime(&a3);
for b in 0..8 {
s[o][b] = xt0[b] ^ xt1[b] ^ a1[b] ^ a2[b] ^ a3[b];
s[o + 1][b] = a0[b] ^ xt1[b] ^ xt2[b] ^ a2[b] ^ a3[b];
s[o + 2][b] = a0[b] ^ a1[b] ^ xt2[b] ^ xt3[b] ^ a3[b];
s[o + 3][b] = xt0[b] ^ a0[b] ^ a1[b] ^ a2[b] ^ xt3[b];
}
}
}
fn ctr_group_planes(ctr0: u32, e: usize, n: usize) -> [[u64; 8]; 16] {
let mut grp = [[0u64; 8]; 16];
for i in 0..64usize {
let lane = e * 64 + i;
if lane < n {
let ctr = ctr0.wrapping_add(lane as u32).to_be_bytes();
for k in 0..4 {
for (b, slot) in grp[12 + k].iter_mut().enumerate() {
*slot |= u64::from((ctr[k] >> b) & 1) << i;
}
}
}
}
grp
}
fn encrypt_ctr_batch_p<P: Plane>(
rk_planes: &[[Planes; 16]],
base: [u8; 16],
n: usize,
out: &mut [u8],
) {
debug_assert!(n > 0 && n <= P::LANES && out.len() >= n * 16);
let mut st = [[P::broadcast(0); 8]; 16];
for (g, byte) in base[..12].iter().enumerate() {
for (b, plane) in st[g].iter_mut().enumerate() {
*plane = P::broadcast(u64::from((byte >> b) & 1).wrapping_neg());
}
}
let ctr0 = u32::from_be_bytes(base[12..16].try_into().unwrap());
let ng = P::LANES / 64;
let mut groups = [[[0u64; 8]; 16]; 8]; for (e, slot) in groups.iter_mut().enumerate().take(ng) {
*slot = ctr_group_planes(ctr0, e, n);
}
for g in 12..16 {
for (b, slot) in st[g].iter_mut().enumerate() {
*slot = P::from_lane_groups(|e| groups[e][g][b]);
}
}
let rk0 = &rk_planes[0];
for (sg, rg) in st.iter_mut().zip(rk0.iter()) {
for (s, r) in sg.iter_mut().zip(rg.iter()) {
*s ^= P::broadcast(*r);
}
}
bs_rounds(&mut st, rk_planes);
for e in 0..ng {
let grp = st.map(|g8| {
let mut grp8 = [0u64; 8];
for (slot, p) in grp8.iter_mut().zip(g8.iter()) {
*slot = p.to_lane_group(e);
}
grp8
});
for i in 0..64usize {
let lane = e * 64 + i;
if lane >= n {
continue; }
for (g, group) in grp.iter().enumerate() {
let mut byte = 0u8;
for (b, plane) in group.iter().enumerate() {
byte |= (((plane >> i) & 1) as u8) << b;
}
out[lane * 16 + g] = byte;
}
}
}
}
fn shift_rows(s: &mut [u8; 16]) {
let t = *s;
for row in 1..4 {
for col in 0..4 {
s[4 * col + row] = t[4 * ((col + row) % 4) + row];
}
}
}
fn inv_shift_rows(s: &mut [u8; 16]) {
let t = *s;
for row in 1..4 {
for col in 0..4 {
s[4 * ((col + row) % 4) + row] = t[4 * col + row];
}
}
}
fn mix_columns(s: &mut [u8; 16]) {
for c in 0..4 {
let o = 4 * c;
let (a0, a1, a2, a3) = (s[o], s[o + 1], s[o + 2], s[o + 3]);
s[o] = xtime(a0) ^ xtime(a1) ^ a1 ^ a2 ^ a3;
s[o + 1] = a0 ^ xtime(a1) ^ xtime(a2) ^ a2 ^ a3;
s[o + 2] = a0 ^ a1 ^ xtime(a2) ^ xtime(a3) ^ a3;
s[o + 3] = xtime(a0) ^ a0 ^ a1 ^ a2 ^ xtime(a3);
}
}
fn inv_mix_columns(s: &mut [u8; 16]) {
for c in 0..4 {
let o = 4 * c;
let (a0, a1, a2, a3) = (s[o], s[o + 1], s[o + 2], s[o + 3]);
s[o] = gf_mul(a0, 14) ^ gf_mul(a1, 11) ^ gf_mul(a2, 13) ^ gf_mul(a3, 9);
s[o + 1] = gf_mul(a0, 9) ^ gf_mul(a1, 14) ^ gf_mul(a2, 11) ^ gf_mul(a3, 13);
s[o + 2] = gf_mul(a0, 13) ^ gf_mul(a1, 9) ^ gf_mul(a2, 14) ^ gf_mul(a3, 11);
s[o + 3] = gf_mul(a0, 11) ^ gf_mul(a1, 13) ^ gf_mul(a2, 9) ^ gf_mul(a3, 14);
}
}
aes_impl!(Aes128, 4, 10, "AES-128 块密码实例。");
aes_impl!(Aes192, 6, 12, "AES-192 块密码实例。");
aes_impl!(Aes256, 8, 14, "AES-256 块密码实例。");
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sbox_known_values_and_bijection() {
assert_eq!(sbox(0x00), 0x63);
assert_eq!(sbox(0x01), 0x7c);
assert_eq!(sbox(0x53), 0xed);
assert_eq!(sbox(0xff), 0x16);
for (x, &official) in SBOX.iter().enumerate() {
assert_eq!(sbox(x as u8), official, "sbox({x:#04x})");
}
let mut seen = [false; 256];
for x in 0..=255u8 {
assert_eq!(inv_sbox(sbox(x)), x, "round trip at {x}");
seen[sbox(x) as usize] = true;
}
assert!(seen.iter().all(|&s| s), "sbox must be a bijection");
}
#[test]
fn fips197_appendix_c_kats() {
let mut key = [0u8; 16];
for (i, b) in key.iter_mut().enumerate() {
*b = i as u8;
}
let aes = Aes128::new(&key);
let mut block = [
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd,
0xee, 0xff,
];
aes.encrypt_block(&mut block);
assert_eq!(
block,
[
0x69, 0xc4, 0xe0, 0xd8, 0x6a, 0x7b, 0x04, 0x30, 0xd8, 0xcd, 0xb7, 0x80, 0x70, 0xb4,
0xc5, 0x5a
]
);
aes.decrypt_block(&mut block);
assert_eq!(
block,
[
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd,
0xee, 0xff
]
);
let mut key = [0u8; 32];
for (i, b) in key.iter_mut().enumerate() {
*b = i as u8;
}
let aes = Aes256::new(&key);
let mut block = [
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd,
0xee, 0xff,
];
aes.encrypt_block(&mut block);
assert_eq!(
block,
[
0x8e, 0xa2, 0xb7, 0xca, 0x51, 0x67, 0x45, 0xbf, 0xea, 0xfc, 0x49, 0x90, 0x4b, 0x49,
0x60, 0x89
]
);
aes.decrypt_block(&mut block);
assert_eq!(
block,
[
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd,
0xee, 0xff,
]
);
}
#[test]
fn bitslice_circuits_match_scalar_exhaustive() {
let planes_of = |v: u8| -> Planes {
let mut p = [0u64; 8];
for (b, plane) in p.iter_mut().enumerate() {
*plane = u64::from((v >> b) & 1).wrapping_neg();
}
p
};
let byte_of = |p: &Planes| -> u8 {
let mut v = 0u8;
for (b, plane) in p.iter().enumerate() {
v |= ((plane & 1) as u8) << b;
}
v
};
for v in 0..=255u8 {
let a = planes_of(v);
assert_eq!(byte_of(&bs_sq(&a)), gf_mul(v, v), "sq({v:#04x})");
assert_eq!(byte_of(&bs_xtime(&a)), xtime(v), "xtime({v:#04x})");
let mut s = a;
bs_sbox(&mut s);
assert_eq!(byte_of(&s), sbox(v), "sbox({v:#04x})");
}
for x in 0..=255u8 {
for y in 0..=255u8 {
assert_eq!(
byte_of(&bs_mul(&planes_of(x), &planes_of(y))),
gf_mul(x, y),
"mul({x:#04x},{y:#04x})"
);
}
}
let mut lanes = [0u8; 64];
let mut seed = 0x9E37_79B9u32;
for v in lanes.iter_mut() {
seed ^= seed << 13;
seed ^= seed >> 17;
seed ^= seed << 5;
*v = seed as u8;
}
let mut group = [0u64; 8];
for (lane, &v) in lanes.iter().enumerate() {
for (b, plane) in group.iter_mut().enumerate() {
*plane |= u64::from((v >> b) & 1) << lane;
}
}
bs_sbox(&mut group);
for (lane, &v) in lanes.iter().enumerate() {
let mut got = 0u8;
for (b, plane) in group.iter().enumerate() {
got |= (((plane >> lane) & 1) as u8) << b;
}
assert_eq!(got, sbox(v), "lane {lane}");
}
}
#[test]
fn sub_bytes_matches_scalar() {
let mut seed = 0x243F_6A88u32;
let mut next = || {
seed ^= seed << 13;
seed ^= seed >> 17;
seed ^= seed << 5;
seed
};
let mut states = vec![[0u8; 16], [0xff; 16], [0x53; 16]];
for _ in 0..64 {
states.push(core::array::from_fn(|_| next() as u8));
}
for st in &states {
let mut fwd = *st;
sub_bytes(&mut fwd);
let mut inv = *st;
inv_sub_bytes(&mut inv);
for g in 0..16 {
assert_eq!(fwd[g], sbox(st[g]), "sub_bytes state={st:?} g={g}");
assert_eq!(inv[g], inv_sbox(st[g]), "inv_sub_bytes state={st:?} g={g}");
}
}
}
#[test]
fn encrypt_ctr_batch_matches_scalar() {
let mut key = [0u8; 16];
for (i, b) in key.iter_mut().enumerate() {
*b = i as u8;
}
let aes = Aes128::new(&key);
let mut key256 = [0u8; 32];
for (i, b) in key256.iter_mut().enumerate() {
*b = (i * 7) as u8;
}
let aes256 = Aes256::new(&key256);
let mut base = [
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0x00, 0x00,
0x00, 0x01,
];
#[cfg(feature = "simd")]
let sizes = [
1usize, 2, 3, 63, 64, 65, 100, 128, 129, 200, 256, 257, 300, 511, 512,
];
#[cfg(not(feature = "simd"))]
let sizes = [1usize, 2, 3, 63, 64];
for n in sizes {
let mut fast = vec![0u8; CTR_BATCH_BLOCKS * 16];
aes.encrypt_ctr_batch(base, n, &mut fast);
let mut expect = vec![0u8; n * 16];
let ctr0 = u32::from_be_bytes(base[12..16].try_into().unwrap());
for i in 0..n {
let mut blk = base;
blk[12..16].copy_from_slice(&ctr0.wrapping_add(i as u32).to_be_bytes());
aes.encrypt_block(&mut blk);
expect[i * 16..(i + 1) * 16].copy_from_slice(&blk);
}
assert_eq!(&fast[..n * 16], &expect, "n={n} aes128");
let mut expect256 = vec![0u8; n * 16];
for i in 0..n {
let mut blk = base;
blk[12..16].copy_from_slice(&ctr0.wrapping_add(i as u32).to_be_bytes());
aes256.encrypt_block(&mut blk);
expect256[i * 16..(i + 1) * 16].copy_from_slice(&blk);
}
aes256.encrypt_ctr_batch(base, n, &mut fast);
assert_eq!(&fast[..n * 16], &expect256, "n={n} aes256");
}
base[12..16].copy_from_slice(&0xFFFF_FFFDu32.to_be_bytes());
let mut fast = vec![0u8; CTR_BATCH_BLOCKS * 16];
aes.encrypt_ctr_batch(base, 64, &mut fast);
let ctr0 = u32::from_be_bytes(base[12..16].try_into().unwrap());
for i in 0..64usize {
let mut blk = base;
blk[12..16].copy_from_slice(&ctr0.wrapping_add(i as u32).to_be_bytes());
aes.encrypt_block(&mut blk);
assert_eq!(&fast[i * 16..(i + 1) * 16], &blk, "wrap i={i}");
}
#[cfg(feature = "simd")]
{
let mut base = [
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xff, 0xff,
0xff, 0x00,
];
let mut fast = vec![0u8; 512 * 16];
aes.encrypt_ctr_batch(base, 512, &mut fast);
let ctr0 = u32::from_be_bytes(base[12..16].try_into().unwrap());
for i in 0..512usize {
let mut blk = base;
blk[12..16].copy_from_slice(&ctr0.wrapping_add(i as u32).to_be_bytes());
aes.encrypt_block(&mut blk);
assert_eq!(&fast[i * 16..(i + 1) * 16], &blk, "wrap512 i={i}");
}
base[12..16].copy_from_slice(&0xFFFF_FFFDu32.to_be_bytes());
aes.encrypt_ctr_batch(base, 129, &mut fast);
let ctr0 = u32::from_be_bytes(base[12..16].try_into().unwrap());
for i in 0..129usize {
let mut blk = base;
blk[12..16].copy_from_slice(&ctr0.wrapping_add(i as u32).to_be_bytes());
aes.encrypt_block(&mut blk);
assert_eq!(&fast[i * 16..(i + 1) * 16], &blk, "wrap129 i={i}");
}
}
}
}