libdpf/party/aes_mmo_ref.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

170 lines
4.5 KiB
C++

/// @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 <array>
#include <cstdint>
#include <cstring>
#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<std::uint8_t>(p ^ a);
const bool hi = (a & 0x80u) != 0;
a = static_cast<std::uint8_t>(a << 1);
if (hi)
a = static_cast<std::uint8_t>(a ^ 0x1b);
b = static_cast<std::uint8_t>(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<std::uint8_t>((b << 1) | (b >> 7));
s = static_cast<std::uint8_t>(s ^ b);
}
return static_cast<std::uint8_t>(s ^ 0x63);
}
inline std::uint8_t xtime(std::uint8_t x)
{
return static_cast<std::uint8_t>((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<std::uint8_t>(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<std::uint8_t>(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<std::uint8_t>(xa ^ xb ^ b ^ c0 ^ d);
s[i + 1] = static_cast<std::uint8_t>(a ^ xb ^ xc ^ c0 ^ d);
s[i + 2] = static_cast<std::uint8_t>(a ^ b ^ xc ^ xd ^ d);
s[i + 3] = static_cast<std::uint8_t>(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<std::uint8_t>(s[i] ^ rk[i]);
}
inline std::array<std::uint8_t, 16> block_bytes(simde__m128i v)
{
std::array<std::uint8_t, 16> 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<simde__m128i, 11> & 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<std::uint8_t>(rk0[0] ^ static_cast<std::uint8_t>(pos));
rk0[1] = static_cast<std::uint8_t>(rk0[1] ^ static_cast<std::uint8_t>(pos >> 8));
rk0[2] = static_cast<std::uint8_t>(rk0[2] ^ static_cast<std::uint8_t>(pos >> 16));
rk0[3] = static_cast<std::uint8_t>(rk0[3] ^ static_cast<std::uint8_t>(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<std::size_t>(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