libdpf/include/dpf/prg_aes.hpp

532 lines
20 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @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/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;
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static void require_block_aligned(const void * p) noexcept
{
assert(p == nullptr
|| reinterpret_cast<std::uintptr_t>(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<block_type, 2>{output0, output1};
HEDLEY_PRAGMA(GCC diagnostic pop)
}
/// @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;
}
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
{
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
{
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
{
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__