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};
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() % 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        // 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#[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    /// The whole point of this backend is that it agrees with the portable one
219    /// exactly. This is the differential test that makes the acceleration safe
220    /// to trust: the portable path is validated against FIPS 197 vectors, and
221    /// this asserts bit-for-bit equality with it.
222    #[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            // SAFETY: guarded by the runtime feature check above.
231            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    /// The batch path must agree with the single-block path, including at the
253    /// boundary where the eight-way loop hands off to the remainder.
254    #[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        // SAFETY: guarded by the runtime feature check above.
261        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        // SAFETY: guarded by the runtime feature check above.
285        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    /// Decryption inverts encryption within the accelerated backend itself.
291    #[test]
292    fn decryption_inverts_encryption() {
293        if !available() {
294            return;
295        }
296        let sched = portable::Schedule::expand(&[0x42u8; 24]).unwrap();
297        // SAFETY: guarded by the runtime feature check above.
298        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}