Translation

This post is also available in Simplified Chinese.

Notice

This article was originally published on Anquanke: original article

AES-NI is Intel’s x86-64 SIMD extension for accelerating AES. Anyone familiar with SIMD probably knows it exists, but AES’s asymmetric structure and AES-NI’s unusual design require considerable detail and theory to use correctly. Using easyRE from N1CTF 2021 as an example, this article summarizes my understanding; corrections are welcome.

AES Structure

AES-128 is a ten-round 4×4 substitution-permutation network whose final round omits MixColumns.

AES encryption and decryption structure

Although AES has ten rounds, an AddRoundKey occurs before round one, so there are eleven round keys. Encryption begins and ends with AddRoundKey, a design called key whitening. The reason is straightforward: the other three operations are fixed and key-independent. At either boundary anyone could invert them, so they would add no security.

AESENC and AESENCLAST

These are AES-NI’s encryption instructions and the easiest to understand. The Intel Intrinsics Guide shows that AESENC applies ShiftRows, SubBytes, MixColumns, and AddRoundKey. Since SubBytes operates independently on bytes, it commutes with ShiftRows; AESENC therefore implements one ordinary round in the diagram.

AESENCLAST applies ShiftRows, SubBytes, and AddRoundKey, implementing the final round.

The initial round-key XOR uses PXOR, giving the following complete encryption (pt is plaintext, k[x] a round key, and ct ciphertext):

1
2
3
4
5
6
7
pxor pt, k[0]
aesenc pt, k[1]
aesenc pt, k[2]
...
aesenc pt, k[n-1]
aesenclast pt, k[n]
movdqa ct, pt

Nine AESENC rounds plus one AESENCLAST are easy to remember; the easily overlooked detail is the direct PXOR with round key zero.

AES Decryption and Equivalent Inverse Cipher

AES’s asymmetric design is deceptive. The decryption side of the diagram also consists of whitening, nine ordinary rounds, and one final round.

Naively reversing encryption suggests decrypting the final round first and ordinary rounds afterward, yet the diagram clearly does not do that.

Ignoring round boundaries, decryption performs the four inverse operations in reverse order. There are many ways to group that sequence into rounds, however. Under the conventional grouping shown above, a decryption round is not the inverse of an encryption round—AES’s first counterintuitive feature.

Here a decryption round contains InvShiftRows, InvSubBytes, AddRoundKey, and InvMixColumns; the final round again omits the MixColumns operation.

AES was originally Rijndael. Its original proposal also defines an “equivalent inverse cipher” (section 5.3.3), swapping AddRoundKey and InvMixColumns so decryption mirrors encryption with AddRoundKey last. InvSubBytes and InvShiftRows themselves commute.

AES equivalent inverse-cipher structure

This swap is not inherently equivalent. InvMixColumns multiplies each four-byte column by a 4×4 matrix over GF(2^8), while AddRoundKey XORs each byte. XOR is addition in GF(2^8); distributivity therefore requires multiplying the round key by the same matrix when moving AddRoundKey after InvMixColumns.

Round key zero is XORed directly and the last key belongs to the final decryption round, so neither participates in the swap. Thus equivalent decryption reverses the encryption keys and applies InvMixColumns to keys 1 through n−1, producing decryption keys.

Encryption and equivalent-decryption rounds have an elegant symmetry but use different round keys—AES’s second counterintuitive feature.

AESDEC, AESDECLAST, and AESIMC

Intel also uses equivalent decryption, as described in the AES-NI white paper. AESDEC is not the inverse of AESENC, nor is AESDECLAST the inverse of AESENCLAST. Complete decryption is:

1
2
3
4
5
6
7
pxor ct, k[n]
aesdec ct, k'[n-1]
aesdec ct, k'[n-2]
...
aesdec ct, k'[1]
aesdeclast ct, k[0]
movdqa pt, ct

Here k[0] and k[n] are unchanged, while k'[1] through k'[n-1] are encryption keys transformed by InvMixColumns. Intel provides AESIMC specifically to perform this single operation.

AESKEYGENASSIST and PCLMULQDQ

AESKEYGENASSIST supports key expansion; see page 19 of the white paper.

PCLMULQDQ, Carry-Less Multiplication Quadword, multiplies two polynomials over GF(2^128). It is not formally part of AES-NI, but besides accelerating CRC32 it computes GCM’s GMAC and therefore often appears in SIMD cryptography. Libsodium’s AES-256-GCM implementation is an excellent example.

