libdpf/party/oblivious_hash.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

633 lines
22 KiB
C++
Raw 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 party/oblivious_hash.hpp
/// @brief Correction-seed hash that keeps the path prefix shared.
/// @details Matches `detail::vdpf::hash_node` without either party learning
/// the prefix. The opened value is only `H(s0) XOR H(s1)`.
/// SubBytes uses the Boyar–Peralta AES S-box (ePrint 2011/332;
/// Yale CMT SLP with 32 ANDs). Bit×bit ANDs use bit Beaver triples
/// and are scheduled by multiplicative depth (six exchanges per
/// SubBytes), so one level costs `8×16×10×32 = 40960` bit-ANDs.
#ifndef LIBDPF_PARTY_OBLIVIOUS_HASH_HPP__
#define LIBDPF_PARTY_OBLIVIOUS_HASH_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <vector>
#include "simde/simde/x86/avx2.h"
#include "aes_mmo_ref.hpp"
#include "aes_sbox_bp.hpp"
#include "dpf/doerner_shelat.hpp"
#include "dpf/prg_aes.hpp"
#include "dpf/verifiable.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/net/sink_exchange.hpp"
/* aes_bp is file-scope */
/// @brief AND-triples consumed by one level of `oblivious_cs`.
inline constexpr std::size_t hash_level_and_count() noexcept
{
// Eight AES (two seeds × four MMO lanes) × 16 bytes × 10 SubBytes
// × 32 Boyar–Peralta ANDs.
return 8u * 16u * 10u * aes_bp::and_count;
}
inline std::vector<bit_and_pad_msg> sample_and_tape(std::size_t n, int party)
{
::dpf::detail::urandom_pad_rng pads;
std::vector<bit_and_pad_msg> out(n);
for (std::size_t i = 0; i < n; ++i)
{
const auto full = ::dpf::detail::ds_sample_bit_and(pads);
out[i] = bit_pad_for(full, party);
}
return out;
}
inline std::vector<std::uint8_t> pack_mask_bits(
const std::uint8_t * d, const std::uint8_t * e, std::size_t n)
{
std::vector<std::uint8_t> out((2u * n + 7u) / 8u);
for (std::size_t i = 0; i < n; ++i)
{
const std::size_t db = 2u * i;
const std::size_t eb = db + 1u;
if (d[i] & 1u)
out[db / 8u] = static_cast<std::uint8_t>(
out[db / 8u] | static_cast<std::uint8_t>(1u << (db % 8u)));
if (e[i] & 1u)
out[eb / 8u] = static_cast<std::uint8_t>(
out[eb / 8u] | static_cast<std::uint8_t>(1u << (eb % 8u)));
}
return out;
}
inline void unpack_mask_bits(const std::vector<std::uint8_t> & packed,
std::uint8_t * d, std::uint8_t * e, std::size_t n)
{
if (packed.size() != (2u * n + 7u) / 8u)
throw std::runtime_error("oblivious AND packed size");
for (std::size_t i = 0; i < n; ++i)
{
const std::size_t db = 2u * i;
const std::size_t eb = db + 1u;
d[i] = static_cast<std::uint8_t>(
(packed[db / 8u] >> (db % 8u)) & 1u);
e[i] = static_cast<std::uint8_t>(
(packed[eb / 8u] >> (eb % 8u)) & 1u);
}
}
/// @brief One exchange of packed `(x⊕α, y⊕β)` masks for `n` bit-ANDs.
template <int Me>
std::vector<std::uint8_t> batch_bit_and(net::trio & net,
const std::uint8_t * x, const std::uint8_t * y, std::size_t n,
const bit_and_pad_msg * triples)
{
const role peer = Me == 0 ? role::p1 : role::p0;
std::vector<std::uint8_t> d(n), e(n);
for (std::size_t i = 0; i < n; ++i)
{
d[i] = static_cast<std::uint8_t>(x[i] ^ triples[i].a);
e[i] = static_cast<std::uint8_t>(y[i] ^ triples[i].b);
}
const auto mine = pack_mask_bits(d.data(), e.data(), n);
std::vector<std::uint8_t> peer_bytes;
if constexpr (Me == 0)
{
net.send_bytes_to(peer, net::msg::delta, mine.data(), mine.size());
peer_bytes = net.recv_bytes_from(peer, net::msg::delta);
}
else
{
peer_bytes = net.recv_bytes_from(peer, net::msg::delta);
net.send_bytes_to(peer, net::msg::delta, mine.data(), mine.size());
}
std::vector<std::uint8_t> pd(n), pe(n);
unpack_mask_bits(peer_bytes, pd.data(), pe.data(), n);
// Standard Beaver: z = [c] ⊕ (d∧[b]) ⊕ ([a]∧e) ⊕ (d∧e), and the public
// `d∧e` term is added on one party only. Adding it on both cancels under
// XOR and the S-box (and every correction seed) diverges from `hash_node`.
std::vector<std::uint8_t> shares(n);
for (std::size_t i = 0; i < n; ++i)
{
const std::uint8_t od = static_cast<std::uint8_t>(d[i] ^ pd[i]);
const std::uint8_t oe = static_cast<std::uint8_t>(e[i] ^ pe[i]);
shares[i] = detail::ds_bit_and_party(
od, oe, triples[i].a, triples[i].b, triples[i].c, Me == 0);
}
return shares;
}
/// @brief Packed bit-AND layer on a RoundSink. Slot must fit `4 + packed`.
template <int Me>
std::vector<std::uint8_t> batch_bit_and_on_sink(dpf::net::RoundSink & sink,
std::uint16_t & round, std::size_t index, const std::uint8_t * x,
const std::uint8_t * y, std::size_t n, const bit_and_pad_msg * triples)
{
std::vector<std::uint8_t> d(n), e(n);
for (std::size_t i = 0; i < n; ++i)
{
d[i] = static_cast<std::uint8_t>(x[i] ^ triples[i].a);
e[i] = static_cast<std::uint8_t>(y[i] ^ triples[i].b);
}
const auto mine = pack_mask_bits(d.data(), e.data(), n);
if (round >= sink.rounds())
throw std::runtime_error("batch_bit_and_on_sink: out of rounds");
const std::size_t slot = sink.slot_bytes(round);
const std::size_t need = sizeof(std::uint32_t) + mine.size();
if (need > slot)
throw std::runtime_error("batch_bit_and_on_sink: slot too small");
std::vector<std::uint8_t> buf(slot, 0);
const std::uint32_t nbytes = static_cast<std::uint32_t>(mine.size());
std::memcpy(buf.data(), &nbytes, 4);
if (!mine.empty())
std::memcpy(buf.data() + 4, mine.data(), mine.size());
sink.submit(round, index, buf.data(), buf.size());
sink.flush_round(round);
sink.poll();
if (!sink.peer_ready(round, index))
throw std::runtime_error("batch_bit_and_on_sink: peer missing");
std::vector<std::uint8_t> peer_buf(slot);
sink.read_peer(round, index, peer_buf.data(), slot);
std::uint32_t pn = 0;
std::memcpy(&pn, peer_buf.data(), 4);
if (pn != nbytes)
throw std::runtime_error("batch_bit_and_on_sink: peer size");
std::vector<std::uint8_t> peer_bytes(pn);
if (pn != 0)
std::memcpy(peer_bytes.data(), peer_buf.data() + 4, pn);
++round;
std::vector<std::uint8_t> pd(n), pe(n);
unpack_mask_bits(peer_bytes, pd.data(), pe.data(), n);
std::vector<std::uint8_t> shares(n);
for (std::size_t i = 0; i < n; ++i)
{
const std::uint8_t od = static_cast<std::uint8_t>(d[i] ^ pd[i]);
const std::uint8_t oe = static_cast<std::uint8_t>(e[i] ^ pe[i]);
shares[i] = detail::ds_bit_and_party(
od, oe, triples[i].a, triples[i].b, triples[i].c, Me == 0);
}
return shares;
}
/// @brief Shared SubBytes via Boyar–Peralta (32 ANDs per byte, depth-batched).
template <int Me>
void sub_bytes_shared(net::trio & net, std::uint8_t * state, std::size_t bytes,
const bit_and_pad_msg * & tape)
{
std::vector<std::uint8_t> wire(bytes * aes_bp::wire_count);
auto at = [&](std::size_t b, std::size_t w) -> std::uint8_t & {
return wire[b * aes_bp::wire_count + w];
};
for (std::size_t b = 0; b < bytes; ++b)
{
for (int i = 0; i < 8; ++i)
at(b, static_cast<std::size_t>(i)) =
static_cast<std::uint8_t>((state[b] >> (7 - i)) & 1u);
}
std::array<bool, aes_bp::wire_count> ready{};
for (std::size_t i = 0; i < 8; ++i)
ready[i] = true;
std::array<bool, aes_bp::op_count> done{};
std::size_t finished = 0;
while (finished < aes_bp::op_count)
{
bool progress = true;
while (progress)
{
progress = false;
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
{
if (done[oi])
continue;
const auto kind = aes_bp::ops[oi][0];
if (kind == 1)
continue;
const auto dst = aes_bp::ops[oi][1];
const auto a = aes_bp::ops[oi][2];
const auto bb = aes_bp::ops[oi][3];
if (!ready[a] || !ready[bb])
continue;
for (std::size_t b = 0; b < bytes; ++b)
{
auto v = static_cast<std::uint8_t>(at(b, a) ^ at(b, bb));
if (kind == 2)
{
if constexpr (Me == 0)
v = static_cast<std::uint8_t>(v ^ 1u);
}
at(b, dst) = v;
}
ready[dst] = true;
done[oi] = true;
++finished;
progress = true;
}
}
std::vector<std::size_t> ands;
ands.reserve(aes_bp::and_count);
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
{
if (done[oi] || aes_bp::ops[oi][0] != 1)
continue;
const auto a = aes_bp::ops[oi][2];
const auto bb = aes_bp::ops[oi][3];
if (ready[a] && ready[bb])
ands.push_back(oi);
}
if (ands.empty())
{
if (finished != aes_bp::op_count)
throw std::runtime_error("oblivious SubBytes stuck");
break;
}
const std::size_t n = ands.size() * bytes;
std::vector<std::uint8_t> xs(n), ys(n);
for (std::size_t ai = 0; ai < ands.size(); ++ai)
{
const auto a = aes_bp::ops[ands[ai]][2];
const auto bb = aes_bp::ops[ands[ai]][3];
for (std::size_t b = 0; b < bytes; ++b)
{
xs[ai * bytes + b] = at(b, a);
ys[ai * bytes + b] = at(b, bb);
}
}
const auto prod = batch_bit_and<Me>(net, xs.data(), ys.data(), n, tape);
tape += n;
for (std::size_t ai = 0; ai < ands.size(); ++ai)
{
const auto dst = aes_bp::ops[ands[ai]][1];
for (std::size_t b = 0; b < bytes; ++b)
at(b, dst) = static_cast<std::uint8_t>(prod[ai * bytes + b] & 1u);
ready[dst] = true;
done[ands[ai]] = true;
++finished;
}
}
for (std::size_t b = 0; b < bytes; ++b)
{
std::uint8_t out = 0;
for (int i = 0; i < 8; ++i)
out = static_cast<std::uint8_t>(
out | (at(b, aes_bp::out_wire[static_cast<std::size_t>(i)])
<< (7 - i)));
state[b] = out;
}
}
/// @brief SubBytes with AND layers on a RoundSink (multi-instance friendly).
template <int Me>
void sub_bytes_shared_on_sink(dpf::net::RoundSink & sink, std::uint16_t & round,
std::size_t index, std::uint8_t * state, std::size_t bytes,
const bit_and_pad_msg * & tape)
{
std::vector<std::uint8_t> wire(bytes * aes_bp::wire_count);
auto at = [&](std::size_t b, std::size_t w) -> std::uint8_t & {
return wire[b * aes_bp::wire_count + w];
};
for (std::size_t b = 0; b < bytes; ++b)
{
for (int i = 0; i < 8; ++i)
at(b, static_cast<std::size_t>(i)) =
static_cast<std::uint8_t>((state[b] >> (7 - i)) & 1u);
}
std::array<bool, aes_bp::wire_count> ready{};
for (std::size_t i = 0; i < 8; ++i)
ready[i] = true;
std::array<bool, aes_bp::op_count> done{};
std::size_t finished = 0;
while (finished < aes_bp::op_count)
{
bool progress = true;
while (progress)
{
progress = false;
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
{
if (done[oi])
continue;
const auto kind = aes_bp::ops[oi][0];
if (kind == 1)
continue;
const auto dst = aes_bp::ops[oi][1];
const auto a = aes_bp::ops[oi][2];
const auto bb = aes_bp::ops[oi][3];
if (!ready[a] || !ready[bb])
continue;
for (std::size_t b = 0; b < bytes; ++b)
{
auto v = static_cast<std::uint8_t>(at(b, a) ^ at(b, bb));
if (kind == 2)
{
if constexpr (Me == 0)
v = static_cast<std::uint8_t>(v ^ 1u);
}
at(b, dst) = v;
}
ready[dst] = true;
done[oi] = true;
++finished;
progress = true;
}
}
std::vector<std::size_t> ands;
ands.reserve(aes_bp::and_count);
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
{
if (done[oi] || aes_bp::ops[oi][0] != 1)
continue;
const auto a = aes_bp::ops[oi][2];
const auto bb = aes_bp::ops[oi][3];
if (ready[a] && ready[bb])
ands.push_back(oi);
}
if (ands.empty())
{
if (finished != aes_bp::op_count)
throw std::runtime_error("oblivious SubBytes stuck");
break;
}
const std::size_t n = ands.size() * bytes;
std::vector<std::uint8_t> xs(n), ys(n);
for (std::size_t ai = 0; ai < ands.size(); ++ai)
{
const auto a = aes_bp::ops[ands[ai]][2];
const auto bb = aes_bp::ops[ands[ai]][3];
for (std::size_t b = 0; b < bytes; ++b)
{
xs[ai * bytes + b] = at(b, a);
ys[ai * bytes + b] = at(b, bb);
}
}
const auto prod = batch_bit_and_on_sink<Me>(
sink, round, index, xs.data(), ys.data(), n, tape);
tape += n;
for (std::size_t ai = 0; ai < ands.size(); ++ai)
{
const auto dst = aes_bp::ops[ands[ai]][1];
for (std::size_t b = 0; b < bytes; ++b)
at(b, dst) = static_cast<std::uint8_t>(prod[ai * bytes + b] & 1u);
ready[dst] = true;
done[ands[ai]] = true;
++finished;
}
}
for (std::size_t b = 0; b < bytes; ++b)
{
std::uint8_t out = 0;
for (int i = 0; i < 8; ++i)
out = static_cast<std::uint8_t>(
out | (at(b, aes_bp::out_wire[static_cast<std::size_t>(i)])
<< (7 - i)));
state[b] = out;
}
}
inline void shift_rows_bytes(std::uint8_t s[16])
{
::dpf::party::aes_ref::shift_rows(s);
}
inline void mix_columns_bytes(std::uint8_t s[16])
{
::dpf::party::aes_ref::mix_columns(s);
}
inline void xor_bytes(std::uint8_t * dst, const std::uint8_t * src, std::size_t n)
{
for (std::size_t i = 0; i < n; ++i)
dst[i] = static_cast<std::uint8_t>(dst[i] ^ src[i]);
}
inline std::array<std::uint8_t, 16> m128_bytes(simde__m128i v)
{
std::array<std::uint8_t, 16> b{};
std::memcpy(b.data(), &v, 16);
return b;
}
/// @brief Public correction seed for one level. Prefix shares stay local.
template <int Me>
cs_block oblivious_cs(net::trio & net, std::size_t level,
std::uint64_t prefix_share, simde__m128i my_seed,
const bit_and_pad_msg * tape)
{
const auto schedule = prg::aes128_key(simde_mm_setzero_si128());
const auto rk = [&](int r) {
return m128_bytes(schedule.rd_key[static_cast<std::size_t>(r)]);
};
constexpr std::size_t blocks = 8;
std::uint8_t state[blocks][16]{};
std::uint8_t feed[blocks][16]{};
auto place_prefix = [&](std::uint8_t * dst) {
for (int i = 0; i < 8; ++i)
dst[i] = static_cast<std::uint8_t>(
dst[i] ^ static_cast<std::uint8_t>(prefix_share >> (8 * i)));
};
for (int owner = 0; owner < 2; ++owner)
{
const bool mine = (owner == Me);
for (int pos = 0; pos < 4; ++pos)
{
const std::size_t bi = static_cast<std::size_t>(owner * 4 + pos);
if (mine)
{
auto bytes = m128_bytes(my_seed);
const std::uint64_t tag = 0x5600ull | (level & 0xffffu);
for (int i = 0; i < 8; ++i)
bytes[8 + i] = static_cast<std::uint8_t>(
bytes[8 + i]
^ static_cast<std::uint8_t>(tag >> (8 * i)));
place_prefix(bytes.data());
std::memcpy(feed[bi], bytes.data(), 16);
std::memcpy(state[bi], bytes.data(), 16);
}
else
{
std::uint8_t bytes[16]{};
place_prefix(bytes);
std::memcpy(feed[bi], bytes, 16);
std::memcpy(state[bi], bytes, 16);
}
if constexpr (Me == 0)
{
auto key0 = rk(0);
key0[0] = static_cast<std::uint8_t>(
key0[0] ^ static_cast<std::uint8_t>(pos));
key0[1] = static_cast<std::uint8_t>(
key0[1]
^ static_cast<std::uint8_t>(static_cast<unsigned>(pos) >> 8));
key0[2] = static_cast<std::uint8_t>(
key0[2]
^ static_cast<std::uint8_t>(
static_cast<unsigned>(pos) >> 16));
key0[3] = static_cast<std::uint8_t>(
key0[3]
^ static_cast<std::uint8_t>(
static_cast<unsigned>(pos) >> 24));
xor_bytes(state[bi], key0.data(), 16);
}
}
}
auto all_bytes = [&]() -> std::uint8_t * { return &state[0][0]; };
const bit_and_pad_msg * cursor = tape;
for (int round = 1; round <= 10; ++round)
{
sub_bytes_shared<Me>(net, all_bytes(), blocks * 16, cursor);
for (std::size_t b = 0; b < blocks; ++b)
shift_rows_bytes(state[b]);
if (round < 10)
{
for (std::size_t b = 0; b < blocks; ++b)
mix_columns_bytes(state[b]);
}
if constexpr (Me == 0)
{
const auto key = rk(round);
for (std::size_t b = 0; b < blocks; ++b)
xor_bytes(state[b], key.data(), 16);
}
}
if (cursor != tape + hash_level_and_count())
throw std::runtime_error("oblivious hash consumed the wrong AND count");
for (std::size_t b = 0; b < blocks; ++b)
xor_bytes(state[b], feed[b], 16);
cs_block mine{};
for (int pos = 0; pos < 4; ++pos)
{
std::uint8_t mixed[16]{};
xor_bytes(mixed, state[static_cast<std::size_t>(pos)], 16);
xor_bytes(mixed, state[static_cast<std::size_t>(4 + pos)], 16);
std::memcpy(&mine[static_cast<std::size_t>(pos)], mixed, 16);
}
return open_cs<Me>(net, mine);
}
/// @brief `oblivious_cs` with AND layers on a RoundSink; final open stays on trio.
template <int Me>
cs_block oblivious_cs_on_sink(net::trio & net, dpf::net::RoundSink & sink,
std::uint16_t & round, std::size_t index, std::size_t level,
std::uint64_t prefix_share, simde__m128i my_seed,
const bit_and_pad_msg * tape)
{
const auto schedule = prg::aes128_key(simde_mm_setzero_si128());
const auto rk = [&](int r) {
return m128_bytes(schedule.rd_key[static_cast<std::size_t>(r)]);
};
constexpr std::size_t blocks = 8;
std::uint8_t state[blocks][16]{};
std::uint8_t feed[blocks][16]{};
auto place_prefix = [&](std::uint8_t * dst) {
for (int i = 0; i < 8; ++i)
dst[i] = static_cast<std::uint8_t>(
dst[i] ^ static_cast<std::uint8_t>(prefix_share >> (8 * i)));
};
for (int owner = 0; owner < 2; ++owner)
{
const bool mine = (owner == Me);
for (int pos = 0; pos < 4; ++pos)
{
const std::size_t bi = static_cast<std::size_t>(owner * 4 + pos);
if (mine)
{
auto bytes = m128_bytes(my_seed);
const std::uint64_t tag = 0x5600ull | (level & 0xffffu);
for (int i = 0; i < 8; ++i)
bytes[8 + i] = static_cast<std::uint8_t>(
bytes[8 + i]
^ static_cast<std::uint8_t>(tag >> (8 * i)));
place_prefix(bytes.data());
std::memcpy(feed[bi], bytes.data(), 16);
std::memcpy(state[bi], bytes.data(), 16);
}
else
{
std::uint8_t bytes[16]{};
place_prefix(bytes);
std::memcpy(feed[bi], bytes, 16);
std::memcpy(state[bi], bytes, 16);
}
if constexpr (Me == 0)
{
auto key0 = rk(0);
key0[0] = static_cast<std::uint8_t>(
key0[0] ^ static_cast<std::uint8_t>(pos));
key0[1] = static_cast<std::uint8_t>(
key0[1]
^ static_cast<std::uint8_t>(static_cast<unsigned>(pos) >> 8));
key0[2] = static_cast<std::uint8_t>(
key0[2]
^ static_cast<std::uint8_t>(
static_cast<unsigned>(pos) >> 16));
key0[3] = static_cast<std::uint8_t>(
key0[3]
^ static_cast<std::uint8_t>(
static_cast<unsigned>(pos) >> 24));
xor_bytes(state[bi], key0.data(), 16);
}
}
}
auto all_bytes = [&]() -> std::uint8_t * { return &state[0][0]; };
const bit_and_pad_msg * cursor = tape;
for (int aes_round = 1; aes_round <= 10; ++aes_round)
{
sub_bytes_shared_on_sink<Me>(
sink, round, index, all_bytes(), blocks * 16, cursor);
for (std::size_t b = 0; b < blocks; ++b)
shift_rows_bytes(state[b]);
if (aes_round < 10)
{
for (std::size_t b = 0; b < blocks; ++b)
mix_columns_bytes(state[b]);
}
if constexpr (Me == 0)
{
const auto key = rk(aes_round);
for (std::size_t b = 0; b < blocks; ++b)
xor_bytes(state[b], key.data(), 16);
}
}
if (cursor != tape + hash_level_and_count())
throw std::runtime_error("oblivious hash consumed the wrong AND count");
for (std::size_t b = 0; b < blocks; ++b)
xor_bytes(state[b], feed[b], 16);
cs_block mine{};
for (int pos = 0; pos < 4; ++pos)
{
std::uint8_t mixed[16]{};
xor_bytes(mixed, state[static_cast<std::size_t>(pos)], 16);
xor_bytes(mixed, state[static_cast<std::size_t>(4 + pos)], 16);
std::memcpy(&mine[static_cast<std::size_t>(pos)], mixed, 16);
}
return open_cs<Me>(net, mine);
}
#endif