Skip to main content

ic_cipher/aes/
bitslice.rs

1//! A bitsliced AES encryption path, four blocks at a time.
2//!
3//! # Why
4//!
5//! [`super::portable`] computes one S-box at a time, and an S-box is an
6//! inversion in GF(2^8): seven squarings and six multiplications, with each
7//! multiplication a bit-serial loop. That is a few hundred operations per byte,
8//! and it is why the portable backend ran some twenty-five times behind
9//! RustCrypto's software AES, which is fixsliced.
10//!
11//! Bitslicing pays the same algebra once for sixty-four bytes instead of once
12//! per byte. The state is held transposed: eight `u64` planes, where plane `i`
13//! bit `j` is bit `i` of byte `j`. A field multiplication is then sixty-four
14//! `AND`s and a reduction on whole words, and each of those words carries four
15//! blocks' worth of work.
16//!
17//! # Constant time
18//!
19//! Nothing here indexes memory with a value derived from the key or the
20//! plaintext, and nothing branches on one. The S-box is computed, as it is in
21//! the byte-at-a-time path -- that property is the reason both exist rather
22//! than a lookup table.
23//!
24//! # What it does not do
25//!
26//! Decryption. The inverse S-box and inverse MixColumns are a separate piece of
27//! work, and the modes that move volume -- CTR, GCM, GCM-SIV -- only ever run
28//! the forward direction. Decryption stays on the byte-at-a-time path.
29//!
30//! # Trusting it
31//!
32//! Every piece below was derived mechanically from the byte-at-a-time code
33//! rather than transcribed, and each is checked against it: the field
34//! operations over their entire domain, the round operations differentially.
35
36use super::portable::{Schedule, BLOCK_LEN};
37
38/// Blocks processed together. Four blocks of sixteen bytes fill a `u64` plane.
39pub const LANES: usize = 4;
40
41/// Bytes in a group.
42pub const GROUP: usize = LANES * BLOCK_LEN;
43
44/// The state, transposed: `planes[i]` bit `j` is bit `i` of byte `j`.
45type Planes = [u64; 8];
46
47/// Which input bits each output bit of a squaring draws from.
48///
49/// Squaring is linear over GF(2), so it is a fixed 8x8 matrix. This one was
50/// produced by squaring each basis element with the byte-at-a-time multiply and
51/// reading off the result, not copied from a reference.
52const SQUARE_TERMS: [&[usize]; 8] = [
53    &[0, 4, 6],
54    &[4, 6, 7],
55    &[1, 5],
56    &[4, 5, 6, 7],
57    &[2, 4, 7],
58    &[5, 6],
59    &[3, 5],
60    &[6, 7],
61];
62
63/// Transpose the 8x8 bit matrix packed in a `u64`, sending bit `8j + i` to bit
64/// `8i + j`.
65///
66/// Three masked swaps rather than sixty-four bit tests. Getting the bytes into
67/// and out of plane form is pure overhead -- it computes nothing -- so it is
68/// worth not doing it a bit at a time: the naive form was about a tenth of the
69/// whole encryption.
70///
71/// `the_byte_transpose_matches_a_naive_one` checks it against the obvious
72/// double loop, on every single-bit input and on random words.
73#[inline(always)]
74fn transpose8(mut x: u64) -> u64 {
75    x = (x & 0xAA55_AA55_AA55_AA55)
76        | ((x & 0x00AA_00AA_00AA_00AA) << 7)
77        | ((x >> 7) & 0x00AA_00AA_00AA_00AA);
78    x = (x & 0xCCCC_3333_CCCC_3333)
79        | ((x & 0x0000_CCCC_0000_CCCC) << 14)
80        | ((x >> 14) & 0x0000_CCCC_0000_CCCC);
81    x = (x & 0xF0F0_F0F0_0F0F_0F0F)
82        | ((x & 0x0000_0000_F0F0_F0F0) << 28)
83        | ((x >> 28) & 0x0000_0000_F0F0_F0F0);
84    x
85}
86
87/// Transpose sixty-four bytes into eight bit-planes.
88///
89/// Eight bytes at a time: one `transpose8` turns eight bytes into eight bytes
90/// where the `i`th holds bit `i` of each, which is one byte of each plane.
91fn transpose_in(bytes: &[u8]) -> Planes {
92    let mut p = [0u64; 8];
93    for (w, chunk) in bytes.chunks_exact(8).enumerate() {
94        let mut word = [0u8; 8];
95        word.copy_from_slice(chunk);
96        let t = transpose8(u64::from_le_bytes(word));
97        for (i, plane) in p.iter_mut().enumerate() {
98            *plane |= ((t >> (8 * i)) & 0xff) << (8 * w);
99        }
100    }
101    p
102}
103
104/// Transpose eight bit-planes back into sixty-four bytes.
105fn transpose_out(p: &Planes, out: &mut [u8]) {
106    for (w, chunk) in out.chunks_exact_mut(8).enumerate() {
107        let mut t = 0u64;
108        for (i, plane) in p.iter().enumerate() {
109            t |= ((plane >> (8 * w)) & 0xff) << (8 * i);
110        }
111        chunk.copy_from_slice(&transpose8(t).to_le_bytes());
112    }
113}
114
115/// `x^2` in GF(2^8), on planes.
116fn square(a: &Planes) -> Planes {
117    let mut out = [0u64; 8];
118    for (i, slot) in out.iter_mut().enumerate() {
119        let mut v = 0u64;
120        for &j in SQUARE_TERMS[i] {
121            v ^= a[j];
122        }
123        *slot = v;
124    }
125    out
126}
127
128/// `a * b` in GF(2^8), on planes.
129///
130/// Schoolbook into fifteen coefficients, then reduced with
131/// `x^8 = x^4 + x^3 + x + 1`, which sends `x^k` to
132/// `x^(k-4) + x^(k-5) + x^(k-7) + x^(k-8)`. Taking `k` downwards means a
133/// coefficient that lands at or above eight is reduced in its own turn.
134fn mul(a: &Planes, b: &Planes) -> Planes {
135    let mut t = [0u64; 15];
136    for i in 0..8 {
137        for j in 0..8 {
138            t[i + j] ^= a[i] & b[j];
139        }
140    }
141    let mut k = 14;
142    while k >= 8 {
143        let v = t[k];
144        t[k - 4] ^= v;
145        t[k - 5] ^= v;
146        t[k - 7] ^= v;
147        t[k - 8] ^= v;
148        k -= 1;
149    }
150    let mut out = [0u64; 8];
151    out.copy_from_slice(&t[..8]);
152    out
153}
154
155/// `x^254`, which is the inverse for non-zero `x` and zero for zero.
156fn inv(a: &Planes) -> Planes {
157    let mut r = *a;
158    let mut bit = 6i32;
159    while bit >= 0 {
160        r = square(&r);
161        if bit > 0 {
162            r = mul(&r, a);
163        }
164        bit -= 1;
165    }
166    r
167}
168
169/// The AES forward S-box.
170///
171/// The affine step is `y ^ rotl(y,1) ^ rotl(y,2) ^ rotl(y,3) ^ rotl(y,4) ^
172/// 0x63`. Rotating a byte permutes its bits, so on planes it is a rotation of
173/// the plane *indices* and costs nothing but the xors. The constant is a plane
174/// of all ones wherever its bit is set, which is a complement.
175fn sbox(a: &Planes) -> Planes {
176    let y = inv(a);
177    let mut out = [0u64; 8];
178    for (i, slot) in out.iter_mut().enumerate() {
179        *slot = y[i] ^ y[(i + 7) % 8] ^ y[(i + 6) % 8] ^ y[(i + 5) % 8] ^ y[(i + 4) % 8];
180    }
181    // 0x63 = 0b0110_0011.
182    for i in [0, 1, 5, 6] {
183        out[i] = !out[i];
184    }
185    out
186}
187
188/// `x * 2` in GF(2^8), on planes: a shift of the plane indices, with the
189/// overflow folded back in through `0x1b`.
190fn xtime(a: &Planes) -> Planes {
191    [
192        a[7],
193        a[0] ^ a[7],
194        a[1],
195        a[2] ^ a[7],
196        a[3] ^ a[7],
197        a[4],
198        a[5],
199        a[6],
200    ]
201}
202
203/// Low bit of each nibble; a nibble is one four-byte AES column.
204const NIBBLE_LOW: u64 = 0x1111_1111_1111_1111;
205
206/// Rotate the bytes of each column by one, so position `r` takes what was at
207/// `r + 1`.
208///
209/// A byte is one bit in a plane and a column is four consecutive bytes, so a
210/// column is a nibble and this is a nibble-wise rotation.
211fn rotate_column(v: u64) -> u64 {
212    ((v >> 1) & 0x7777_7777_7777_7777) | ((v & NIBBLE_LOW) << 3)
213}
214
215/// MixColumns.
216///
217/// Written as `xtime(a) ^ xtime(R a) ^ R a ^ R^2 a ^ R^3 a`, where `R` is
218/// [`rotate_column`]. Those are the same four output expressions the
219/// byte-at-a-time version has, with the position within the column folded into
220/// `R` so all four are computed at once.
221fn mix_columns(a: &Planes) -> Planes {
222    let r1 = a.map(rotate_column);
223    let r2 = r1.map(rotate_column);
224    let r3 = r2.map(rotate_column);
225    let xa = xtime(a);
226    let xr1 = xtime(&r1);
227    let mut out = [0u64; 8];
228    for (i, slot) in out.iter_mut().enumerate() {
229        *slot = xa[i] ^ xr1[i] ^ r1[i] ^ r2[i] ^ r3[i];
230    }
231    out
232}
233
234/// ShiftRows.
235///
236/// Row `r` of the column-major state occupies the byte positions congruent to
237/// `r` modulo four, and rotates towards lower column indices by `r`. A byte is
238/// one bit and a block is sixteen bits, so that is a rotation by `4r` within
239/// each sixteen-bit group, applied to the bits that row owns.
240fn shift_rows(a: &Planes) -> Planes {
241    let mut out = [0u64; 8];
242    for (slot, &v) in out.iter_mut().zip(a.iter()) {
243        // Row 0 does not move.
244        let mut acc = v & NIBBLE_LOW;
245        for r in 1..4u32 {
246            let row = v & (NIBBLE_LOW << r);
247            let s = 4 * r;
248            // A bit whose position within its sixteen-bit group is below `s`
249            // wraps to the top of that group; the rest simply move down.
250            let m = (1u64 << s) - 1;
251            let low_mask = m | (m << 16) | (m << 32) | (m << 48);
252            let lo = row & low_mask;
253            let hi = row & !low_mask;
254            acc |= (hi >> s) | (lo << (16 - s));
255        }
256        *slot = acc;
257    }
258    out
259}
260
261/// The round keys, transposed once so the round loop does not transpose them.
262pub struct RoundKeys {
263    planes: [Planes; 15],
264    rounds: usize,
265}
266
267impl RoundKeys {
268    /// Transpose every round key of `sched`.
269    ///
270    /// A round key is the same sixteen bytes for all four lanes, so it is
271    /// repeated across the group before transposing. Done once per call rather
272    /// than once per group, which is why it is worth doing at all.
273    pub fn new(sched: &Schedule) -> Self {
274        let mut planes = [[0u64; 8]; 15];
275        for (r, slot) in planes.iter_mut().enumerate().take(sched.rounds + 1) {
276            let rk = sched.round_key(r);
277            let mut wide = [0u8; GROUP];
278            for lane in 0..LANES {
279                wide[lane * BLOCK_LEN..(lane + 1) * BLOCK_LEN].copy_from_slice(rk);
280            }
281            *slot = transpose_in(&wide);
282        }
283        Self {
284            planes,
285            rounds: sched.rounds,
286        }
287    }
288}
289
290/// Encrypt exactly [`GROUP`] bytes in place.
291pub fn encrypt_group(keys: &RoundKeys, data: &mut [u8]) {
292    debug_assert_eq!(data.len(), GROUP);
293    let mut s = transpose_in(data);
294
295    for (slot, k) in s.iter_mut().zip(keys.planes[0].iter()) {
296        *slot ^= k;
297    }
298    for r in 1..keys.rounds {
299        s = sbox(&s);
300        s = shift_rows(&s);
301        s = mix_columns(&s);
302        for (slot, k) in s.iter_mut().zip(keys.planes[r].iter()) {
303            *slot ^= k;
304        }
305    }
306    s = sbox(&s);
307    s = shift_rows(&s);
308    for (slot, k) in s.iter_mut().zip(keys.planes[keys.rounds].iter()) {
309        *slot ^= k;
310    }
311
312    transpose_out(&s, data);
313}
314
315#[cfg(test)]
316mod tests {
317    use super::*;
318    use crate::gf;
319
320    /// Pack 64 byte values into planes, run `f`, and read the bytes back.
321    fn through<F: Fn(&Planes) -> Planes>(vals: &[u8; GROUP], f: F) -> [u8; GROUP] {
322        let out_planes = f(&transpose_in(vals));
323        let mut out = [0u8; GROUP];
324        transpose_out(&out_planes, &mut out);
325        out
326    }
327
328    /// `transpose8` against the obvious double loop.
329    ///
330    /// The masked-swap form is the one place here where the code does not look
331    /// like what it computes, so it is checked against a version that does.
332    #[test]
333    fn the_byte_transpose_matches_a_naive_one() {
334        fn naive(x: u64) -> u64 {
335            let mut r = 0u64;
336            for j in 0..8 {
337                for i in 0..8 {
338                    if (x >> (8 * j + i)) & 1 == 1 {
339                        r |= 1 << (8 * i + j);
340                    }
341                }
342            }
343            r
344        }
345        for b in 0..64 {
346            let v = 1u64 << b;
347            assert_eq!(transpose8(v), naive(v), "single bit {b}");
348        }
349        let mut state = 0x9e37_79b9_7f4a_7c15u64;
350        for _ in 0..20_000 {
351            state ^= state >> 12;
352            state ^= state << 25;
353            state ^= state >> 27;
354            let v = state.wrapping_mul(0x2545_f491_4f6c_dd1d);
355            assert_eq!(transpose8(v), naive(v), "word {v:#018x}");
356        }
357    }
358
359    /// The transpose is its own inverse, for every byte pattern that matters.
360    ///
361    /// Everything below reads its answer back through `transpose_out`, so a
362    /// transpose that lost or moved a bit would make the other tests agree
363    /// about the wrong thing.
364    #[test]
365    fn transposing_round_trips() {
366        let mut v = [0u8; GROUP];
367        for (i, slot) in v.iter_mut().enumerate() {
368            *slot = (i as u8).wrapping_mul(7).wrapping_add(3);
369        }
370        assert_eq!(through(&v, |p| *p), v);
371
372        // One bit set at a time, across every byte and every bit, so a
373        // transpose that swapped two positions cannot hide behind a pattern.
374        for byte in 0..GROUP {
375            for bit in 0..8 {
376                let mut one = [0u8; GROUP];
377                one[byte] = 1 << bit;
378                assert_eq!(through(&one, |p| *p), one, "byte {byte} bit {bit}");
379            }
380        }
381    }
382
383    /// Bitsliced squaring against the byte-at-a-time multiply, over the whole
384    /// domain.
385    #[test]
386    fn squaring_matches_the_byte_at_a_time_path() {
387        for base in (0..=255u16).step_by(GROUP) {
388            let mut vals = [0u8; GROUP];
389            for (i, slot) in vals.iter_mut().enumerate() {
390                *slot = (base as usize + i).min(255) as u8;
391            }
392            let got = through(&vals, square);
393            for (i, &v) in vals.iter().enumerate() {
394                assert_eq!(got[i], gf::mul(v, v), "square({v:#04x})");
395            }
396        }
397    }
398
399    /// Bitsliced multiplication against the byte-at-a-time one, over all
400    /// 65536 pairs.
401    ///
402    /// Not a sample. GF(2^8) has 256 elements, so this is every pair of inputs
403    /// the routine can ever be given, checked against the implementation the
404    /// FIPS-197 vectors already validate.
405    #[test]
406    fn field_multiply_matches_the_byte_at_a_time_one() {
407        for a in 0..=255u8 {
408            let a_vals = [a; GROUP];
409            let a_planes = transpose_in(&a_vals);
410            for chunk in 0..(256 / GROUP) {
411                let mut b_vals = [0u8; GROUP];
412                for (i, slot) in b_vals.iter_mut().enumerate() {
413                    *slot = (chunk * GROUP + i) as u8;
414                }
415                let planes = mul(&a_planes, &transpose_in(&b_vals));
416                let mut got = [0u8; GROUP];
417                transpose_out(&planes, &mut got);
418                for (i, &b) in b_vals.iter().enumerate() {
419                    assert_eq!(got[i], gf::mul(a, b), "{a:#04x} * {b:#04x}");
420                }
421            }
422        }
423    }
424
425    /// The S-box, over all 256 inputs.
426    #[test]
427    fn sbox_matches_the_byte_at_a_time_path() {
428        for chunk in 0..(256 / GROUP) {
429            let mut vals = [0u8; GROUP];
430            for (i, slot) in vals.iter_mut().enumerate() {
431                *slot = (chunk * GROUP + i) as u8;
432            }
433            let got = through(&vals, sbox);
434            for (i, &v) in vals.iter().enumerate() {
435                assert_eq!(got[i], gf::sbox(v), "sbox({v:#04x})");
436            }
437        }
438    }
439
440    /// Four blocks through the bitsliced path must equal four blocks through
441    /// the byte-at-a-time one.
442    ///
443    /// This is the test that matters: ShiftRows and MixColumns are not exposed
444    /// separately, and a rotation applied to the wrong axis would still produce
445    /// a permutation, still round-trip through the transpose, and still look
446    /// like AES from the outside. Only agreement with the implementation the
447    /// published vectors validate rules that out.
448    ///
449    /// Every key length, and byte patterns rather than one fixed buffer, since
450    /// the lanes must stay independent: a bug that mixed lane 1 into lane 2
451    /// would be invisible if all four lanes held the same block.
452    #[test]
453    fn four_blocks_match_the_byte_at_a_time_path() {
454        for key_len in [16usize, 24, 32] {
455            let key: Vec<u8> = (0..key_len).map(|i| (i * 11 + 5) as u8).collect();
456            let sched = Schedule::expand(&key).unwrap();
457            let keys = RoundKeys::new(&sched);
458
459            for case in 0..64u32 {
460                let mut data = [0u8; GROUP];
461                for (i, slot) in data.iter_mut().enumerate() {
462                    *slot = (i as u32)
463                        .wrapping_mul(case.wrapping_add(1))
464                        .wrapping_add(case) as u8;
465                }
466                // Distinct lanes: otherwise cross-lane contamination is
467                // indistinguishable from correct behaviour.
468                let mut want = data;
469                for block in want.chunks_exact_mut(BLOCK_LEN) {
470                    super::super::portable::encrypt_block(&sched, block).unwrap();
471                }
472                let mut got = data;
473                encrypt_group(&keys, &mut got);
474                assert_eq!(got, want, "key_len {key_len}, case {case}");
475            }
476        }
477    }
478
479    /// One lane at a time, with the other three zeroed.
480    ///
481    /// A cross-lane leak that happens to cancel on structured data will not
482    /// cancel here: three of the four blocks have a known answer of their own,
483    /// and any bleed from the fourth shows up in them.
484    #[test]
485    fn lanes_do_not_leak_into_each_other() {
486        let key = [0x42u8; 32];
487        let sched = Schedule::expand(&key).unwrap();
488        let keys = RoundKeys::new(&sched);
489
490        for lane in 0..LANES {
491            let mut data = [0u8; GROUP];
492            for k in 0..BLOCK_LEN {
493                data[lane * BLOCK_LEN + k] = (k as u8).wrapping_mul(37).wrapping_add(1);
494            }
495            let mut want = data;
496            for block in want.chunks_exact_mut(BLOCK_LEN) {
497                super::super::portable::encrypt_block(&sched, block).unwrap();
498            }
499            let mut got = data;
500            encrypt_group(&keys, &mut got);
501            assert_eq!(got, want, "lane {lane}");
502        }
503    }
504}