/// @file dpf/prg_aes.hpp /// @brief Fixed-key AES Matyas–Meyer–Oseas PRG. /// @author Ryan Henry /// @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 #include #include #include #include #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 #endif template 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(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{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(~static_cast(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(__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(__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(__builtin_assume_aligned(seeds, alignof(block_type))); block_type * HEDLEY_RESTRICT out = static_cast(__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(__builtin_assume_aligned(seeds, alignof(block_type))); block_type * HEDLEY_RESTRICT out = static_cast(__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 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; 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; 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; using aes256 = aes; 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__