libdpf/include/dpf/prg_aes.hpp

640 lines
25 KiB
C++
Raw Normal View History

/// @file dpf/prg_aes.hpp
/// @brief Fixed-key AES Matyas–Meyer–Oseas PRG.
/// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details.
#ifndef LIBDPF_INCLUDE_DPF_PRG_AES_HPP__
#define LIBDPF_INCLUDE_DPF_PRG_AES_HPP__
#include <cassert>
#include <cstddef>
#include <cstdint>
#include <array>
#include <stdexcept>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "portable-snippets/exact-int/exact-int.h"
#include "dpf/prg_count.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
namespace prg
{
#ifdef __ARM_NEON
#else
#define simde_mm_aesenc_si128(a, RoundKey) _mm_aesenc_si128(a, RoundKey)
#define simde_mm_aesenclast_si128(a, RoundKey) _mm_aesenclast_si128(a, RoundKey)
#define simde_mm_aeskeygenassist_si128(a, inn8) _mm_aeskeygenassist_si128(a, inn8)
#endif
#if defined(__VAES__) && defined(__AVX2__)
#define DPF_PRG_AES_HAS_VAES 1
#include <immintrin.h>
#endif
template <typename AesKey>
struct aes final
{
using block_type = simde__m128i;
static constexpr primitive counted_as =
AesKey::rounds == 14 ? primitive::aes256 : primitive::aes128;
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static void require_block_aligned(const void * p) noexcept
{
(void)p;
assert(p == nullptr
|| reinterpret_cast<std::uintptr_t>(p) % alignof(block_type) == 0);
}
/// @brief One AES block in Matyas–Meyer–Oseas mode.
/// @param seed the PRG seed
/// @param pos the block counter mixed into the first round key
/// @return the 128-bit block
/// \complexity One block: `key.rounds - 1` `aesenc` calls plus one `aesenclast`. `rounds` is 10 for `aes128` and 14 for `aes256`.
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static block_type eval(block_type seed, psnip_uint32_t pos) noexcept
{
note_eval(1, counted_as);
block_type rd_key0 = simde_mm_xor_si128(key.rd_key[0],
simde_mm_set_epi64x(0, pos));
block_type output = simde_mm_xor_si128(seed, rd_key0);
for (std::size_t j = 1; j < key.rounds; ++j)
{
output = simde_mm_aesenc_si128(output, key.rd_key[j]);
}
output = simde_mm_aesenclast_si128(output, key.rd_key[key.rounds]);
output = simde_mm_xor_si128(output, seed);
return output;
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static auto eval01(block_type seed) noexcept
{
note_eval(2, counted_as);
block_type rd_key00 = key.rd_key[0];
block_type rd_key01 = simde_mm_xor_si128(rd_key00,
simde_mm_set_epi64x(0, 1));
block_type output0 = simde_mm_xor_si128(seed, rd_key00);
block_type output1 = simde_mm_xor_si128(seed, rd_key01);
HEDLEY_PRAGMA(GCC unroll(14))
for (std::size_t j = 1; j < key.rounds; ++j)
{
output0 = simde_mm_aesenc_si128(output0, key.rd_key[j]);
output1 = simde_mm_aesenc_si128(output1, key.rd_key[j]);
}
output0 = simde_mm_aesenclast_si128(output0, key.rd_key[key.rounds]);
output1 = simde_mm_aesenclast_si128(output1, key.rd_key[key.rounds]);
output0 = simde_mm_xor_si128(output0, seed);
output1 = simde_mm_xor_si128(output1, seed);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::array<block_type, 2>{output0, output1};
HEDLEY_PRAGMA(GCC diagnostic pop)
}
/// @brief Four independent MMO blocks, round-major so `aesenc` throughput
/// covers the latency of a single `eval`.
/// @param seeds four PRG seeds
/// @param pos lane index for each seed (`0` or `1` for a tree child)
/// @param out four blocks, same order as `seeds`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static void eval_indep4(const block_type seeds[4], const psnip_uint32_t pos[4],
block_type out[4]) noexcept
{
note_eval(4, counted_as);
block_type o0, o1, o2, o3;
{
const block_type k0 = key.rd_key[0];
o0 = simde_mm_xor_si128(seeds[0],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[0])));
o1 = simde_mm_xor_si128(seeds[1],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[1])));
o2 = simde_mm_xor_si128(seeds[2],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[2])));
o3 = simde_mm_xor_si128(seeds[3],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[3])));
}
HEDLEY_PRAGMA(GCC unroll(14))
for (std::size_t j = 1; j < key.rounds; ++j)
{
const block_type rk = key.rd_key[j];
o0 = simde_mm_aesenc_si128(o0, rk);
o1 = simde_mm_aesenc_si128(o1, rk);
o2 = simde_mm_aesenc_si128(o2, rk);
o3 = simde_mm_aesenc_si128(o3, rk);
}
const block_type last = key.rd_key[key.rounds];
o0 = simde_mm_xor_si128(simde_mm_aesenclast_si128(o0, last), seeds[0]);
o1 = simde_mm_xor_si128(simde_mm_aesenclast_si128(o1, last), seeds[1]);
o2 = simde_mm_xor_si128(simde_mm_aesenclast_si128(o2, last), seeds[2]);
o3 = simde_mm_xor_si128(simde_mm_aesenclast_si128(o3, last), seeds[3]);
out[0] = o0;
out[1] = o1;
out[2] = o2;
out[3] = o3;
}
/// @brief Eight independent MMO blocks. Zen's AES unit retires two
/// `aesenc`s per cycle at four-cycle latency, so eight in flight
/// fills it; four does not.
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static void eval_indep8(const block_type seeds[8], const psnip_uint32_t pos[8],
block_type out[8]) noexcept
{
note_eval(8, counted_as);
const block_type k0 = key.rd_key[0];
block_type o0 = simde_mm_xor_si128(seeds[0],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[0])));
block_type o1 = simde_mm_xor_si128(seeds[1],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[1])));
block_type o2 = simde_mm_xor_si128(seeds[2],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[2])));
block_type o3 = simde_mm_xor_si128(seeds[3],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[3])));
block_type o4 = simde_mm_xor_si128(seeds[4],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[4])));
block_type o5 = simde_mm_xor_si128(seeds[5],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[5])));
block_type o6 = simde_mm_xor_si128(seeds[6],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[6])));
block_type o7 = simde_mm_xor_si128(seeds[7],
simde_mm_xor_si128(k0, simde_mm_set_epi64x(0, pos[7])));
HEDLEY_PRAGMA(GCC unroll(14))
for (std::size_t j = 1; j < key.rounds; ++j)
{
const block_type rk = key.rd_key[j];
o0 = simde_mm_aesenc_si128(o0, rk);
o1 = simde_mm_aesenc_si128(o1, rk);
o2 = simde_mm_aesenc_si128(o2, rk);
o3 = simde_mm_aesenc_si128(o3, rk);
o4 = simde_mm_aesenc_si128(o4, rk);
o5 = simde_mm_aesenc_si128(o5, rk);
o6 = simde_mm_aesenc_si128(o6, rk);
o7 = simde_mm_aesenc_si128(o7, rk);
}
const block_type last = key.rd_key[key.rounds];
out[0] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o0, last), seeds[0]);
out[1] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o1, last), seeds[1]);
out[2] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o2, last), seeds[2]);
out[3] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o3, last), seeds[3]);
out[4] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o4, last), seeds[4]);
out[5] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o5, last), seeds[5]);
out[6] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o6, last), seeds[6]);
out[7] = simde_mm_xor_si128(simde_mm_aesenclast_si128(o7, last), seeds[7]);
}
/// @brief Round-major multi-block MMO. Positions use the same lane as
/// `eval` / `eval01` (`set_epi64x(0, pos)`). The first AddRoundKey
/// includes `rd_key[0]` so this matches the one-block `eval` for any
/// key, not only the all-zero key this PRG currently installs.
/// @param seed the PRG seed
/// @param output the destination. Unused when `count` is 0
/// @param count the number of blocks
/// @param pos the 0-based index
/// @throws std::invalid_argument if `prg lane index is out of range`
HEDLEY_ALWAYS_INLINE
static void eval(block_type seed, block_type * HEDLEY_RESTRICT output,
psnip_uint32_t count, psnip_uint32_t pos = 0)
{
if (HEDLEY_UNLIKELY(count == 0))
{
return;
}
// `pos + i` is a uint32 add. A span that passes UINT32_MAX must
// fail the same way buffered_prg does, rather than wrap to 0.
if (count > 1 &&
pos > static_cast<psnip_uint32_t>(~static_cast<psnip_uint32_t>(0)) - (count - 1u))
{
throw std::invalid_argument("prg lane index is out of range");
}
require_block_aligned(output);
block_type * HEDLEY_RESTRICT out =
static_cast<block_type *>(__builtin_assume_aligned(output,
alignof(block_type)));
if (HEDLEY_LIKELY(count == 1))
{
out[0] = eval(seed, pos);
return;
}
if (count == 2 && pos == 0)
{
auto kids = eval01(seed);
out[0] = kids[0];
out[1] = kids[1];
return;
}
note_eval(count, counted_as);
auto whitened = simde_mm_xor_si128(seed, key.rd_key[0]);
DPF_UNROLL_LOOP
for (psnip_uint32_t i = 0; i < count; ++i)
{
out[i] = simde_mm_xor_si128(whitened,
simde_mm_set_epi64x(0, pos + i));
}
DPF_UNROLL_LOOP
for (std::size_t j = 1; j < key.rounds; ++j)
{
DPF_UNROLL_LOOP
for (psnip_uint32_t i = 0; i < count; ++i)
{
out[i] = simde_mm_aesenc_si128(out[i], key.rd_key[j]);
}
}
DPF_UNROLL_LOOP
for (psnip_uint32_t i = 0; i < count; ++i)
{
out[i] = simde_mm_aesenclast_si128(out[i],
key.rd_key[key.rounds]);
out[i] = simde_mm_xor_si128(out[i], seed);
}
}
/// @brief Four independent `eval01` calls as one 8-block round-major AES.
/// @details `left[i] == eval(seeds[i], 0)`, `right[i] == eval(seeds[i], 1)`.
/// @param seeds the root seeds
/// @param left the `left`
/// @param right the `right`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2, 3)
static void eval01_x4(const block_type * HEDLEY_RESTRICT seeds,
block_type * HEDLEY_RESTRICT left,
block_type * HEDLEY_RESTRICT right) noexcept
{
note_eval(8, counted_as);
require_block_aligned(seeds);
require_block_aligned(left);
require_block_aligned(right);
const block_type * HEDLEY_RESTRICT s =
static_cast<const block_type *>(__builtin_assume_aligned(seeds,
alignof(block_type)));
block_type rk0 = key.rd_key[0];
block_type rk1 = simde_mm_xor_si128(rk0, simde_mm_set_epi64x(0, 1));
block_type blk[8], feed[8];
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
feed[2*i] = feed[2*i + 1] = s[i];
blk[2*i] = simde_mm_xor_si128(s[i], rk0);
blk[2*i + 1] = simde_mm_xor_si128(s[i], rk1);
}
aes_mmo_rounds_x8(blk, feed);
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
left[i] = blk[2*i];
right[i] = blk[2*i + 1];
}
}
/// @brief Four independent `eval(seed, pos)` as one 4-block round-major AES.
/// @param seeds the root seeds
/// @param output the destination. Unused when `count` is 0
/// @param pos the 0-based index
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2)
static void eval_x4(const block_type * HEDLEY_RESTRICT seeds,
block_type * HEDLEY_RESTRICT output, psnip_uint32_t pos) noexcept
{
note_eval(4, counted_as);
require_block_aligned(seeds);
require_block_aligned(output);
const block_type * HEDLEY_RESTRICT s =
static_cast<const block_type *>(__builtin_assume_aligned(seeds,
alignof(block_type)));
block_type * HEDLEY_RESTRICT out =
static_cast<block_type *>(__builtin_assume_aligned(output,
alignof(block_type)));
block_type rk0 = simde_mm_xor_si128(key.rd_key[0],
simde_mm_set_epi64x(0, pos));
block_type blk[4];
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
blk[i] = simde_mm_xor_si128(s[i], rk0);
}
aes_mmo_rounds_x4(blk, s);
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
out[i] = blk[i];
}
}
/// @brief Eight independent `eval(seed, pos)` as one 8-block round-major AES.
/// @param seeds the root seeds
/// @param output the destination. Unused when `count` is 0
/// @param pos the 0-based index
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2)
static void eval_x8(const block_type * HEDLEY_RESTRICT seeds,
block_type * HEDLEY_RESTRICT output, psnip_uint32_t pos) noexcept
{
note_eval(8, counted_as);
require_block_aligned(seeds);
require_block_aligned(output);
const block_type * HEDLEY_RESTRICT s =
static_cast<const block_type *>(__builtin_assume_aligned(seeds,
alignof(block_type)));
block_type * HEDLEY_RESTRICT out =
static_cast<block_type *>(__builtin_assume_aligned(output,
alignof(block_type)));
block_type rk0 = simde_mm_xor_si128(key.rd_key[0],
simde_mm_set_epi64x(0, pos));
block_type blk[8];
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 8; ++i)
{
blk[i] = simde_mm_xor_si128(s[i], rk0);
}
aes_mmo_rounds_x8(blk, s);
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 8; ++i)
{
out[i] = blk[i];
}
}
/// @brief Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`).
/// @tparam T value type
/// @tparam Party party index, `0` or `1`
/// @param seed the PRG seed
/// @param pos the 0-based index
/// @return Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`)
/// @see `prg.hpp`
template <typename T, std::size_t Party>
HEDLEY_NO_THROW
static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept;
private:
static const AesKey key;
/// @brief `blk[i]` is already `seed[i] XOR rd_key[0] XOR pos_i`. Runs AES
/// rounds 1..last and the MMO feed-forward `XOR seed[i]`.
/// @param blk the `blk`
/// @param seed the PRG seed
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2)
static void aes_mmo_rounds_x4(block_type * HEDLEY_RESTRICT blk,
const block_type * HEDLEY_RESTRICT seed) noexcept
{
block_type b0 = blk[0], b1 = blk[1], b2 = blk[2], b3 = blk[3];
HEDLEY_PRAGMA(GCC unroll(14))
for (std::size_t j = 1; j < key.rounds; ++j)
{
const block_type rk = key.rd_key[j];
b0 = simde_mm_aesenc_si128(b0, rk);
b1 = simde_mm_aesenc_si128(b1, rk);
b2 = simde_mm_aesenc_si128(b2, rk);
b3 = simde_mm_aesenc_si128(b3, rk);
}
const block_type last = key.rd_key[key.rounds];
blk[0] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b0, last), seed[0]);
blk[1] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b1, last), seed[1]);
blk[2] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b2, last), seed[2]);
blk[3] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b3, last), seed[3]);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2)
static void aes_mmo_rounds_x8(block_type * HEDLEY_RESTRICT blk,
const block_type * HEDLEY_RESTRICT seed) noexcept
{
#if defined(DPF_PRG_AES_HAS_VAES) && defined(__AVX512F__)
// 4 blocks per VAES instruction. Same round keys as the 128-bit path.
__m512i v0 = _mm512_castsi128_si512(blk[0]);
v0 = _mm512_inserti32x4(v0, blk[1], 1);
v0 = _mm512_inserti32x4(v0, blk[2], 2);
v0 = _mm512_inserti32x4(v0, blk[3], 3);
__m512i v1 = _mm512_castsi128_si512(blk[4]);
v1 = _mm512_inserti32x4(v1, blk[5], 1);
v1 = _mm512_inserti32x4(v1, blk[6], 2);
v1 = _mm512_inserti32x4(v1, blk[7], 3);
HEDLEY_PRAGMA(GCC unroll(14))
for (std::size_t j = 1; j < key.rounds; ++j)
{
const __m512i rk = _mm512_broadcast_i32x4(key.rd_key[j]);
v0 = _mm512_aesenc_epi128(v0, rk);
v1 = _mm512_aesenc_epi128(v1, rk);
}
const __m512i last = _mm512_broadcast_i32x4(key.rd_key[key.rounds]);
v0 = _mm512_aesenclast_epi128(v0, last);
v1 = _mm512_aesenclast_epi128(v1, last);
blk[0] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v0, 0), seed[0]);
blk[1] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v0, 1), seed[1]);
blk[2] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v0, 2), seed[2]);
blk[3] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v0, 3), seed[3]);
blk[4] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v1, 0), seed[4]);
blk[5] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v1, 1), seed[5]);
blk[6] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v1, 2), seed[6]);
blk[7] = simde_mm_xor_si128(_mm512_extracti32x4_epi32(v1, 3), seed[7]);
#elif defined(DPF_PRG_AES_HAS_VAES)
__m256i v0 = _mm256_insertf128_si256(_mm256_castsi128_si256(blk[0]), blk[1], 1);
__m256i v1 = _mm256_insertf128_si256(_mm256_castsi128_si256(blk[2]), blk[3], 1);
__m256i v2 = _mm256_insertf128_si256(_mm256_castsi128_si256(blk[4]), blk[5], 1);
__m256i v3 = _mm256_insertf128_si256(_mm256_castsi128_si256(blk[6]), blk[7], 1);
HEDLEY_PRAGMA(GCC unroll(14))
for (std::size_t j = 1; j < key.rounds; ++j)
{
const __m256i rk = _mm256_broadcastsi128_si256(key.rd_key[j]);
v0 = _mm256_aesenc_epi128(v0, rk);
v1 = _mm256_aesenc_epi128(v1, rk);
v2 = _mm256_aesenc_epi128(v2, rk);
v3 = _mm256_aesenc_epi128(v3, rk);
}
const __m256i last = _mm256_broadcastsi128_si256(key.rd_key[key.rounds]);
v0 = _mm256_aesenclast_epi128(v0, last);
v1 = _mm256_aesenclast_epi128(v1, last);
v2 = _mm256_aesenclast_epi128(v2, last);
v3 = _mm256_aesenclast_epi128(v3, last);
blk[0] = simde_mm_xor_si128(_mm256_castsi256_si128(v0), seed[0]);
blk[1] = simde_mm_xor_si128(_mm256_extracti128_si256(v0, 1), seed[1]);
blk[2] = simde_mm_xor_si128(_mm256_castsi256_si128(v1), seed[2]);
blk[3] = simde_mm_xor_si128(_mm256_extracti128_si256(v1, 1), seed[3]);
blk[4] = simde_mm_xor_si128(_mm256_castsi256_si128(v2), seed[4]);
blk[5] = simde_mm_xor_si128(_mm256_extracti128_si256(v2, 1), seed[5]);
blk[6] = simde_mm_xor_si128(_mm256_castsi256_si128(v3), seed[6]);
blk[7] = simde_mm_xor_si128(_mm256_extracti128_si256(v3, 1), seed[7]);
#else
block_type b0 = blk[0], b1 = blk[1], b2 = blk[2], b3 = blk[3];
block_type b4 = blk[4], b5 = blk[5], b6 = blk[6], b7 = blk[7];
HEDLEY_PRAGMA(GCC unroll(14))
for (std::size_t j = 1; j < key.rounds; ++j)
{
const block_type rk = key.rd_key[j];
b0 = simde_mm_aesenc_si128(b0, rk);
b1 = simde_mm_aesenc_si128(b1, rk);
b2 = simde_mm_aesenc_si128(b2, rk);
b3 = simde_mm_aesenc_si128(b3, rk);
b4 = simde_mm_aesenc_si128(b4, rk);
b5 = simde_mm_aesenc_si128(b5, rk);
b6 = simde_mm_aesenc_si128(b6, rk);
b7 = simde_mm_aesenc_si128(b7, rk);
}
const block_type last = key.rd_key[key.rounds];
blk[0] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b0, last), seed[0]);
blk[1] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b1, last), seed[1]);
blk[2] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b2, last), seed[2]);
blk[3] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b3, last), seed[3]);
blk[4] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b4, last), seed[4]);
blk[5] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b5, last), seed[5]);
blk[6] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b6, last), seed[6]);
blk[7] = simde_mm_xor_si128(simde_mm_aesenclast_si128(b7, last), seed[7]);
#endif
}
}; // struct aes
#define EXPAND_ASSIST(v1, v2, v3, v4, shuff_const, aes_const) \
v2 = simde_mm_aeskeygenassist_si128(v4, aes_const); \
v3 = simde_mm_castps_si128(_mm_shuffle_ps( \
simde_mm_castsi128_ps(v3), \
simde_mm_castsi128_ps(v1), 16)); \
v1 = simde_mm_xor_si128(v1, v3); \
v3 = simde_mm_castps_si128(simde_mm_shuffle_ps( \
simde_mm_castsi128_ps(v3), \
simde_mm_castsi128_ps(v1), 140)); \
v1 = simde_mm_xor_si128(v1, v3); \
v2 = simde_mm_shuffle_epi32(v2, shuff_const); \
v1 = simde_mm_xor_si128(v1, v2)
struct aes128_key
{
public:
static constexpr std::size_t rounds = 10;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using rd_key_array = std::array<simde__m128i, rounds+1>;
HEDLEY_PRAGMA(GCC diagnostic pop)
const rd_key_array rd_key;
explicit aes128_key(const simde__m128i & userkey)
: rd_key{compute_round_keys(userkey)} { }
private:
rd_key_array compute_round_keys(const simde__m128i & userkey)
{
rd_key_array rd_key;
simde__m128i x0, x1, x2;
rd_key[0] = x0 = userkey;
x2 = simde_mm_setzero_si128();
EXPAND_ASSIST(x0, x1, x2, x0, 255, 1);
rd_key[1] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 2);
rd_key[2] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 4);
rd_key[3] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 8);
rd_key[4] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 16);
rd_key[5] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 32);
rd_key[6] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 64);
rd_key[7] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 128);
rd_key[8] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 27);
rd_key[9] = x0;
EXPAND_ASSIST(x0, x1, x2, x0, 255, 54);
rd_key[10] = x0;
return rd_key;
}
}; // struct aes128_key
struct aes256_key
{
public:
static constexpr std::size_t rounds = 14;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using rd_key_array = std::array<simde__m128i, rounds+1>;
HEDLEY_PRAGMA(GCC diagnostic pop)
const rd_key_array rd_key;
explicit aes256_key(const simde__m256i & userkey)
: rd_key{compute_round_keys(userkey)} { }
private:
rd_key_array compute_round_keys(const simde__m256i & userkey)
{
rd_key_array rd_key;
simde__m128i x0, x1, x2, x3;
rd_key[0] = x0 = simde_mm256_extracti128_si256(userkey, 0);
rd_key[1] = x3 = simde_mm256_extracti128_si256(userkey, 1);
x2 = simde_mm_setzero_si128();
EXPAND_ASSIST(x0, x1, x2, x3, 255, 1);
rd_key[2] = x0;
EXPAND_ASSIST(x3, x1, x2, x0, 170, 1);
rd_key[3] = x3;
EXPAND_ASSIST(x0, x1, x2, x3, 255, 2);
rd_key[4] = x0;
EXPAND_ASSIST(x3, x1, x2, x0, 170, 2);
rd_key[5] = x3;
EXPAND_ASSIST(x0, x1, x2, x3, 255, 4);
rd_key[6] = x0;
EXPAND_ASSIST(x3, x1, x2, x0, 170, 4);
rd_key[7] = x3;
EXPAND_ASSIST(x0, x1, x2, x3, 255, 8);
rd_key[8] = x0;
EXPAND_ASSIST(x3, x1, x2, x0, 170, 8);
rd_key[9] = x3;
EXPAND_ASSIST(x0, x1, x2, x3, 255, 16);
rd_key[10] = x0;
EXPAND_ASSIST(x3, x1, x2, x0, 170, 16);
rd_key[11] = x3;
EXPAND_ASSIST(x0, x1, x2, x3, 255, 32);
rd_key[12] = x0;
EXPAND_ASSIST(x3, x1, x2, x0, 170, 32);
rd_key[13] = x3;
EXPAND_ASSIST(x0, x1, x2, x3, 255, 64);
rd_key[14] = x0;
return rd_key;
}
}; // struct aes256_key
using aes128 = aes<aes128_key>;
using aes256 = aes<aes256_key>;
template <>
const aes128_key aes128::key = aes128_key(simde__m128i{0, 0});
template <>
const aes256_key aes256::key = aes256_key(simde__m256i{0, 0, 0, 0});
} // namespace prg
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_PRG_AES_HPP__