/// @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 #include #include #include #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 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(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(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(i)] = n.xor_( a.b[static_cast(i)], b.b[static_cast(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(i)] = n.not_(x.b[static_cast(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 w{}; for (int i = 0; i < 8; ++i) w[static_cast(i)] = in.b[static_cast(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(i)] = w[aes_bp::out_wire[static_cast(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((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( 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((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(p ^ a); const bool hi = (a & 0x80u) != 0; a = static_cast(a << 1); if (hi) a = static_cast(a ^ 0x1bu); b = static_cast(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((b << 1) | (b >> 7)); s = static_cast(s ^ b); } return static_cast(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(t[0] ^ rcon); rcon = xtime_byte(rcon); } for (int i = 0; i < 4; ++i) { w[n] = static_cast(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(rk[0][0] ^ static_cast(pos)); rk[0][1] = static_cast(rk[0][1] ^ static_cast(pos >> 8)); rk[0][2] = static_cast(rk[0][2] ^ static_cast(pos >> 16)); rk[0][3] = static_cast(rk[0][3] ^ static_cast(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; word words[44]; for (int i = 0; i < 4; ++i) for (int j = 0; j < 4; ++j) words[i][static_cast(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(j)] = sbox(n, temp[static_cast(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(j)] = xor_byte(n, words[wi - 4][static_cast(j)], temp[static_cast(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(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 aes128(session & s, int me, net::channel & link, const std::uint8_t block_share[128], const std::uint8_t key_share[128]) { std::vector 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 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(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(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( raw[8 + i] ^ static_cast(tag >> (8 * i))); } for (int i = 0; i < 8; ++i) raw[i] = static_cast( raw[i] ^ static_cast(prefix_share >> (8 * i))); block blk; std::memcpy(&blk, raw, 16); const std::size_t off = static_cast(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 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__