1use super::portable::{Schedule, BLOCK_LEN};
31use ic_core::{ensure, Result, Zeroize};
32
33#[cfg(target_arch = "x86")]
34use core::arch::x86::*;
35#[cfg(target_arch = "x86_64")]
36use core::arch::x86_64::*;
37
38pub const PARALLEL_BLOCKS: usize = 8;
45
46#[derive(Clone, Copy)]
48pub struct Keys {
49 enc: [__m128i; 15],
50 dec: [__m128i; 15],
51 rounds: usize,
52}
53
54impl Keys {
55 #[target_feature(enable = "aes")]
62 pub unsafe fn load(sched: &Schedule) -> Keys {
63 let rounds = sched.rounds;
64 let mut enc = [_mm_setzero_si128(); 15];
65 let mut dec = [_mm_setzero_si128(); 15];
66
67 for (r, slot) in enc.iter_mut().enumerate().take(rounds + 1) {
68 let rk = sched.round_key(r);
69 *slot = unsafe { _mm_loadu_si128(rk.as_ptr() as *const __m128i) };
72 }
73
74 dec[0] = enc[rounds];
77 for i in 1..rounds {
78 dec[i] = _mm_aesimc_si128(enc[rounds - i]);
79 }
80 dec[rounds] = enc[0];
81
82 Keys { enc, dec, rounds }
83 }
84
85 #[target_feature(enable = "aes")]
91 #[inline]
92 unsafe fn encrypt(&self, block: __m128i) -> __m128i {
93 let mut b = _mm_xor_si128(block, self.enc[0]);
94 for r in 1..self.rounds {
95 b = _mm_aesenc_si128(b, self.enc[r]);
96 }
97 _mm_aesenclast_si128(b, self.enc[self.rounds])
98 }
99
100 #[target_feature(enable = "aes")]
106 #[inline]
107 unsafe fn decrypt(&self, block: __m128i) -> __m128i {
108 let mut b = _mm_xor_si128(block, self.dec[0]);
109 for r in 1..self.rounds {
110 b = _mm_aesdec_si128(b, self.dec[r]);
111 }
112 _mm_aesdeclast_si128(b, self.dec[self.rounds])
113 }
114}
115
116#[target_feature(enable = "aes")]
122pub unsafe fn encrypt_block(keys: &Keys, block: &mut [u8]) -> Result<()> {
123 ensure!(block.len() == BLOCK_LEN, InvalidLength, "aes block");
124 unsafe {
127 let b = _mm_loadu_si128(block.as_ptr() as *const __m128i);
128 let out = keys.encrypt(b);
129 _mm_storeu_si128(block.as_mut_ptr() as *mut __m128i, out);
130 }
131 Ok(())
132}
133
134#[target_feature(enable = "aes")]
140pub unsafe fn decrypt_block(keys: &Keys, block: &mut [u8]) -> Result<()> {
141 ensure!(block.len() == BLOCK_LEN, InvalidLength, "aes block");
142 unsafe {
144 let b = _mm_loadu_si128(block.as_ptr() as *const __m128i);
145 let out = keys.decrypt(b);
146 _mm_storeu_si128(block.as_mut_ptr() as *mut __m128i, out);
147 }
148 Ok(())
149}
150
151#[target_feature(enable = "aes")]
157pub unsafe fn encrypt_blocks(keys: &Keys, data: &mut [u8]) -> Result<()> {
158 ensure!(
159 data.len().is_multiple_of(BLOCK_LEN),
160 InvalidLength,
161 "aes batch must be block-aligned"
162 );
163
164 let mut chunks = data.chunks_exact_mut(BLOCK_LEN * PARALLEL_BLOCKS);
165 for chunk in &mut chunks {
166 let p = chunk.as_mut_ptr();
167 let mut b = [_mm_setzero_si128(); PARALLEL_BLOCKS];
168
169 unsafe {
172 for (i, slot) in b.iter_mut().enumerate() {
173 *slot = _mm_loadu_si128(p.add(i * BLOCK_LEN) as *const __m128i);
174 }
175 }
176
177 for slot in b.iter_mut() {
180 *slot = _mm_xor_si128(*slot, keys.enc[0]);
181 }
182 for r in 1..keys.rounds {
183 let rk = keys.enc[r];
184 for slot in b.iter_mut() {
185 *slot = _mm_aesenc_si128(*slot, rk);
186 }
187 }
188 let last = keys.enc[keys.rounds];
189 for slot in b.iter_mut() {
190 *slot = _mm_aesenclast_si128(*slot, last);
191 }
192
193 unsafe {
195 for (i, slot) in b.iter().enumerate() {
196 _mm_storeu_si128(p.add(i * BLOCK_LEN) as *mut __m128i, *slot);
197 }
198 }
199 }
200
201 for block in chunks.into_remainder().chunks_exact_mut(BLOCK_LEN) {
203 unsafe { encrypt_block(keys, block)? };
205 }
206 Ok(())
207}
208
209#[target_feature(enable = "aes")]
229pub unsafe fn ctr32_xor(keys: &Keys, counter: &mut [u8; BLOCK_LEN], data: &mut [u8]) {
230 let mut n = u32::from_be_bytes([counter[12], counter[13], counter[14], counter[15]]);
231 let mut blocks = [[0u8; BLOCK_LEN]; PARALLEL_BLOCKS];
232 for block in blocks.iter_mut() {
233 block[..12].copy_from_slice(&counter[..12]);
234 }
235
236 for chunk in data.chunks_mut(BLOCK_LEN * PARALLEL_BLOCKS) {
237 for (i, block) in blocks.iter_mut().enumerate() {
238 block[12..].copy_from_slice(&n.wrapping_add(i as u32).to_be_bytes());
239 }
240 let mut b = [_mm_setzero_si128(); PARALLEL_BLOCKS];
241 for (slot, block) in b.iter_mut().zip(blocks.iter()) {
242 *slot = _mm_xor_si128(
244 unsafe { _mm_loadu_si128(block.as_ptr() as *const __m128i) },
245 keys.enc[0],
246 );
247 }
248 for r in 1..keys.rounds {
249 let rk = keys.enc[r];
250 for slot in b.iter_mut() {
251 *slot = _mm_aesenc_si128(*slot, rk);
252 }
253 }
254 let last = keys.enc[keys.rounds];
255 for slot in b.iter_mut() {
256 *slot = _mm_aesenclast_si128(*slot, last);
257 }
258
259 let p = chunk.as_mut_ptr();
260 if chunk.len() == BLOCK_LEN * PARALLEL_BLOCKS {
261 for (i, ks) in b.iter().enumerate() {
262 unsafe {
265 let d = _mm_loadu_si128(p.add(i * BLOCK_LEN) as *const __m128i);
266 _mm_storeu_si128(p.add(i * BLOCK_LEN) as *mut __m128i, _mm_xor_si128(d, *ks));
267 }
268 }
269 } else {
270 let mut stream = [0u8; BLOCK_LEN * PARALLEL_BLOCKS];
273 for (i, ks) in b.iter().enumerate() {
274 unsafe {
276 _mm_storeu_si128(stream.as_mut_ptr().add(i * BLOCK_LEN) as *mut __m128i, *ks)
277 };
278 }
279 for (d, k) in chunk.iter_mut().zip(stream.iter()) {
280 *d ^= k;
281 }
282 stream.zeroize();
283 }
284 n = n.wrapping_add(chunk.len().div_ceil(BLOCK_LEN) as u32);
285 }
286 counter[12..].copy_from_slice(&n.to_be_bytes());
287}
288
289#[cfg(test)]
290mod tests {
291 use super::*;
292 use crate::aes::portable;
293
294 fn available() -> bool {
295 std::arch::is_x86_feature_detected!("aes")
296 }
297
298 #[test]
303 fn matches_the_portable_backend_on_every_key_length() {
304 if !available() {
305 return;
306 }
307 for key_len in [16usize, 24, 32] {
308 let key: Vec<u8> = (0..key_len).map(|i| (i * 7 + 1) as u8).collect();
309 let sched = portable::Schedule::expand(&key).unwrap();
310 let keys = unsafe { Keys::load(&sched) };
312
313 for seed in 0..64u8 {
314 let original: [u8; 16] =
315 core::array::from_fn(|i| seed ^ (i as u8).wrapping_mul(31));
316
317 let mut a = original;
318 let mut b = original;
319 portable::encrypt_block(&sched, &mut a).unwrap();
320 unsafe { encrypt_block(&keys, &mut b).unwrap() };
321 assert_eq!(a, b, "encrypt mismatch, key_len {key_len}, seed {seed}");
322
323 let mut a = original;
324 let mut b = original;
325 portable::decrypt_block(&sched, &mut a).unwrap();
326 unsafe { decrypt_block(&keys, &mut b).unwrap() };
327 assert_eq!(a, b, "decrypt mismatch, key_len {key_len}, seed {seed}");
328 }
329 }
330 }
331
332 #[test]
335 fn batch_matches_single_block_at_every_length() {
336 if !available() {
337 return;
338 }
339 let sched = portable::Schedule::expand(&[0x2bu8; 32]).unwrap();
340 let keys = unsafe { Keys::load(&sched) };
342
343 for blocks in 0..=(PARALLEL_BLOCKS * 2 + 3) {
344 let data: Vec<u8> = (0..blocks * BLOCK_LEN).map(|i| (i * 13) as u8).collect();
345
346 let mut batched = data.clone();
347 unsafe { encrypt_blocks(&keys, &mut batched).unwrap() };
348
349 let mut one_at_a_time = data.clone();
350 for block in one_at_a_time.chunks_exact_mut(BLOCK_LEN) {
351 unsafe { encrypt_block(&keys, block).unwrap() };
352 }
353
354 assert_eq!(batched, one_at_a_time, "{blocks} blocks");
355 }
356 }
357
358 #[test]
359 fn batch_rejects_unaligned_input() {
360 if !available() {
361 return;
362 }
363 let sched = portable::Schedule::expand(&[0u8; 16]).unwrap();
364 let keys = unsafe { Keys::load(&sched) };
366 let mut data = [0u8; BLOCK_LEN + 1];
367 assert!(unsafe { encrypt_blocks(&keys, &mut data) }.is_err());
368 }
369
370 #[test]
372 fn decryption_inverts_encryption() {
373 if !available() {
374 return;
375 }
376 let sched = portable::Schedule::expand(&[0x42u8; 24]).unwrap();
377 let keys = unsafe { Keys::load(&sched) };
379 let original: [u8; 16] = core::array::from_fn(|i| i as u8);
380 let mut block = original;
381 unsafe {
382 encrypt_block(&keys, &mut block).unwrap();
383 assert_ne!(block, original);
384 decrypt_block(&keys, &mut block).unwrap();
385 }
386 assert_eq!(block, original);
387 }
388}