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>
170 lines
4.5 KiB
C++
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
|