Skip to main content

aes/backends/x86_aes/
expand.rs

1#![allow(unsafe_op_in_unsafe_fn)]
2
3#[cfg(target_arch = "x86")]
4use core::arch::x86::*;
5#[cfg(target_arch = "x86_64")]
6use core::arch::x86_64::*;
7
8pub(crate) type RoundKeys<const RK: usize> = [__m128i; RK];
9
10#[inline]
11#[target_feature(enable = "aes")]
12pub(super) unsafe fn aes128_expand_key(key: &[u8; 16]) -> RoundKeys<11> {
13    unsafe fn expand_round<const R: i32>(keys: &mut RoundKeys<11>, pos: usize) {
14        let mut t1 = keys[pos - 1];
15        let mut t2;
16        let mut t3;
17
18        t2 = _mm_aeskeygenassist_si128(t1, R);
19        t2 = _mm_shuffle_epi32(t2, 0xff);
20        t3 = _mm_slli_si128(t1, 0x4);
21        t1 = _mm_xor_si128(t1, t3);
22        t3 = _mm_slli_si128(t3, 0x4);
23        t1 = _mm_xor_si128(t1, t3);
24        t3 = _mm_slli_si128(t3, 0x4);
25        t1 = _mm_xor_si128(t1, t3);
26        t1 = _mm_xor_si128(t1, t2);
27
28        keys[pos] = t1;
29    }
30
31    let mut keys = [_mm_setzero_si128(); 11];
32    let k = _mm_loadu_si128(key.as_ptr().cast());
33    keys[0] = k;
34
35    let kr = &mut keys;
36    expand_round::<0x01>(kr, 1);
37    expand_round::<0x02>(kr, 2);
38    expand_round::<0x04>(kr, 3);
39    expand_round::<0x08>(kr, 4);
40    expand_round::<0x10>(kr, 5);
41    expand_round::<0x20>(kr, 6);
42    expand_round::<0x40>(kr, 7);
43    expand_round::<0x80>(kr, 8);
44    expand_round::<0x1B>(kr, 9);
45    expand_round::<0x36>(kr, 10);
46
47    keys
48}
49
50#[inline]
51#[target_feature(enable = "aes")]
52pub(super) unsafe fn aes192_expand_key(key: &[u8; 24]) -> RoundKeys<13> {
53    unsafe fn unpack_hilo(a: __m128i, b: __m128i) -> __m128i {
54        let a = _mm_shuffle_epi32(a, 0b01_00_11_10);
55        _mm_unpacklo_epi64(a, b)
56    }
57
58    #[target_feature(enable = "aes")]
59    unsafe fn expand_round<const R: i32>(mut t1: __m128i, mut t3: __m128i) -> (__m128i, __m128i) {
60        let (mut t2, mut t4);
61
62        t2 = _mm_aeskeygenassist_si128(t3, R);
63        t2 = _mm_shuffle_epi32(t2, 0x55);
64        t4 = _mm_slli_si128(t1, 0x4);
65        t1 = _mm_xor_si128(t1, t4);
66        t4 = _mm_slli_si128(t4, 0x4);
67        t1 = _mm_xor_si128(t1, t4);
68        t4 = _mm_slli_si128(t4, 0x4);
69        t1 = _mm_xor_si128(t1, t4);
70        t1 = _mm_xor_si128(t1, t2);
71        t2 = _mm_shuffle_epi32(t1, 0xff);
72        t4 = _mm_slli_si128(t3, 0x4);
73        t3 = _mm_xor_si128(t3, t4);
74        t3 = _mm_xor_si128(t3, t2);
75
76        (t1, t3)
77    }
78
79    let mut keys = [_mm_setzero_si128(); 13];
80    // We are being extra pedantic here to remove out-of-bound access.
81    // This should be optimized into movups, movsd sequence.
82    let (k0, k1l) = {
83        let mut t = [0u8; 32];
84        t[..key.len()].copy_from_slice(key);
85        (
86            _mm_loadu_si128(t.as_ptr().cast()),
87            _mm_loadu_si128(t.as_ptr().offset(16).cast()),
88        )
89    };
90
91    keys[0] = k0;
92
93    let (k1_2, k2r) = expand_round::<0x01>(k0, k1l);
94    keys[1] = _mm_unpacklo_epi64(k1l, k1_2);
95    keys[2] = unpack_hilo(k1_2, k2r);
96
97    let (k3, k4l) = expand_round::<0x02>(k1_2, k2r);
98    keys[3] = k3;
99
100    let (k4_5, k5r) = expand_round::<0x04>(k3, k4l);
101    let k4 = _mm_unpacklo_epi64(k4l, k4_5);
102    let k5 = unpack_hilo(k4_5, k5r);
103    keys[4] = k4;
104    keys[5] = k5;
105
106    let (k6, k7l) = expand_round::<0x08>(k4_5, k5r);
107    keys[6] = k6;
108
109    let (k7_8, k8r) = expand_round::<0x10>(k6, k7l);
110    keys[7] = _mm_unpacklo_epi64(k7l, k7_8);
111    keys[8] = unpack_hilo(k7_8, k8r);
112
113    let (k9, k10l) = expand_round::<0x20>(k7_8, k8r);
114    keys[9] = k9;
115
116    let (k10_11, k11r) = expand_round::<0x40>(k9, k10l);
117    keys[10] = _mm_unpacklo_epi64(k10l, k10_11);
118    keys[11] = unpack_hilo(k10_11, k11r);
119
120    let (k12, _) = expand_round::<0x80>(k10_11, k11r);
121    keys[12] = k12;
122
123    keys
124}
125
126#[inline]
127#[target_feature(enable = "aes")]
128pub(super) unsafe fn aes256_expand_key(key: &[u8; 32]) -> RoundKeys<15> {
129    unsafe fn expand_round<const R: i32>(keys: &mut RoundKeys<15>, pos: usize) {
130        let mut t1 = keys[pos - 2];
131        let mut t2;
132        let mut t3 = keys[pos - 1];
133        let mut t4;
134
135        t2 = _mm_aeskeygenassist_si128(t3, R);
136        t2 = _mm_shuffle_epi32(t2, 0xff);
137        t4 = _mm_slli_si128(t1, 0x4);
138        t1 = _mm_xor_si128(t1, t4);
139        t4 = _mm_slli_si128(t4, 0x4);
140        t1 = _mm_xor_si128(t1, t4);
141        t4 = _mm_slli_si128(t4, 0x4);
142        t1 = _mm_xor_si128(t1, t4);
143        t1 = _mm_xor_si128(t1, t2);
144
145        keys[pos] = t1;
146
147        t4 = _mm_aeskeygenassist_si128(t1, 0x00);
148        t2 = _mm_shuffle_epi32(t4, 0xaa);
149        t4 = _mm_slli_si128(t3, 0x4);
150        t3 = _mm_xor_si128(t3, t4);
151        t4 = _mm_slli_si128(t4, 0x4);
152        t3 = _mm_xor_si128(t3, t4);
153        t4 = _mm_slli_si128(t4, 0x4);
154        t3 = _mm_xor_si128(t3, t4);
155        t3 = _mm_xor_si128(t3, t2);
156
157        keys[pos + 1] = t3;
158    }
159
160    unsafe fn expand_round_last<const R: i32>(keys: &mut RoundKeys<15>, pos: usize) {
161        let mut t1 = keys[pos - 2];
162        let mut t2;
163        let t3 = keys[pos - 1];
164        let mut t4;
165
166        t2 = _mm_aeskeygenassist_si128(t3, R);
167        t2 = _mm_shuffle_epi32(t2, 0xff);
168        t4 = _mm_slli_si128(t1, 0x4);
169        t1 = _mm_xor_si128(t1, t4);
170        t4 = _mm_slli_si128(t4, 0x4);
171        t1 = _mm_xor_si128(t1, t4);
172        t4 = _mm_slli_si128(t4, 0x4);
173        t1 = _mm_xor_si128(t1, t4);
174        t1 = _mm_xor_si128(t1, t2);
175
176        keys[pos] = t1;
177    }
178
179    let mut keys = [_mm_setzero_si128(); 15];
180
181    let kp = key.as_ptr().cast::<__m128i>();
182    keys[0] = _mm_loadu_si128(kp);
183    keys[1] = _mm_loadu_si128(kp.add(1));
184
185    let k = &mut keys;
186    expand_round::<0x01>(k, 2);
187    expand_round::<0x02>(k, 4);
188    expand_round::<0x04>(k, 6);
189    expand_round::<0x08>(k, 8);
190    expand_round::<0x10>(k, 10);
191    expand_round::<0x20>(k, 12);
192    expand_round_last::<0x40>(k, 14);
193
194    keys
195}
196
197#[inline]
198#[target_feature(enable = "aes")]
199pub(super) unsafe fn inv_expanded_keys<const N: usize>(keys: &[__m128i; N]) -> [__m128i; N] {
200    let mut inv_keys: [__m128i; N] = [_mm_setzero_si128(); N];
201    inv_keys[0] = keys[N - 1];
202    for i in 1..N - 1 {
203        inv_keys[i] = _mm_aesimc_si128(keys[N - 1 - i]);
204    }
205    inv_keys[N - 1] = keys[0];
206    inv_keys
207}