Advanced AES-NI Uses

Isolating AES Operations

I initially wondered why Intel chose equivalent decryption, which requires an extra AESIMC step to generate decryption keys. The white paper revealed the elegance of the design.

Page 34 shows how AES-NI can isolate the individual AES operations:

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
Isolating ShiftRows
 PSHUFB xmm0, 0x0b06010c07020d08030e09040f0a0500
Isolating InvShiftRows
 PSHUFB xmm0, 0x0306090c0f0205080b0e0104070a0d00
Isolating MixColumns
 AESDECLAST xmm0, 0x00000000000000000000000000000000
 AESENC xmm0, 0x00000000000000000000000000000000
Isolating InvMixColumns
 AESENCLAST xmm0, 0x00000000000000000000000000000000
 AESDEC xmm0, 0x00000000000000000000000000000000
Isolating SubBytes
 PSHUFB xmm0, 0x0306090c0f0205080b0e0104070a0d00
 AESENCLAST xmm0, 0x00000000000000000000000000000000
Isolating InvSubBytes
 PSHUFB xmm0, 0x0b06010c07020d08030e09040f0a0500
 AESDECLAST xmm0, 0x00000000000000000000000000000000

ShiftRows is directly expressible with SSSE3 PSHUFB. SubBytes reverses the shuffle and performs a final round with a zero key, canceling the other two operations. MixColumns combines encryption and decryption so the final-round behavior leaves only MixColumns. The composition is remarkably clever.

AESIMC converts an encryption key to an equivalent-decryption key. To reverse this without a direct MixColumns instruction, combine AESDECLAST and AESENC as shown above.

The Intrinsics Guide shows that on Skylake, AESIMC has twice the latency and reciprocal throughput of AESENC. I suspect it internally composes AESENCLAST and AESDEC similarly.

Accelerating Other Algorithms

AES-NI’s flexibility supports larger substitution-permutation networks. The Rijndael proposal defines 128-, 192-, and 256-bit block sizes—not key sizes—but only Rijndael-128 became AES. Intel’s white paper implements other variants, including Rijndael-256:

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
#include <emmintrin.h>
#include <smmintrin.h>
void Rijndael256_encrypt(unsigned char* in,
                         unsigned char* out,
                         unsigned char* Key_Schedule,
                         unsigned long long length,
                         int number_of_rounds) {
  __m128i tmp1, tmp2, data1, data2;
  __m128i RIJNDAEL256_MASK =
      _mm_set_epi32(0x03020d0c, 0x0f0e0908, 0x0b0a0504, 0x07060100);
  __m128i BLEND_MASK =
      _mm_set_epi32(0x80000000, 0x80800000, 0x80800000, 0x80808000);
  __m128i* KS = (__m128i*)Key_Schedule;
  int i, j;
  for (i = 0; i < length / 32; i++) { /* loop over the data blocks */
    data1 = _mm_loadu_si128(&((__m128i*)in)[i * 2 + 0]); /* load data block */
    data2 = _mm_loadu_si128(&((__m128i*)in)[i * 2 + 1]);
    data1 = _mm_xor_si128(data1, KS[0]); /* round 0 (initial xor) */
    data2 = _mm_xor_si128(data2, KS[1]);
    /* Do number_of_rounds-1 AES rounds */
    for (j = 1; j < number_of_rounds; j++) {
      /*Blend to compensate for the shift rows shifts bytes between two
      128 bit blocks*/
      tmp1 = _mm_blendv_epi8(data1, data2, BLEND_MASK);
      tmp2 = _mm_blendv_epi8(data2, data1, BLEND_MASK);
      /*Shuffle that compensates for the additional shift in rows 3 and 4
      as opposed to rijndael128 (AES)*/
      tmp1 = _mm_shuffle_epi8(tmp1, RIJNDAEL256_MASK);
      tmp2 = _mm_shuffle_epi8(tmp2, RIJNDAEL256_MASK);
      /*This is the encryption step that includes sub bytes, shift rows,
      mix columns, xor with round key*/
      data1 = _mm_aesenc_si128(tmp1, KS[j * 2]);
      data2 = _mm_aesenc_si128(tmp2, KS[j * 2 + 1]);
    }
    tmp1 = _mm_blendv_epi8(data1, data2, BLEND_MASK);
    tmp2 = _mm_blendv_epi8(data2, data1, BLEND_MASK);
    tmp1 = _mm_shuffle_epi8(tmp1, RIJNDAEL256_MASK);
    tmp2 = _mm_shuffle_epi8(tmp2, RIJNDAEL256_MASK);
    tmp1 = _mm_aesenclast_si128(tmp1, KS[j * 2 + 0]); /*last AES round */
    tmp2 = _mm_aesenclast_si128(tmp2, KS[j * 2 + 1]);
    _mm_storeu_si128(&((__m128i*)out)[i * 2 + 0], tmp1);
    _mm_storeu_si128(&((__m128i*)out)[i * 2 + 1], tmp2);
  }
}

