/// @file dpf/prg_aes.hpp /// @brief /// @details /// @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 "hedley/hedley.h" #include "simde/simde/x86/avx2.h" #include "portable-snippets/exact-int/exact-int.h" #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; HEDLEY_NO_THROW HEDLEY_ALWAYS_INLINE static void require_block_aligned(const void * p) noexcept { assert(p == nullptr || reinterpret_cast(p) % alignof(block_type) == 0); } HEDLEY_NO_THROW HEDLEY_ALWAYS_INLINE HEDLEY_PURE static block_type eval(block_type seed, psnip_uint32_t pos) noexcept { 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 HEDLEY_PURE static auto eval01(block_type seed) noexcept { 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) } /// 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. HEDLEY_NO_THROW HEDLEY_ALWAYS_INLINE static void eval(block_type seed, block_type * HEDLEY_RESTRICT output, psnip_uint32_t count, psnip_uint32_t pos = 0) noexcept { if (HEDLEY_UNLIKELY(count == 0)) { return; } 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; } 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); } } /// Four independent `eval01` calls as one 8-block round-major AES. /// `left[i] == eval(seeds[i], 0)`, `right[i] == eval(seeds[i], 1)`. HEDLEY_NO_THROW HEDLEY_ALWAYS_INLINE static void eval01_x4(const block_type * HEDLEY_RESTRICT seeds, block_type * HEDLEY_RESTRICT left, block_type * HEDLEY_RESTRICT right) noexcept { 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]; } } /// Four independent `eval(seed, pos)` as one 4-block round-major AES. HEDLEY_NO_THROW HEDLEY_ALWAYS_INLINE static void eval_x4(const block_type * HEDLEY_RESTRICT seeds, block_type * HEDLEY_RESTRICT output, psnip_uint32_t pos) noexcept { 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]; } } /// Eight independent `eval(seed, pos)` as one 8-block round-major AES. HEDLEY_NO_THROW HEDLEY_ALWAYS_INLINE static void eval_x8(const block_type * HEDLEY_RESTRICT seeds, block_type * HEDLEY_RESTRICT output, psnip_uint32_t pos) noexcept { 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]; } } /// Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`). template static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept; private: static const AesKey key; /// `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]`. HEDLEY_NO_THROW HEDLEY_ALWAYS_INLINE 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 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__