Skip to main content

ic_cipher/aes/
x86.rs

1//! The x86-64 AES-NI backend.
2//!
3//! # What is and is not reimplemented here
4//!
5//! Only the *round function* is accelerated. Key expansion stays in
6//! [`portable::Schedule`][super::portable::Schedule], and this backend loads
7//! its output into SIMD registers. Expansion happens once per key and is not on
8//! the hot path, so a second implementation of it would buy nothing and risk a
9//! divergence — particularly for AES-192, whose SIMD key schedule is the
10//! fiddliest part of a typical AES-NI implementation.
11//!
12//! Decryption uses the equivalent inverse cipher: the encryption round keys are
13//! passed through `AESIMC` and reversed at construction, which is what lets
14//! `AESDEC` run the inverse rounds in the same shape as the forward ones.
15//!
16//! # Safety
17//!
18//! These functions are `unsafe` because `#[target_feature]` requires it: a
19//! caller must not invoke them on a CPU without AES-NI. The intrinsics
20//! themselves are safe once that feature is enabled in scope, so the only
21//! genuinely unsafe operations are the raw-pointer loads and stores — and those
22//! are the only things wrapped in an `unsafe` block.
23//!
24//! # Constant-time properties
25//!
26//! `AESENC` and friends are single instructions with data-independent latency,
27//! so this backend is constant-time for the same reason the portable one is —
28//! and it never touches a lookup table at all.
29
30use 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
38/// How many blocks the batch path processes at once.
39///
40/// `AESENC` has a latency of around four cycles but is fully pipelined, so
41/// interleaving eight independent blocks keeps the unit busy instead of
42/// stalling on each result. This is where most of the speedup over the portable
43/// backend comes from in CTR and GCM.
44pub const PARALLEL_BLOCKS: usize = 8;
45
46/// AES round keys held in SIMD registers.
47#[derive(Clone, Copy)]
48pub struct Keys {
49    enc: [__m128i; 15],
50    dec: [__m128i; 15],
51    rounds: usize,
52}
53
54impl Keys {
55    /// Load an expanded schedule into SIMD registers.
56    ///
57    /// # Safety
58    ///
59    /// The caller must ensure the `aes` and `sse2` target features are
60    /// available on this CPU.
61    #[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            // SAFETY: `round_key` always returns exactly BLOCK_LEN bytes, and
70            // the unaligned load has no alignment requirement.
71            *slot = unsafe { _mm_loadu_si128(rk.as_ptr() as *const __m128i) };
72        }
73
74        // Equivalent inverse cipher: reverse the round keys and apply AESIMC to
75        // everything except the first and last.
76        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    /// Encrypt one block in registers.
86    ///
87    /// # Safety
88    ///
89    /// Requires the `aes` target feature.
90    #[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    /// Decrypt one block in registers.
101    ///
102    /// # Safety
103    ///
104    /// Requires the `aes` target feature.
105    #[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/// Encrypt one block in place.
117///
118/// # Safety
119///
120/// Requires the `aes` target feature.
121#[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    // SAFETY: `block` is exactly BLOCK_LEN bytes, checked above; the loads and
125    // stores are unaligned and so have no alignment requirement.
126    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/// Decrypt one block in place.
135///
136/// # Safety
137///
138/// Requires the `aes` target feature.
139#[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    // SAFETY: as above.
143    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/// Encrypt a whole number of blocks in place, eight at a time.
152///
153/// # Safety
154///
155/// Requires the `aes` target feature.
156#[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        // SAFETY: the chunk is exactly PARALLEL_BLOCKS * BLOCK_LEN bytes, so
170        // every offset below is in bounds, and the loads are unaligned.
171        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        // Interleave the rounds across all eight blocks so the pipeline stays
178        // full rather than waiting on each AESENC.
179        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        // SAFETY: same bounds as the loads above.
194        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    // Whatever is left over is fewer than PARALLEL_BLOCKS blocks.
202    for block in chunks.into_remainder().chunks_exact_mut(BLOCK_LEN) {
203        // SAFETY: the caller established the `aes` feature for this call.
204        unsafe { encrypt_block(keys, block)? };
205    }
206    Ok(())
207}
208
209/// XOR `data` with the CTR keystream starting at `counter`, advancing the
210/// counter's last 32 bits big-endian and wrapping within them -- GCM's
211/// `inc32` -- eight blocks at a time.
212///
213/// GCM used to reach AES through `encrypt_blocks` eight blocks per call, with
214/// the counters built and the keystream XORed byte by byte around each call:
215/// 8.5 GiB/s where one call over the whole buffer runs at 15. Here the
216/// counters are written straight into the blocks, the keystream never leaves
217/// registers on the whole-block path, and there is one call per message.
218///
219/// The counter blocks are assembled with plain stores rather than a byte
220/// shuffle, so this needs nothing beyond the `aes` feature the caller already
221/// checked; building eight 16-byte blocks is noise beside eighty `AESENC`s.
222///
223/// On return `counter` is the block after the last one used.
224///
225/// # Safety
226///
227/// Requires the `aes` target feature.
228#[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            // SAFETY: `block` is exactly 16 bytes and the load is unaligned.
243            *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                // SAFETY: the chunk is exactly PARALLEL_BLOCKS blocks, so every
263                // offset is in bounds; loads and stores are unaligned.
264                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            // The final, shorter chunk: through a buffer, and only as far as
271            // the data goes. The keystream past the end is discarded.
272            let mut stream = [0u8; BLOCK_LEN * PARALLEL_BLOCKS];
273            for (i, ks) in b.iter().enumerate() {
274                // SAFETY: `stream` holds PARALLEL_BLOCKS blocks.
275                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    /// The whole point of this backend is that it agrees with the portable one
299    /// exactly. This is the differential test that makes the acceleration safe
300    /// to trust: the portable path is validated against FIPS 197 vectors, and
301    /// this asserts bit-for-bit equality with it.
302    #[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            // SAFETY: guarded by the runtime feature check above.
311            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    /// The batch path must agree with the single-block path, including at the
333    /// boundary where the eight-way loop hands off to the remainder.
334    #[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        // SAFETY: guarded by the runtime feature check above.
341        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        // SAFETY: guarded by the runtime feature check above.
365        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    /// Decryption inverts encryption within the accelerated backend itself.
371    #[test]
372    fn decryption_inverts_encryption() {
373        if !available() {
374            return;
375        }
376        let sched = portable::Schedule::expand(&[0x42u8; 24]).unwrap();
377        // SAFETY: guarded by the runtime feature check above.
378        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}