Rijndael-256 is an 8×4 network. Byte-level SubBytes and AddRoundKey work normally, as does per-column MixColumns; only ShiftRows needs SSE4.1 PBLENDVB and SSSE3 PSHUFB to adjust offsets. Since 8×4 is twice 4×4, every round uses two AESENC instructions and ends with two AESENCLASTs. This orderly, elegant code is part of what makes computers fascinating.

SM4’s nonlinear transform τ is also an S-box over GF(2^8), differing from AES only in its generating polynomial. The two fields are isomorphic, so algebraic transformations map their elements. Markku-Juhani O. Saarinen used this to accelerate SM4 with AES-NI; see sm4ni.

N1CTF 2021 easyRe

The challenge is available here. Its encryption function repeatedly encrypts and shuffles the plaintext in xmm0:

easyRe encryption function

Although the program uses V-prefixed AVX2 instructions, it touches only XMM registers, so decryption can use SSE alone. Parse the function with Capstone into an expression tree whose input is a leaf and ciphertext is the root. Rotate the tree left and right until the input becomes the root; the resulting expression is the decryption formula.

During rotation, inverses of VPXOR, VPADDQ, and VPSUBQ are straightforward. VPSHUFD permutes four 32-bit XMM elements and can permute them back. For VAESENC, extract the entire VAESENC/VAESENCLAST block, invert the intermediate round keys with VAESIMC, and generate the reverse decryption tree. Since AES’s key zero is applied by VPXOR, a VAESENC not preceded by VPXOR can be treated as XOR with an all-zero key, making the final VAESDECLAST key zero. Handle VAESDEC blocks similarly, using the VAESDECLAST plus VAESENC composition above to obtain MixColumns and transform round keys.

I wrote a JIT from the expression tree; compiling and running its generated code produces the flag:

  1
  2
  3
  4
  5
  6
  7
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
import sys
import capstone
import binascii

sys.setrecursionlimit(0x100000)

class Node:
    def __init__(self):
        self.emitted = False
        self.parent = None

    def __str__(self):
        return "v"+hex(id(self))

    def emit(self, f):
        if self.emitted:
            return
        self.emitted = True
        self.do_emit(f)

class Constant(Node):
    def __init__(self, c):
        super().__init__()
        self.c = c

    def do_emit(self, f):
        if self.c == 0:
            f.write("__m128i {}=_mm_setzero_si128();\n".format(self))
        else:
            f.write("__m128i {}=_mm_set_epi64x({}ULL,{}ULL);\n".format(
                self, hex(self.c >> 64), hex(self.c & ((1 << 64)-1))))

class Binary(Node):
    def __init__(self, a, b):
        super().__init__()
        self.a = a
        self.b = b
        a.parent = self
        b.parent = self

class Add(Binary):
    def __init__(self, a, b):
        super().__init__(a, b)

    def do_emit(self, f):
        self.a.emit(f)
        self.b.emit(f)
        f.write("__m128i {}=_mm_add_epi64({},{});\n".format(
            self, self.a, self.b))

class Sub(Binary):
    def __init__(self, a, b):
        super().__init__(a, b)

    def do_emit(self, f):
        self.a.emit(f)
        self.b.emit(f)
        f.write("__m128i {}=_mm_sub_epi64({},{});\n".format(
            self, self.a, self.b))

class Xor(Binary):
    def __init__(self, a, b):
        super().__init__(a, b)

    def do_emit(self, f):
        self.a.emit(f)
        self.b.emit(f)
        f.write("__m128i {}=_mm_xor_si128({},{});\n".format(
            self, self.a, self.b))

