1use super::portable::{Schedule, BLOCK_LEN};
31use ic_core::{ensure, Result};
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() % BLOCK_LEN == 0,
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#[cfg(test)]
210mod tests {
211 use super::*;
212 use crate::aes::portable;
213
214 fn available() -> bool {
215 std::arch::is_x86_feature_detected!("aes")
216 }
217
218 #[test]
223 fn matches_the_portable_backend_on_every_key_length() {
224 if !available() {
225 return;
226 }
227 for key_len in [16usize, 24, 32] {
228 let key: Vec<u8> = (0..key_len).map(|i| (i * 7 + 1) as u8).collect();
229 let sched = portable::Schedule::expand(&key).unwrap();
230 let keys = unsafe { Keys::load(&sched) };
232
233 for seed in 0..64u8 {
234 let original: [u8; 16] =
235 core::array::from_fn(|i| seed ^ (i as u8).wrapping_mul(31));
236
237 let mut a = original;
238 let mut b = original;
239 portable::encrypt_block(&sched, &mut a).unwrap();
240 unsafe { encrypt_block(&keys, &mut b).unwrap() };
241 assert_eq!(a, b, "encrypt mismatch, key_len {key_len}, seed {seed}");
242
243 let mut a = original;
244 let mut b = original;
245 portable::decrypt_block(&sched, &mut a).unwrap();
246 unsafe { decrypt_block(&keys, &mut b).unwrap() };
247 assert_eq!(a, b, "decrypt mismatch, key_len {key_len}, seed {seed}");
248 }
249 }
250 }
251
252 #[test]
255 fn batch_matches_single_block_at_every_length() {
256 if !available() {
257 return;
258 }
259 let sched = portable::Schedule::expand(&[0x2bu8; 32]).unwrap();
260 let keys = unsafe { Keys::load(&sched) };
262
263 for blocks in 0..=(PARALLEL_BLOCKS * 2 + 3) {
264 let data: Vec<u8> = (0..blocks * BLOCK_LEN).map(|i| (i * 13) as u8).collect();
265
266 let mut batched = data.clone();
267 unsafe { encrypt_blocks(&keys, &mut batched).unwrap() };
268
269 let mut one_at_a_time = data.clone();
270 for block in one_at_a_time.chunks_exact_mut(BLOCK_LEN) {
271 unsafe { encrypt_block(&keys, block).unwrap() };
272 }
273
274 assert_eq!(batched, one_at_a_time, "{blocks} blocks");
275 }
276 }
277
278 #[test]
279 fn batch_rejects_unaligned_input() {
280 if !available() {
281 return;
282 }
283 let sched = portable::Schedule::expand(&[0u8; 16]).unwrap();
284 let keys = unsafe { Keys::load(&sched) };
286 let mut data = [0u8; BLOCK_LEN + 1];
287 assert!(unsafe { encrypt_blocks(&keys, &mut data) }.is_err());
288 }
289
290 #[test]
292 fn decryption_inverts_encryption() {
293 if !available() {
294 return;
295 }
296 let sched = portable::Schedule::expand(&[0x42u8; 24]).unwrap();
297 let keys = unsafe { Keys::load(&sched) };
299 let original: [u8; 16] = core::array::from_fn(|i| i as u8);
300 let mut block = original;
301 unsafe {
302 encrypt_block(&keys, &mut block).unwrap();
303 assert_ne!(block, original);
304 decrypt_block(&keys, &mut block).unwrap();
305 }
306 assert_eq!(block, original);
307 }
308}