/// @file party/aes_mmo_ref.hpp /// @brief Portable AES-128 MMO matching `dpf::prg::aes128::eval` on the zero key. /// @details Used to check an oblivious evaluation against the hardware PRG. /// The cipher key is the all-zero key installed by `prg::aes128`. #ifndef LIBDPF_PARTY_AES_MMO_REF_HPP__ #define LIBDPF_PARTY_AES_MMO_REF_HPP__ #include #include #include #include "simde/simde/x86/avx2.h" namespace dpf { namespace party { namespace aes_ref { inline std::uint8_t gmul(std::uint8_t a, std::uint8_t b) { std::uint8_t p = 0; for (int i = 0; i < 8; ++i) { if (b & 1u) p = static_cast(p ^ a); const bool hi = (a & 0x80u) != 0; a = static_cast(a << 1); if (hi) a = static_cast(a ^ 0x1b); b = static_cast(b >> 1); } return p; } inline std::uint8_t ginv(std::uint8_t a) { if (a == 0) return 0; // a^254 via square-and-multiply. std::uint8_t r = 1; std::uint8_t b = a; for (int e = 0; e < 7; ++e) { b = gmul(b, b); r = gmul(r, b); } return r; } inline std::uint8_t sbox_at(std::uint8_t x) { std::uint8_t b = ginv(x); std::uint8_t s = b; for (int i = 0; i < 4; ++i) { b = static_cast((b << 1) | (b >> 7)); s = static_cast(s ^ b); } return static_cast(s ^ 0x63); } inline std::uint8_t xtime(std::uint8_t x) { return static_cast((x << 1) ^ ((x & 0x80) ? 0x1b : 0)); } inline const std::uint8_t * sbox() { static std::uint8_t table[256]; static bool ready = false; if (!ready) { for (int i = 0; i < 256; ++i) table[i] = sbox_at(static_cast(i)); ready = true; } return table; } inline void sub_bytes(std::uint8_t s[16]) { const std::uint8_t * box = sbox(); for (int i = 0; i < 16; ++i) s[i] = box[s[i]]; } inline void shift_rows(std::uint8_t s[16]) { std::uint8_t t = s[1]; s[1] = s[5]; s[5] = s[9]; s[9] = s[13]; s[13] = t; t = s[2]; s[2] = s[10]; s[10] = t; t = s[6]; s[6] = s[14]; s[14] = t; t = s[15]; s[15] = s[11]; s[11] = s[7]; s[7] = s[3]; s[3] = t; } inline void mix_columns(std::uint8_t s[16]) { for (int c = 0; c < 4; ++c) { const std::uint8_t i = static_cast(c * 4); const std::uint8_t a = s[i], b = s[i + 1], c0 = s[i + 2], d = s[i + 3]; const std::uint8_t xa = xtime(a), xb = xtime(b), xc = xtime(c0), xd = xtime(d); s[i] = static_cast(xa ^ xb ^ b ^ c0 ^ d); s[i + 1] = static_cast(a ^ xb ^ xc ^ c0 ^ d); s[i + 2] = static_cast(a ^ b ^ xc ^ xd ^ d); s[i + 3] = static_cast(xa ^ a ^ b ^ c0 ^ xd); } } inline void add_round_key(std::uint8_t s[16], const std::uint8_t rk[16]) { for (int i = 0; i < 16; ++i) s[i] = static_cast(s[i] ^ rk[i]); } inline std::array block_bytes(simde__m128i v) { std::array b{}; std::memcpy(b.data(), &v, 16); return b; } /// @brief One MMO block: AES-128 encrypt `msg` under the zero key, then XOR `msg`. inline simde__m128i encrypt_mmo(simde__m128i msg, std::uint32_t pos, const std::array & round_keys) { auto s = block_bytes(msg); auto rk0 = block_bytes(round_keys[0]); // eval XORs set_epi64x(0, pos) into rd_key[0] before the first AddRoundKey. rk0[0] = static_cast(rk0[0] ^ static_cast(pos)); rk0[1] = static_cast(rk0[1] ^ static_cast(pos >> 8)); rk0[2] = static_cast(rk0[2] ^ static_cast(pos >> 16)); rk0[3] = static_cast(rk0[3] ^ static_cast(pos >> 24)); add_round_key(s.data(), rk0.data()); for (int r = 1; r < 10; ++r) { sub_bytes(s.data()); shift_rows(s.data()); mix_columns(s.data()); auto rk = block_bytes(round_keys[static_cast(r)]); add_round_key(s.data(), rk.data()); } sub_bytes(s.data()); shift_rows(s.data()); auto last = block_bytes(round_keys[10]); add_round_key(s.data(), last.data()); simde__m128i out{}; std::memcpy(&out, s.data(), 16); return simde_mm_xor_si128(out, msg); } } // namespace aes_ref } // namespace party } // namespace dpf #endif