class Aes(Node):
    def __init__(self, base, key, is_enc, is_last):
        super().__init__()
        self.base = base
        self.key = key
        self.is_enc, self.is_last = is_enc, is_last
        base.parent = self
        key.parent = self

    def do_emit(self, f):
        self.base.emit(f)
        self.key.emit(f)
        f.write("__m128i {}=_mm_aes{}{}_si128({},{});\n".format(
            self, "enc" if self.is_enc else "dec", "last" if self.is_last else "", self.base, self.key))

class Aesimc(Node):
    def __init__(self, a, is_imc):
        super().__init__()
        self.a = a
        self.is_imc = is_imc
        a.parent = self

    def do_emit(self, f):
        self.a.emit(f)
        if self.is_imc:
            f.write("__m128i {}=_mm_aesimc_si128({});\n".format(self, self.a))
        else:
            f.write("__m128i {}=_mm_aesenc_si128(_mm_aesdeclast_si128({},zero),zero);\n".format(
                self, self.a))

class Shuffle(Node):
    def __init__(self, a, x):
        super().__init__()
        self.a = a
        self.x = x
        a.parent = self

    def do_emit(self, f):
        self.a.emit(f)
        f.write("__m128i {}=_mm_shuffle_epi32({},{});\n".format(
            self, self.a, hex(self.x)))

def flip(root):
    parent = root.parent
    if isinstance(parent, Constant):
        return parent
    elif isinstance(parent, Xor):
        if root == parent.a:
            return Xor(parent.b, flip(parent))
        else:
            return Xor(parent.a, flip(parent))
    elif isinstance(parent, Add):
        if root == parent.a:
            return Sub(flip(parent), parent.b)
        else:
            return Sub(flip(parent), parent.a)
    elif isinstance(parent, Sub):
        if root == parent.a:
            return Add(flip(parent), parent.b)
        else:
            return Sub(parent.a, flip(parent))
    elif isinstance(parent, Shuffle):
        x = parent.x
        shuffle = []
        for i in range(4):
            shuffle.append(x & 3)
            x >>= 2
        assert set(shuffle) == set({0, 1, 2, 3})
        x = 0
        for i in range(4):
            x <<= 2
            x += shuffle.index(3-i)
        return Shuffle(flip(parent), x)
    elif isinstance(parent, Aesimc):
        return Aesimc(flip(parent), not parent.is_imc)
    elif isinstance(parent, Aes):
        keys = [parent]
        p = parent.parent
        while True:
            if isinstance(p, Aes):
                keys.append(p)
                if p.is_last:
                    break
                p = p.parent
            else:
                raise ValueError
        keys.reverse()
        r = Xor(flip(p), keys[0].key)
        r_keys = {}
        for i in range(1, len(keys)):
            if id(keys[i].key) not in r_keys:
                r_keys[id(keys[i].key)] = Aesimc(keys[i].key, keys[i].is_enc)
            r = Aes(r, r_keys[id(keys[i].key)],
                    not keys[i].is_enc, False)
        return Aes(r, Constant(0), not keys[0].is_enc, True)
    else:
        raise ValueError

xmmnames = ['xmm{}'.format(i) for i in range(16)]
xmm = [None for i in range(16)]
target = Node()
xmm[0] = target
memory = {}
c = ['2f0fc4f2839a1d5401ead9842fc23d00',
     '24e1c94761c31694cdb7d3a38fb0c100',
     '2af5fcb6d4373ceac4590d4f86956d00',
     'cbc6b50249b0b519a2620a3cc73d9200',
     '60a876c1193162a02a1531a79d6a5900',
     'd083cfb2f3a048c4cf47af9bcaaefa00',
     'eb93d59f3756816e2671cd0d1c73bf00',
     'c32de58cdbcf9fdd7de74f364a594b00',
     '6055580a46572c4e6a591ddd77c0ce00',
     '13bf3e7536d86ce89d81348f6f10e000', ]
for i in range(len(c)):
    memory[0x620-i *
           0x10] = Constant(int.from_bytes(binascii.a2b_hex(c[i]), 'little'))
