crypton-2.1.3: cbits/aes/gcm_vaes512_x86.c
/*
* AES-GCM through VAES and VPCLMULQDQ in their 512-bit form, which takes
* four blocks where the 256-bit form in cbits/aes/gcm_vaes_x86.c takes two
* and the 128-bit one takes one. The instruction rate is the same, so the
* work per group halves again.
*
* This is that file widened and nothing else: the same group of sixteen
* blocks, the same descending powers of H sharing one reduction, the same
* round keys read from memory rather than held in registers. Four blocks to
* a register means the group fills four of them rather than eight, which is
* what leaves room for the group's own ciphertext to be kept for the GHASH
* when encrypting.
*
* Nothing here is borrowed. OpenSSL's and BoringSSL's AVX-512 AES-GCM is
* Apache-2.0 and s2n-bignum has no GCM at all.
*
* The reduction at the end is a copy of the one in gcm_vaes_x86.c rather
* than a call to it: the two files are compiled for different instruction
* sets, and a function compiled for one cannot be inlined into the other.
*/
#include "aes/gcm_vaes512_x86.h"
#ifdef WITH_GCM_VAES512
#include <string.h>
#include <immintrin.h>
#include <aes/gf.h>
#include <aes/block128.h>
#if defined(__clang__) || defined(__GNUC__)
#define V512_TARGET \
__attribute__((target("avx512f,avx512bw,avx512vl,aes,pclmul,vaes,vpclmulqdq")))
#else
#define V512_TARGET
#endif
/*
* Thirty-two blocks to a group, four to a register, so eight registers are
* in flight. The number of registers is what matters as much as the blocks
* per instruction: AES-NI has a latency of four cycles against a throughput
* of one, so it takes eight independent chains to keep two ports busy. Four
* registers of four blocks was written first and measured *slower* than the
* 256-bit path -- the blocks per instruction had doubled and the chains had
* halved.
*
* The table holds sixteen powers of H, so the GHASH of a group is two passes
* of sixteen blocks, the second picking up the tag the first leaves.
*/
#define V512WIDE 8
#define V512HALF 4
#define V512BYTES 512
/*
* The 128-bit multiply of cbits/aes/x86ni.c, done in all four lanes at once.
* Every shuffle and shift here works inside its own 128-bit lane, so the
* four products never mix: what comes out is four independent carry-less
* products, accumulated by the caller and reduced together at the end.
*/
V512_TARGET
static inline void clmul512(__m512i a, __m512i b, __m512i *lo, __m512i *hi)
{
const __m512i bswap = _mm512_set4_epi32(
0x00010203, 0x04050607, 0x08090a0b, 0x0c0d0e0f);
__m512i t3, t4, t5, t6;
a = _mm512_shuffle_epi8(a, bswap);
/* Karatsuba, as in the 128-bit one: three multiplies, not four */
t3 = _mm512_clmulepi64_epi128(a, b, 0x00);
t6 = _mm512_clmulepi64_epi128(a, b, 0x11);
t4 = _mm512_clmulepi64_epi128(
_mm512_xor_si512(a, _mm512_shuffle_epi32(a, _MM_PERM_BADC)),
_mm512_xor_si512(b, _mm512_shuffle_epi32(b, _MM_PERM_BADC)),
0x00);
t4 = _mm512_xor_si512(t4, _mm512_xor_si512(t3, t6));
t5 = _mm512_bslli_epi128(t4, 8);
t4 = _mm512_bsrli_epi128(t4, 8);
*lo = _mm512_xor_si512(t3, t5);
*hi = _mm512_xor_si512(t6, t4);
}
/*
* The reduction of cbits/aes/x86ni.c. By the time it runs the four lanes
* have been folded into one, so there is one 256-bit product to reduce.
*/
V512_TARGET
static inline __m128i gfred512(__m128i t3, __m128i t6)
{
const __m128i bswap = _mm_set_epi8(0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15);
__m128i t2, t4, t5, t7, t8, t9;
t7 = _mm_srli_epi32(t3, 31);
t8 = _mm_srli_epi32(t6, 31);
t3 = _mm_slli_epi32(t3, 1);
t6 = _mm_slli_epi32(t6, 1);
t9 = _mm_srli_si128(t7, 12);
t8 = _mm_slli_si128(t8, 4);
t7 = _mm_slli_si128(t7, 4);
t3 = _mm_or_si128(t3, t7);
t6 = _mm_or_si128(t6, t8);
t6 = _mm_or_si128(t6, t9);
t7 = _mm_slli_epi32(t3, 31);
t8 = _mm_slli_epi32(t3, 30);
t9 = _mm_slli_epi32(t3, 25);
t7 = _mm_xor_si128(t7, t8);
t7 = _mm_xor_si128(t7, t9);
t8 = _mm_srli_si128(t7, 4);
t7 = _mm_slli_si128(t7, 12);
t3 = _mm_xor_si128(t3, t7);
t2 = _mm_srli_epi32(t3, 1);
t4 = _mm_srli_epi32(t3, 2);
t5 = _mm_srli_epi32(t3, 7);
t2 = _mm_xor_si128(t2, t4);
t2 = _mm_xor_si128(t2, t5);
t2 = _mm_xor_si128(t2, t8);
t3 = _mm_xor_si128(t3, t2);
t6 = _mm_xor_si128(t6, t3);
return _mm_shuffle_epi8(t6, bswap);
}
/* the four 128-bit lanes of a register added together */
V512_TARGET
static inline __m128i fold512(__m512i v)
{
__m256i h = _mm256_xor_si256(_mm512_castsi512_si256(v),
_mm512_extracti64x4_epi64(v, 1));
return _mm_xor_si128(_mm256_castsi256_si128(h),
_mm256_extracti128_si256(h, 1));
}
/*
* Sixteen blocks against H^16 .. H^1, one reduction. v[j] holds blocks 4j
* to 4j+3 in its four lanes, so the powers for it are H^(16-4j) down to
* H^(13-4j) -- the table's own order the other way round, hence the four
* 128-bit loads rather than one 512-bit one.
*/
V512_TARGET
static inline __m128i ghash16(__m128i tag, const table_4bit htable,
const __m512i *v, int fromwire)
{
__m512i lo = _mm512_setzero_si512(), hi = _mm512_setzero_si512();
__m512i l, h, b;
int j;
for (j = 0; j < V512HALF; j++) {
const __m128i p0 =
_mm_loadu_si128((const __m128i *) &htable[15 - 4 * j]);
const __m128i p1 =
_mm_loadu_si128((const __m128i *) &htable[14 - 4 * j]);
const __m128i p2 =
_mm_loadu_si128((const __m128i *) &htable[13 - 4 * j]);
const __m128i p3 =
_mm_loadu_si128((const __m128i *) &htable[12 - 4 * j]);
__m512i hp = _mm512_castsi128_si512(p0);
hp = _mm512_inserti32x4(hp, p1, 1);
hp = _mm512_inserti32x4(hp, p2, 2);
hp = _mm512_inserti32x4(hp, p3, 3);
b = fromwire ? _mm512_loadu_si512(v + j) : v[j];
if (j == 0) /* the running tag joins the first block */
b = _mm512_xor_si512(
b, _mm512_inserti32x4(
_mm512_setzero_si512(), tag, 0));
clmul512(b, hp, &l, &h);
lo = _mm512_xor_si512(lo, l);
hi = _mm512_xor_si512(hi, h);
}
/* the four lanes are independent products of the same sum: fold them */
return gfred512(fold512(lo), fold512(hi));
}
#define KK512(r) _mm512_broadcast_i32x4(_mm_loadu_si128(k_ + (r)))
#define AESENC32(K) \
do { \
const __m512i rk = (K); \
v[0] = _mm512_aesenc_epi128(v[0], rk); \
v[1] = _mm512_aesenc_epi128(v[1], rk); \
v[2] = _mm512_aesenc_epi128(v[2], rk); \
v[3] = _mm512_aesenc_epi128(v[3], rk); \
v[4] = _mm512_aesenc_epi128(v[4], rk); \
v[5] = _mm512_aesenc_epi128(v[5], rk); \
v[6] = _mm512_aesenc_epi128(v[6], rk); \
v[7] = _mm512_aesenc_epi128(v[7], rk); \
} while (0)
#define AESLAST32(K) \
do { \
const __m512i rk = (K); \
v[0] = _mm512_aesenclast_epi128(v[0], rk); \
v[1] = _mm512_aesenclast_epi128(v[1], rk); \
v[2] = _mm512_aesenclast_epi128(v[2], rk); \
v[3] = _mm512_aesenclast_epi128(v[3], rk); \
v[4] = _mm512_aesenclast_epi128(v[4], rk); \
v[5] = _mm512_aesenclast_epi128(v[5], rk); \
v[6] = _mm512_aesenclast_epi128(v[6], rk); \
v[7] = _mm512_aesenclast_epi128(v[7], rk); \
} while (0)
#define XOR32(K) \
do { \
const __m512i rk = (K); \
v[0] = _mm512_xor_si512(v[0], rk); \
v[1] = _mm512_xor_si512(v[1], rk); \
v[2] = _mm512_xor_si512(v[2], rk); \
v[3] = _mm512_xor_si512(v[3], rk); \
v[4] = _mm512_xor_si512(v[4], rk); \
v[5] = _mm512_xor_si512(v[5], rk); \
v[6] = _mm512_xor_si512(v[6], rk); \
v[7] = _mm512_xor_si512(v[7], rk); \
} while (0)
/*
* The rounds are written out rather than looped for the reason the 128-bit
* loop gives: the count is a value in the key, and a loop over it leaves the
* round key reached through an index the compiler cannot fold.
*/
V512_TARGET
static inline __attribute__((always_inline)) void
rounds32(__m512i *v, const uint8_t *k, const int nbr)
{
const __m128i *k_ = (const __m128i *) k;
XOR32(KK512(0));
AESENC32(KK512(1)); AESENC32(KK512(2)); AESENC32(KK512(3));
AESENC32(KK512(4)); AESENC32(KK512(5)); AESENC32(KK512(6));
AESENC32(KK512(7)); AESENC32(KK512(8)); AESENC32(KK512(9));
if (nbr > 10) {
AESENC32(KK512(10)); AESENC32(KK512(11));
if (nbr > 12) {
AESENC32(KK512(12)); AESENC32(KK512(13));
}
}
AESLAST32(_mm512_broadcast_i32x4(_mm_loadu_si128(k_ + nbr)));
}
/* sixteen consecutive counters, four to a register. GCM counts in the low
* thirty-two bits and wraps there, which is what _mm_add_epi32 does. */
V512_TARGET
static inline __m128i counters32(__m512i *v, __m128i iv, __m128i one,
__m128i bswap)
{
int j;
for (j = 0; j < V512WIDE; j++) {
__m128i c0, c1, c2, c3;
__m512i c;
iv = _mm_add_epi32(iv, one);
c0 = _mm_shuffle_epi8(iv, bswap);
iv = _mm_add_epi32(iv, one);
c1 = _mm_shuffle_epi8(iv, bswap);
iv = _mm_add_epi32(iv, one);
c2 = _mm_shuffle_epi8(iv, bswap);
iv = _mm_add_epi32(iv, one);
c3 = _mm_shuffle_epi8(iv, bswap);
c = _mm512_castsi128_si512(c0);
c = _mm512_inserti32x4(c, c1, 1);
c = _mm512_inserti32x4(c, c2, 2);
c = _mm512_inserti32x4(c, c3, 3);
v[j] = c;
}
return iv;
}
/*
* Inlined into three callers with the round count a constant in each, which
* folds away the tests inside the group loop -- the same reason the 256-bit
* file gives, where without it AES-256 lost what AES-128 gained.
*/
V512_TARGET
static inline __attribute__((always_inline)) uint32_t
bulk_n(uint8_t *output, aes_gcm *gcm, const aes_key *key,
const uint8_t *input, uint32_t length, int decrypt, const int nbr)
{
const __m128i bswap = _mm_setr_epi8(7,6,5,4,3,2,1,0,15,14,13,12,11,10,9,8);
const __m128i one = _mm_set_epi32(0, 1, 0, 0);
__m512i v[V512WIDE];
__m128i iv, tag;
uint32_t groups = length / V512BYTES;
uint32_t done = 0;
uint32_t g;
int j;
if (groups == 0)
return 0;
iv = _mm_shuffle_epi8(_mm_loadu_si128((const __m128i *) &gcm->civ), bswap);
tag = _mm_loadu_si128((const __m128i *) &gcm->tag);
for (g = 0; g < groups; g++, input += V512BYTES, output += V512BYTES,
done += V512BYTES) {
iv = counters32(v, iv, one, bswap);
rounds32(v, key->data, nbr);
for (j = 0; j < V512WIDE; j++) {
const __m512i in =
_mm512_loadu_si512((const __m512i *) (input + 64 * j));
v[j] = _mm512_xor_si512(v[j], in);
_mm512_storeu_si512((__m512i *) (output + 64 * j), v[j]);
}
/* sixteen blocks to a pass, since that is how many powers of
* H the table holds; the second picks up the tag the first
* leaves */
tag = ghash16(tag, gcm->htable,
decrypt ? (const __m512i *) input : v,
decrypt);
tag = ghash16(tag, gcm->htable,
decrypt ? (const __m512i *) (input + 256)
: v + V512HALF,
decrypt);
}
_mm_storeu_si128((__m128i *) &gcm->civ, _mm_shuffle_epi8(iv, bswap));
_mm_storeu_si128((__m128i *) &gcm->tag, tag);
return done;
}
V512_TARGET
static uint32_t bulk(uint8_t *output, aes_gcm *gcm, const aes_key *key,
const uint8_t *input, uint32_t length, int decrypt)
{
switch (key->nbr) {
case 10:
return bulk_n(output, gcm, key, input, length, decrypt, 10);
case 12:
return bulk_n(output, gcm, key, input, length, decrypt, 12);
case 14:
return bulk_n(output, gcm, key, input, length, decrypt, 14);
default:
return 0; /* not a key length AES has */
}
}
uint32_t crypton_gcm_vaes512_bulk_encrypt(uint8_t *output, aes_gcm *gcm,
const aes_key *key,
const uint8_t *input,
uint32_t length)
{
return bulk(output, gcm, key, input, length, 0);
}
uint32_t crypton_gcm_vaes512_bulk_decrypt(uint8_t *output, aes_gcm *gcm,
const aes_key *key,
const uint8_t *input,
uint32_t length)
{
return bulk(output, gcm, key, input, length, 1);
}
#endif