libdpf/include/dpf/yao_aes.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

550 lines
17 KiB
C++
Raw Permalink 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/yao_aes.hpp
/// @brief Garbled AES-128 and the zero-key Matyas–Meyer–Oseas block.
/// @details Both are `netlist`s on `session::eval`. Input and output bits are
/// XOR shares. Within a byte, bit 0 of the 8-bit group is the MSB,
/// matching the Boyar–Peralta S-box in `aes_sbox_bp.hpp`. Bytes of a
/// block are in memory order: byte 0 is the first byte of an
/// `__m128i`. AES-128 includes the key schedule (6400 ANDs). The MMO
/// block uses the public all-zero key already installed in
/// `prg::aes128`, with `pos` mixed into the first round key the same
/// way `prg::aes128::eval` does, then XORs the input (5120 ANDs).
/// @copyright Copyright (c) 2019-2026 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_YAO_AES_HPP__
#define LIBDPF_INCLUDE_DPF_YAO_AES_HPP__
#include <array>
#include <cstdint>
#include <cstring>
#include <vector>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "dpf/aes_sbox_bp.hpp"
#include "dpf/yao.hpp"
namespace dpf
{
namespace yao
{
/// @brief Eight wires, index 0 = MSB.
struct bits8
{
std::array<bit, 8> b{};
};
/// @brief Pack one byte, MSB at `b[0]`, as eight shared inputs.
HEDLEY_WARN_UNUSED_RESULT
inline bits8 shared_byte(netlist & n)
{
bits8 o;
for (int i = 0; i < 8; ++i)
o.b[static_cast<std::size_t>(i)] = n.shared_in();
return o;
}
inline void out_byte(netlist & n, const bits8 & x)
{
for (int i = 0; i < 8; ++i)
n.out(x.b[static_cast<std::size_t>(i)]);
}
HEDLEY_WARN_UNUSED_RESULT
inline bits8 xor_byte(netlist & n, const bits8 & a, const bits8 & b)
{
bits8 o;
for (int i = 0; i < 8; ++i)
o.b[static_cast<std::size_t>(i)] = n.xor_(
a.b[static_cast<std::size_t>(i)], b.b[static_cast<std::size_t>(i)]);
return o;
}
/// @brief XOR a public byte into `x`. Bit 7 of `k` meets `b[0]`.
HEDLEY_WARN_UNUSED_RESULT
inline bits8 xor_public_byte(netlist & n, const bits8 & in, std::uint8_t k)
{
bits8 x = in;
for (int i = 0; i < 8; ++i)
{
if ((k >> (7 - i)) & 1u)
x.b[static_cast<std::size_t>(i)] =
n.not_(x.b[static_cast<std::size_t>(i)]);
}
return x;
}
/// @brief Boyar–Peralta S-box. 32 ANDs. `in[0]` is the MSB.
HEDLEY_WARN_UNUSED_RESULT
inline bits8 sbox(netlist & n, const bits8 & in)
{
std::array<bit, aes_bp::wire_count> w{};
for (int i = 0; i < 8; ++i)
w[static_cast<std::size_t>(i)] = in.b[static_cast<std::size_t>(i)];
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
{
const auto kind = aes_bp::ops[oi][0];
const auto dst = aes_bp::ops[oi][1];
const auto a = aes_bp::ops[oi][2];
const auto b = aes_bp::ops[oi][3];
const bit wa = w[a];
const bit wb = w[b];
bit d;
if (kind == 0)
d = n.xor_(wa, wb);
else if (kind == 1)
d = n.and_(wa, wb);
else
d = n.xnor_(wa, wb);
w[dst] = d;
}
bits8 o;
for (int i = 0; i < 8; ++i)
o.b[static_cast<std::size_t>(i)] =
w[aes_bp::out_wire[static_cast<std::size_t>(i)]];
return o;
}
/// @brief AES `xtime` (multiply by 2 in GF(2^8)). Linear, no AND.
HEDLEY_WARN_UNUSED_RESULT
inline bits8 xtime(netlist & n, const bits8 & x)
{
const bit msb = x.b[0];
bits8 y;
y.b[0] = x.b[1];
y.b[1] = x.b[2];
y.b[2] = x.b[3];
y.b[3] = n.xor_(x.b[4], msb);
y.b[4] = n.xor_(x.b[5], msb);
y.b[5] = x.b[6];
y.b[6] = n.xor_(x.b[7], msb);
y.b[7] = msb;
return y;
}
inline void shift_rows(bits8 (& s)[16])
{
bits8 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(netlist & n, bits8 (& s)[16])
{
for (int c = 0; c < 4; ++c)
{
const int i = c * 4;
const bits8 a = s[i];
const bits8 b = s[i + 1];
const bits8 c0 = s[i + 2];
const bits8 d = s[i + 3];
const bits8 xa = xtime(n, a);
const bits8 xb = xtime(n, b);
const bits8 xc = xtime(n, c0);
const bits8 xd = xtime(n, d);
s[i] = xor_byte(n, xor_byte(n, xor_byte(n, xor_byte(n, xa, xb), b), c0), d);
s[i + 1] = xor_byte(n, xor_byte(n, xor_byte(n, xor_byte(n, a, xb), xc), c0), d);
s[i + 2] = xor_byte(n, xor_byte(n, xor_byte(n, xor_byte(n, a, b), xc), xd), d);
s[i + 3] = xor_byte(n, xor_byte(n, xor_byte(n, xor_byte(n, xa, a), b), c0), xd);
}
}
inline void sub_bytes(netlist & n, bits8 (& s)[16])
{
for (int i = 0; i < 16; ++i)
s[i] = sbox(n, s[i]);
}
/// @brief 128 bits, byte 0 first, MSB of each byte first.
HEDLEY_NON_NULL(1)
inline void block_to_bits(std::uint8_t dst[128], block v) noexcept
{
std::uint8_t raw[16];
std::memcpy(raw, &v, 16);
for (int i = 0; i < 16; ++i)
{
for (int k = 0; k < 8; ++k)
dst[i * 8 + k] = static_cast<std::uint8_t>((raw[i] >> (7 - k)) & 1u);
}
}
HEDLEY_WARN_UNUSED_RESULT
HEDLEY_NON_NULL(1)
inline block bits_to_block(const std::uint8_t src[128]) noexcept
{
std::uint8_t raw[16]{};
for (int i = 0; i < 16; ++i)
{
for (int k = 0; k < 8; ++k)
raw[i] = static_cast<std::uint8_t>(
raw[i] | (src[i * 8 + k] << (7 - k)));
}
block v;
std::memcpy(&v, raw, 16);
return v;
}
namespace detail
{
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
constexpr std::uint8_t xtime_byte(std::uint8_t x) noexcept
{
return static_cast<std::uint8_t>((x << 1) ^ ((x & 0x80u) ? 0x1bu : 0));
}
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
constexpr std::uint8_t gmul(std::uint8_t a, std::uint8_t b) noexcept
{
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 ^ 0x1bu);
b = static_cast<std::uint8_t>(b >> 1);
}
return p;
}
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
constexpr std::uint8_t sbox_byte(std::uint8_t x) noexcept
{
std::uint8_t inv = 0;
if (x != 0)
{
std::uint8_t r = 1;
std::uint8_t b = x;
for (int e = 0; e < 7; ++e)
{
b = gmul(b, b);
r = gmul(r, b);
}
inv = r;
}
std::uint8_t s = inv;
std::uint8_t b = inv;
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);
}
/// @brief AES-128 key schedule for a public key. Used for the zero key of MMO.
inline void expand_public_key(const std::uint8_t key[16], std::uint8_t rk[11][16])
{
std::uint8_t w[176];
std::memcpy(w, key, 16);
int n = 16;
std::uint8_t rcon = 1;
while (n < 176)
{
std::uint8_t t[4];
std::memcpy(t, w + n - 4, 4);
if (n % 16 == 0)
{
const std::uint8_t tmp = t[0];
t[0] = t[1];
t[1] = t[2];
t[2] = t[3];
t[3] = tmp;
for (int i = 0; i < 4; ++i)
t[i] = sbox_byte(t[i]);
t[0] = static_cast<std::uint8_t>(t[0] ^ rcon);
rcon = xtime_byte(rcon);
}
for (int i = 0; i < 4; ++i)
{
w[n] = static_cast<std::uint8_t>(w[n - 16] ^ t[i]);
++n;
}
}
std::memcpy(rk, w, 176);
}
inline void read_shared_block(netlist & n, bits8 (& s)[16])
{
for (int i = 0; i < 16; ++i)
s[i] = shared_byte(n);
}
inline void emit_block(netlist & n, const bits8 (& s)[16])
{
for (int i = 0; i < 16; ++i)
out_byte(n, s[i]);
}
inline void add_round_key_public(netlist & n, bits8 (& s)[16],
const std::uint8_t (& rk)[16])
{
for (int i = 0; i < 16; ++i)
s[i] = xor_public_byte(n, s[i], rk[i]);
}
inline void add_round_key_secret(netlist & n, bits8 (& s)[16],
const bits8 (& rk)[16])
{
for (int i = 0; i < 16; ++i)
s[i] = xor_byte(n, s[i], rk[i]);
}
inline void aes_rounds_public_key(netlist & n, bits8 (& s)[16],
const std::uint8_t (& rk)[11][16])
{
add_round_key_public(n, s, rk[0]);
for (int r = 1; r < 10; ++r)
{
sub_bytes(n, s);
shift_rows(s);
mix_columns(n, s);
add_round_key_public(n, s, rk[r]);
}
sub_bytes(n, s);
shift_rows(s);
add_round_key_public(n, s, rk[10]);
}
inline void zero_key_schedule(std::uint8_t (& rk)[11][16])
{
const std::uint8_t zero_key[16]{};
expand_public_key(zero_key, rk);
}
/// @brief MMO under the zero key. `state` is overwritten with the digest.
/// `msg` is the shared block; `pos` is mixed into round key 0.
inline void apply_zero_key_mmo(netlist & n, bits8 (& state)[16],
const bits8 (& msg)[16], std::uint32_t pos,
const std::uint8_t (& base_rk)[11][16])
{
for (int i = 0; i < 16; ++i)
state[i] = msg[i];
std::uint8_t rk[11][16];
std::memcpy(rk, base_rk, sizeof(rk));
rk[0][0] = static_cast<std::uint8_t>(rk[0][0] ^ static_cast<std::uint8_t>(pos));
rk[0][1] = static_cast<std::uint8_t>(rk[0][1] ^ static_cast<std::uint8_t>(pos >> 8));
rk[0][2] = static_cast<std::uint8_t>(rk[0][2] ^ static_cast<std::uint8_t>(pos >> 16));
rk[0][3] = static_cast<std::uint8_t>(rk[0][3] ^ static_cast<std::uint8_t>(pos >> 24));
aes_rounds_public_key(n, state, rk);
for (int i = 0; i < 16; ++i)
state[i] = xor_byte(n, state[i], msg[i]);
}
} // namespace detail
/// @brief AES-128 of a shared block under a shared key. 128 block bits, then
/// 128 key bits, then 128 ciphertext bits. 6400 ANDs.
HEDLEY_WARN_UNUSED_RESULT
inline netlist aes128_netlist()
{
netlist n;
bits8 state[16];
bits8 key[16];
detail::read_shared_block(n, state);
detail::read_shared_block(n, key);
using word = std::array<bits8, 4>;
word words[44];
for (int i = 0; i < 4; ++i)
for (int j = 0; j < 4; ++j)
words[i][static_cast<std::size_t>(j)] = key[i * 4 + j];
std::uint8_t rcon = 0x01;
for (int wi = 4; wi < 44; ++wi)
{
word temp = words[wi - 1];
if (wi % 4 == 0)
{
const bits8 rot = temp[0];
temp[0] = temp[1];
temp[1] = temp[2];
temp[2] = temp[3];
temp[3] = rot;
for (int j = 0; j < 4; ++j)
temp[static_cast<std::size_t>(j)] =
sbox(n, temp[static_cast<std::size_t>(j)]);
temp[0] = xor_public_byte(n, temp[0], rcon);
rcon = detail::xtime_byte(rcon);
}
for (int j = 0; j < 4; ++j)
words[wi][static_cast<std::size_t>(j)] = xor_byte(n,
words[wi - 4][static_cast<std::size_t>(j)],
temp[static_cast<std::size_t>(j)]);
}
// AddRoundKey, then SubBytes / ShiftRows / MixColumns for rounds 0..8.
// Round 9 skips MixColumns. Round 10 is the last AddRoundKey.
for (int r = 0; r < 11; ++r)
{
bits8 rk[16];
for (int i = 0; i < 16; ++i)
rk[i] = words[r * 4 + i / 4][static_cast<std::size_t>(i % 4)];
detail::add_round_key_secret(n, state, rk);
if (r == 10)
break;
sub_bytes(n, state);
shift_rows(state);
if (r != 9)
mix_columns(n, state);
}
detail::emit_block(n, state);
return n;
}
/// @brief One MMO block: `AES_{k=0}(msg XOR pos) XOR msg`, matching
/// `prg::aes128::eval`. 128 shared input bits, 128 shared output bits.
/// `pos` is public and baked into the netlist.
HEDLEY_WARN_UNUSED_RESULT
inline netlist aes_mmo_netlist(std::uint32_t pos)
{
netlist n;
bits8 msg[16];
bits8 state[16];
detail::read_shared_block(n, msg);
for (int i = 0; i < 16; ++i)
state[i] = msg[i];
std::uint8_t rk[11][16];
detail::zero_key_schedule(rk);
detail::apply_zero_key_mmo(n, state, msg, pos, rk);
detail::emit_block(n, state);
return n;
}
/// @brief Garbled AES-128. `block_share` and `key_share` are 128 XOR shares
/// each, in `block_to_bits` order. Returns 128 XOR shares of the
/// ciphertext.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> aes128(session & s, int me, net::channel & link,
const std::uint8_t block_share[128], const std::uint8_t key_share[128])
{
std::vector<std::uint8_t> in(256);
std::memcpy(in.data(), block_share, 128);
std::memcpy(in.data() + 128, key_share, 128);
return s.eval(me, aes128_netlist(), in.data(), link);
}
/// @brief Garbled zero-key MMO. `msg_share` is 128 XOR shares in
/// `block_to_bits` order. `pos` is public and must match on both sides.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> aes_mmo(session & s, int me, net::channel & link,
const std::uint8_t msg_share[128], std::uint32_t pos)
{
return s.eval(me, aes_mmo_netlist(pos), msg_share, link);
}
/// @brief Shared bits of one correction-seed level: eight MMO blocks.
/// @details Matches `detail::vdpf::make_cs` / `party/oblivious_hash.hpp`.
/// Two seeds, four public lanes `0..3` each. 40960 ANDs.
/// The netlist does not depend on the level or the prefix. Those
/// are folded into the input shares by `pack_correction_level`.
inline constexpr std::size_t correction_level_in_bits = 8u * 128u;
inline constexpr std::size_t correction_level_out_bits = 4u * 128u;
inline constexpr std::uint32_t correction_level_ands =
static_cast<std::uint32_t>(8u * 16u * 10u * aes_bp::and_count);
/// @brief Eight zero-key MMO blocks, lanes 0..3 for each seed, then the XOR
/// of the two seeds' digests. 1024 input bits, 512 output bits.
HEDLEY_WARN_UNUSED_RESULT
inline netlist correction_level_netlist()
{
netlist n;
bits8 msg[8][16];
for (int b = 0; b < 8; ++b)
detail::read_shared_block(n, msg[b]);
std::uint8_t rk[11][16];
detail::zero_key_schedule(rk);
bits8 dig[8][16];
for (int b = 0; b < 8; ++b)
detail::apply_zero_key_mmo(n, dig[b], msg[b],
static_cast<std::uint32_t>(b & 3), rk);
for (int pos = 0; pos < 4; ++pos)
{
bits8 mixed[16];
for (int i = 0; i < 16; ++i)
mixed[i] = xor_byte(n, dig[pos][i], dig[4 + pos][i]);
detail::emit_block(n, mixed);
}
return n;
}
/// @brief This party's 1024 XOR shares for `correction_level_netlist`.
/// @details Party `me` holds `my_seed` in the clear. `prefix_share` is that
/// party's XOR share of the path prefix. The public tag
/// `0x5600 | (level & 0xffff)` is mixed into the owner's share only,
/// matching `hash_node` and `oblivious_cs`. Blocks are owner 0 lanes
/// 0..3, then owner 1 lanes 0..3.
inline void pack_correction_level(int me, block my_seed,
std::uint64_t prefix_share, std::size_t level,
std::uint8_t (& dst)[correction_level_in_bits])
{
if (me != 0 && me != 1)
throw std::invalid_argument("yao: party");
const std::uint64_t tag = 0x5600ull | (level & 0xffffu);
std::uint8_t seed_bytes[16];
std::memcpy(seed_bytes, &my_seed, 16);
for (int owner = 0; owner < 2; ++owner)
{
for (int pos = 0; pos < 4; ++pos)
{
std::uint8_t raw[16]{};
if (owner == me)
{
std::memcpy(raw, seed_bytes, 16);
for (int i = 0; i < 8; ++i)
raw[8 + i] = static_cast<std::uint8_t>(
raw[8 + i] ^ static_cast<std::uint8_t>(tag >> (8 * i)));
}
for (int i = 0; i < 8; ++i)
raw[i] = static_cast<std::uint8_t>(
raw[i] ^ static_cast<std::uint8_t>(prefix_share >> (8 * i)));
block blk;
std::memcpy(&blk, raw, 16);
const std::size_t off = static_cast<std::size_t>(owner * 4 + pos) * 128u;
block_to_bits(dst + off, blk);
}
}
}
/// @brief Garbled correction-seed level. Returns 512 XOR shares: four blocks,
/// lane 0 first. XOR of the two parties' outputs is `make_cs`.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> correction_level(session & s, int me,
net::channel & link, block my_seed, std::uint64_t prefix_share,
std::size_t level)
{
std::uint8_t in[correction_level_in_bits];
pack_correction_level(me, my_seed, prefix_share, level, in);
return s.eval(me, correction_level_netlist(), in, link);
}
} // namespace yao
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_YAO_AES_HPP__