use super::armv8_rounds::{armv8_decrypt_rounds, armv8_encrypt_rounds};
use super::portable::{Schedule, BLOCK_LEN};
use core::arch::aarch64::*;
use ic_core::{ensure, Result};
pub const PARALLEL_BLOCKS: usize = 4;
#[derive(Clone, Copy)]
pub struct Keys {
enc: [uint8x16_t; 15],
dec: [uint8x16_t; 15],
rounds: usize,
}
impl Keys {
#[target_feature(enable = "neon")]
#[target_feature(enable = "aes")]
pub unsafe fn load(sched: &Schedule) -> Keys {
let rounds = sched.rounds;
let zero = vdupq_n_u8(0);
let mut enc = [zero; 15];
let mut dec = [zero; 15];
for (r, slot) in enc.iter_mut().enumerate().take(rounds + 1) {
let rk = sched.round_key(r);
*slot = unsafe { vld1q_u8(rk.as_ptr()) };
}
dec[0] = enc[rounds];
for i in 1..rounds {
dec[i] = vaesimcq_u8(enc[rounds - i]);
}
dec[rounds] = enc[0];
Keys { enc, dec, rounds }
}
#[target_feature(enable = "neon")]
#[target_feature(enable = "aes")]
#[inline]
unsafe fn encrypt(&self, block: uint8x16_t) -> uint8x16_t {
armv8_encrypt_rounds!(
block,
self.enc,
self.rounds,
vaeseq_u8,
vaesmcq_u8,
veorq_u8
)
}
#[target_feature(enable = "neon")]
#[target_feature(enable = "aes")]
#[inline]
unsafe fn decrypt(&self, block: uint8x16_t) -> uint8x16_t {
armv8_decrypt_rounds!(
block,
self.dec,
self.rounds,
vaesdq_u8,
vaesimcq_u8,
veorq_u8
)
}
}
#[target_feature(enable = "neon")]
#[target_feature(enable = "aes")]
pub unsafe fn encrypt_block(keys: &Keys, block: &mut [u8]) -> Result<()> {
ensure!(block.len() == BLOCK_LEN, InvalidLength, "aes block");
unsafe {
let b = vld1q_u8(block.as_ptr());
let out = keys.encrypt(b);
vst1q_u8(block.as_mut_ptr(), out);
}
Ok(())
}
#[target_feature(enable = "neon")]
#[target_feature(enable = "aes")]
pub unsafe fn decrypt_block(keys: &Keys, block: &mut [u8]) -> Result<()> {
ensure!(block.len() == BLOCK_LEN, InvalidLength, "aes block");
unsafe {
let b = vld1q_u8(block.as_ptr());
let out = keys.decrypt(b);
vst1q_u8(block.as_mut_ptr(), out);
}
Ok(())
}
#[target_feature(enable = "neon")]
#[target_feature(enable = "aes")]
pub unsafe fn encrypt_blocks(keys: &Keys, data: &mut [u8]) -> Result<()> {
ensure!(
data.len() % BLOCK_LEN == 0,
InvalidLength,
"aes block sequence"
);
let mut chunks = data.chunks_exact_mut(BLOCK_LEN * PARALLEL_BLOCKS);
for chunk in chunks.by_ref() {
unsafe {
let mut state = [vdupq_n_u8(0); PARALLEL_BLOCKS];
for (i, slot) in state.iter_mut().enumerate() {
*slot = vld1q_u8(chunk.as_ptr().add(i * BLOCK_LEN));
}
for slot in state.iter_mut() {
*slot = keys.encrypt(*slot);
}
for (i, slot) in state.iter().enumerate() {
vst1q_u8(chunk.as_mut_ptr().add(i * BLOCK_LEN), *slot);
}
}
}
for block in chunks.into_remainder().chunks_exact_mut(BLOCK_LEN) {
unsafe { encrypt_block(keys, block)? };
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::aes::portable;
fn available() -> bool {
std::arch::is_aarch64_feature_detected!("aes")
}
#[test]
fn matches_the_portable_backend_for_every_key_size() {
if !available() {
eprintln!("aarch64 aes extension not available; skipping");
return;
}
for key_len in [16usize, 24, 32] {
let key: Vec<u8> = (0..key_len).map(|i| (i as u8).wrapping_mul(7)).collect();
let sched = portable::Schedule::expand(&key).unwrap();
let keys = unsafe { Keys::load(&sched) };
for seed in 0..16u8 {
let mut want = [0u8; BLOCK_LEN];
for (i, b) in want.iter_mut().enumerate() {
*b = seed.wrapping_mul(31).wrapping_add(i as u8);
}
let mut got = want;
portable::encrypt_block(&sched, &mut want).unwrap();
unsafe { encrypt_block(&keys, &mut got).unwrap() };
assert_eq!(got, want, "encrypt, key_len={key_len} seed={seed}");
unsafe { decrypt_block(&keys, &mut got).unwrap() };
let mut original = [0u8; BLOCK_LEN];
for (i, b) in original.iter_mut().enumerate() {
*b = seed.wrapping_mul(31).wrapping_add(i as u8);
}
assert_eq!(got, original, "decrypt, key_len={key_len} seed={seed}");
}
}
}
#[test]
fn batching_matches_single_blocks() {
if !available() {
eprintln!("aarch64 aes extension not available; skipping");
return;
}
let key = [0x3cu8; 32];
let sched = portable::Schedule::expand(&key).unwrap();
let keys = unsafe { Keys::load(&sched) };
for blocks in [1usize, 3, 4, 5, 9] {
let mut batched: Vec<u8> = (0..blocks * BLOCK_LEN).map(|i| (i % 251) as u8).collect();
let mut single = batched.clone();
unsafe { encrypt_blocks(&keys, &mut batched).unwrap() };
for block in single.chunks_exact_mut(BLOCK_LEN) {
unsafe { encrypt_block(&keys, block).unwrap() };
}
assert_eq!(batched, single, "blocks={blocks}");
}
}
#[test]
fn lengths_are_checked() {
if !available() {
eprintln!("aarch64 aes extension not available; skipping");
return;
}
let sched = portable::Schedule::expand(&[0u8; 16]).unwrap();
let keys = unsafe { Keys::load(&sched) };
let mut short = [0u8; 15];
assert!(unsafe { encrypt_block(&keys, &mut short) }.is_err());
let mut ragged = [0u8; BLOCK_LEN + 1];
assert!(unsafe { encrypt_blocks(&keys, &mut ragged) }.is_err());
}
}