c = capstone.Cs(capstone.CS_ARCH_X86, capstone.CS_MODE_64)
c.detail = True
it = c.disasm(open('easyRe', 'rb').read()[0xc4d:0x19d90], 0x100000C4D)
for ins in it:
    if ins.mnemonic == 'vmovdqa':
        a, b = ins.op_str.split(',')
        a, b = a.strip(), b.strip()
        if a in xmmnames and b not in xmmnames:
            assert b.startswith('xmmword ptr [rbp - ')
            off = int(b[19:b.index(']')], 16)
            assert off in memory
            xmm[xmmnames.index(a)] = memory[off]
        elif b in xmmnames and a not in xmmnames:
            assert a.startswith('xmmword ptr [rbp - ')
            off = int(a[19:a.index(']')], 16)
            memory[off] = xmm[xmmnames.index(b)]
        else:
            xmm[xmmnames.index(a)] = xmm[xmmnames.index(b)]
    elif ins.mnemonic == 'vpxor':
        a, b, c = ins.op_str.split(',')
        a, b, c = a.strip(), b.strip(), c.strip()
        if c in xmmnames:
            xmm[xmmnames.index(a)] = Xor(
                xmm[xmmnames.index(b)], xmm[xmmnames.index(c)])
        else:
            assert c.startswith('xmmword ptr [rbp - ')
            off = int(c[19:c.index(']')], 16)
            xmm[xmmnames.index(a)] = Xor(
                xmm[xmmnames.index(b)], memory[off])
    elif ins.mnemonic == 'vpaddq' or ins.mnemonic == 'vpsubq':
        a, b, c = ins.op_str.split(',')
        a, b, c = a.strip(), b.strip(), c.strip()
        if c in xmmnames:
            xmm[xmmnames.index(a)] = (Add if ins.mnemonic == 'vpaddq' else Sub)(
                xmm[xmmnames.index(b)], xmm[xmmnames.index(c)])
        else:
            assert c.startswith('xmmword ptr [rbp - ')
            off = int(c[19:c.index(']')], 16)
            xmm[xmmnames.index(a)] = (Add if ins.mnemonic == 'vpaddq' else Sub)(
                xmm[xmmnames.index(b)], memory[off])
    elif ins.mnemonic == 'vpshufd':
        a, b, c = ins.op_str.split(',')
        a, b, c = a.strip(), b.strip(), c.strip()
        c = int(c, 16)
        xmm[xmmnames.index(a)] = Shuffle(
            xmm[xmmnames.index(b)], c)
    elif ins.mnemonic == 'vaesenc' or ins.mnemonic == 'vaesdec':
        is_enc = ins.mnemonic == 'vaesenc'
        a, b, c = ins.op_str.split(',')
        a, b, c = a.strip(), b.strip(), c.strip()
        xmm[xmmnames.index(a)] = Aes(xmm[xmmnames.index(b)],
                                     xmm[xmmnames.index(c)], is_enc, False)
    elif ins.mnemonic == 'vaesenclast' or ins.mnemonic == 'vaesdeclast':
        is_enc = ins.mnemonic == 'vaesenclast'
        a, b, c = ins.op_str.split(',')
        a, b, c = a.strip(), b.strip(), c.strip()
        xmm[xmmnames.index(a)] = Aes(xmm[xmmnames.index(b)],
                                     xmm[xmmnames.index(c)], is_enc, True)
    elif ins.mnemonic == 'vaesimc':
        a, b = ins.op_str.split(',')
        a, b = a.strip(), b.strip()
        xmm[xmmnames.index(a)] = Aesimc(xmm[xmmnames.index(b)], True)
    elif ins.mnemonic == 'movabs' or ins.mnemonic == 'mov':
        pass
    else:
        print(ins)
        raise ValueError
xmm[0].parent = Constant(0x79eeb3fa8c39dbd77bc066c7647d0b72)
target = flip(target)

f = open('a.c', 'w')
f.write('''
#include <immintrin.h>
#include <stdio.h>

int main(){
__m128i zero=_mm_setzero_si128();
''')
target.emit(f)
f.write('''char pt[16];
_mm_storeu_si128((__m128i*)pt, {});
fwrite(pt,16,1,stdout);
return 0;
}}
'''.format(target))

Compile with -maes to enable AES-NI and produce SSE code. Adding -march=native enables more extensions and can automatically produce AVX2 plus VAES code; modern compilers are remarkably capable.

Flag: n1ctf{Easy_AVX!}—not easy at all.