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

686 lines
21 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 dpf/yao.hpp
/// @brief Semi-honest garbled netlist for a boolean function of a DPF leaf.
/// @details The key still computes the point, the comparison, the interval, and
/// a public-offset polynomial. This header is the step after a leaf
/// share exists and the next function is a straight-line bit circuit
/// the key does not contain. The circuit this library is built around
/// is the zero-key AES-MMO block the PRG already uses, so a seed or a
/// leaf can enter that block without being opened.
///
/// Free-XOR and half-gates (Zahur, Rosulek, and Evans, ePrint 2014/756).
/// A secret if/else or a k-way switch is `dpf/yao_stack.hpp`: stacked
/// garbling (Heath and Kolesnikov, CRYPTO 2020) and one-hot garbling
/// (Heath and Kolesnikov, CCS 2021). The netlist below is unchanged.
/// Party 0 garbles. Party 1 evaluates. A shared input is the XOR of
/// the two parties' bits. A private input is known to one party.
/// Each output is an XOR share: the garbler's share is the permute
/// bit of the zero-label, and the evaluator's share is the color of
/// the label it holds. XOR of the two shares is the output bit.
///
/// The gate hash is fixed-key AES in the same Matyas–Meyer–Oseas
/// shape as `iknp::detail::ot_hash`, with a distinct tweak. Tables
/// are one-time. Gate ids increase by one per AND and are never
/// reused. There is no Yao domain on the composer. A ring share
/// becomes these bits through `dpf/yao_share.hpp`, then comes back
/// the same way. Comparisons and public-offset LUTs stay on the key.
///
/// `session` keeps the IKNP base OT. The first `eval` that needs a
/// choice label runs Chou–Orlandi; later evals on that session only
/// extend. Do not interleave `iknp::sample` on the same channel.
///
/// Ideal functionality (semi-honest): both parties input the bits named by
/// the netlist; each receives an XOR share of every output wire. Party 0
/// does not learn party 1's private bits or party 1's shares. Party 1 does
/// not learn party 0's private bits, party 0's shares, or `Δ`.
/// @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_HPP__
#define LIBDPF_INCLUDE_DPF_YAO_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "dpf/iknp.hpp"
#include "dpf/prg_aes.hpp"
#include "dpf/random.hpp"
namespace dpf
{
namespace yao
{
using block = simde__m128i;
/// @brief One wire in a `netlist`. Valid only for the netlist that minted it.
struct bit
{
std::uint32_t id = 0;
};
/// @brief Who knows an input wire.
enum class input : unsigned char
{
shared = 0, ///< XOR of the two parties' bits
priv0 = 1, ///< party 0's private bit
priv1 = 2 ///< party 1's private bit
};
namespace detail
{
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
block xor_b(block a, block b) noexcept
{
return simde_mm_xor_si128(a, b);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
block zero_b() noexcept
{
return simde_mm_setzero_si128();
}
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
std::uint8_t lsb(block x) noexcept
{
unsigned char b = 0;
std::memcpy(&b, &x, 1);
return static_cast<std::uint8_t>(b & 1u);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
block with_lsb(block x, std::uint8_t bit) noexcept
{
unsigned char b = 0;
std::memcpy(&b, &x, 1);
b = static_cast<unsigned char>((b & 0xfeu) | (bit & 1u));
std::memcpy(&x, &b, 1);
return x;
}
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
block mask_bit(std::uint8_t bit, block delta) noexcept
{
return (bit & 1u) ? delta : zero_b();
}
/// @brief Correlation-robust hash for one half-gate. Tweak `2·gid` is the
/// left wire and `2·gid+1` is the right wire. Both labels of one wire
/// use the same tweak.
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block cr_hash(block x, std::uint64_t tweak) noexcept
{
const prg::purpose_scope counted(prg::purpose::hash);
const auto mixed = xor_b(x, simde_mm_set_epi64x(
static_cast<std::int64_t>(tweak >> 32),
static_cast<std::int64_t>(static_cast<std::uint32_t>(tweak))));
return prg::aes128::eval(mixed,
static_cast<psnip_uint32_t>(tweak * 0x9E3779B9u) ^ 0xC3A5u);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block random_delta() noexcept
{
return with_lsb(dpf::uniform_sample<block>(), 1);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block random_label() noexcept
{
return dpf::uniform_sample<block>();
}
/// @brief Half-gates AND. `la0` and `lb0` are zero-labels. `table[0]` and
/// `table[1]` are the two ciphertexts. Returns the output zero-label.
/// Matches `halfgates_garble` / `halfgates_eval` in EMP-toolkit.
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block and_garble(block la0, block lb0, block delta, std::uint64_t gid,
block table[2]) noexcept
{
const bool pa = lsb(la0) != 0;
const bool pb = lsb(lb0) != 0;
const block ha0 = cr_hash(la0, gid * 2u);
const block ha1 = cr_hash(xor_b(la0, delta), gid * 2u);
const block hb0 = cr_hash(lb0, gid * 2u + 1u);
const block hb1 = cr_hash(xor_b(lb0, delta), gid * 2u + 1u);
table[0] = xor_b(xor_b(ha0, ha1), pb ? delta : zero_b());
block w0 = xor_b(ha0, pa ? table[0] : zero_b());
const block tmp = xor_b(hb0, hb1);
table[1] = xor_b(tmp, la0);
w0 = xor_b(w0, hb0);
w0 = xor_b(w0, pb ? tmp : zero_b());
return w0;
}
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block and_eval(block a, block b, std::uint64_t gid, const block table[2]) noexcept
{
const bool sa = lsb(a) != 0;
const bool sb = lsb(b) != 0;
block w = xor_b(cr_hash(a, gid * 2u), cr_hash(b, gid * 2u + 1u));
if (sa)
w = xor_b(w, table[0]);
if (sb)
{
w = xor_b(w, table[1]);
w = xor_b(w, a);
}
return w;
}
inline void require_bit(std::uint8_t b, const char * what)
{
if (b > 1u)
throw std::invalid_argument(what);
}
} // namespace detail
/// @brief Straight-line circuit: XOR, AND, XNOR, NOT, and XOR with a public bit.
/// @details XOR, XNOR, and NOT are free. Each AND is one garbled row (32 bytes).
/// Wires are SSA. There is no secret shift, no memory, and no loop.
class netlist
{
public:
enum class op : unsigned char
{
xor_ = 0,
and_ = 1,
not_ = 2
};
struct gate
{
op code = op::xor_;
std::uint32_t a = 0;
std::uint32_t b = 0;
};
HEDLEY_WARN_UNUSED_RESULT
bit shared_in()
{
return push_in(input::shared);
}
/// @param owner 0 or 1
HEDLEY_WARN_UNUSED_RESULT
bit priv_in(unsigned owner)
{
if (owner > 1u)
throw std::invalid_argument("yao: input owner");
return push_in(owner == 0 ? input::priv0 : input::priv1);
}
HEDLEY_WARN_UNUSED_RESULT
bit xor_(bit a, bit b)
{
check(a);
check(b);
gates_.push_back(gate{op::xor_, a.id, b.id});
return bit{nwire() - 1u};
}
HEDLEY_WARN_UNUSED_RESULT
bit and_(bit a, bit b)
{
check(a);
check(b);
gates_.push_back(gate{op::and_, a.id, b.id});
++n_and_;
return bit{nwire() - 1u};
}
HEDLEY_WARN_UNUSED_RESULT
bit xnor_(bit a, bit b)
{
return not_(xor_(a, b));
}
HEDLEY_WARN_UNUSED_RESULT
bit not_(bit a)
{
check(a);
gates_.push_back(gate{op::not_, a.id, 0});
return bit{nwire() - 1u};
}
/// @brief XOR `x` with a public bit. A zero bit returns `x`.
HEDLEY_WARN_UNUSED_RESULT
bit xor_public(bit x, std::uint8_t public_bit)
{
check(x);
if ((public_bit & 1u) == 0)
return x;
return not_(x);
}
void out(bit x)
{
check(x);
outs_.push_back(x.id);
}
std::uint32_t n_in() const noexcept
{
return static_cast<std::uint32_t>(kinds_.size());
}
std::uint32_t n_out() const noexcept
{
return static_cast<std::uint32_t>(outs_.size());
}
std::uint32_t n_and() const noexcept { return n_and_; }
std::uint32_t nwire() const noexcept
{
return static_cast<std::uint32_t>(kinds_.size() + gates_.size());
}
/// @brief Bytes of garbled AND rows (`32 · n_and`).
std::size_t table_bytes() const noexcept
{
return static_cast<std::size_t>(n_and_) * 32u;
}
input kind_at(std::uint32_t i) const
{
return kinds_.at(i);
}
const std::vector<gate> & gates() const noexcept { return gates_; }
const std::vector<std::uint32_t> & outputs() const noexcept { return outs_; }
private:
bit push_in(input k)
{
if (!gates_.empty() || !outs_.empty())
throw std::logic_error("yao: inputs are declared before gates");
kinds_.push_back(k);
return bit{n_in() - 1u};
}
void check(bit x) const
{
if (x.id >= nwire())
throw std::invalid_argument("yao: wire");
}
std::vector<input> kinds_;
std::vector<gate> gates_;
std::vector<std::uint32_t> outs_;
std::uint32_t n_and_ = 0;
};
namespace detail
{
struct garbled
{
block delta{};
std::vector<block> z;
std::vector<block> tables;
std::vector<block> direct;
std::vector<block> ot0;
std::vector<block> ot1;
std::vector<std::uint8_t> share;
};
inline void garble_gates(const netlist & nl, block delta, std::vector<block> & z,
std::vector<block> & tables)
{
tables.clear();
tables.reserve(static_cast<std::size_t>(nl.n_and()) * 2u);
std::uint64_t gid = 0;
const auto & gates = nl.gates();
const std::uint32_t base = nl.n_in();
for (std::uint32_t gi = 0; gi < gates.size(); ++gi)
{
const auto & g = gates[gi];
const std::uint32_t dst = base + gi;
switch (g.code)
{
case netlist::op::xor_:
z[dst] = xor_b(z[g.a], z[g.b]);
break;
case netlist::op::not_:
z[dst] = xor_b(z[g.a], delta);
break;
case netlist::op::and_:
{
block row[2];
z[dst] = and_garble(z[g.a], z[g.b], delta, gid, row);
tables.push_back(row[0]);
tables.push_back(row[1]);
++gid;
break;
}
}
}
}
inline std::vector<std::uint8_t> eval_active(const netlist & nl,
const block * tables, std::vector<block> & w)
{
std::uint64_t gid = 0;
std::size_t ti = 0;
const auto & gates = nl.gates();
const std::uint32_t base = nl.n_in();
const std::size_t ntable = static_cast<std::size_t>(nl.n_and()) * 2u;
for (std::uint32_t gi = 0; gi < gates.size(); ++gi)
{
const auto & g = gates[gi];
const std::uint32_t dst = base + gi;
switch (g.code)
{
case netlist::op::xor_:
w[dst] = xor_b(w[g.a], w[g.b]);
break;
case netlist::op::not_:
w[dst] = w[g.a];
break;
case netlist::op::and_:
if (ti + 2u > ntable)
throw std::runtime_error("yao: truncated garbled tables");
w[dst] = and_eval(w[g.a], w[g.b], gid, tables + ti);
ti += 2u;
++gid;
break;
}
}
std::vector<std::uint8_t> share(nl.n_out());
const auto & outs = nl.outputs();
for (std::uint32_t i = 0; i < outs.size(); ++i)
share[i] = lsb(w[outs[i]]);
return share;
}
inline void check_party_bits(const netlist & nl, int me, const std::uint8_t * in)
{
if (me != 0 && me != 1)
throw std::invalid_argument("yao: party");
if (nl.n_in() != 0 && in == nullptr)
throw std::invalid_argument("yao: missing inputs");
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
const bool used = k == input::shared
|| (k == input::priv0 && me == 0)
|| (k == input::priv1 && me == 1);
if (used)
require_bit(in[i], "yao: input bit");
}
}
inline garbled garble_party0(const netlist & nl, const std::uint8_t * p0)
{
check_party_bits(nl, 0, p0);
garbled g;
g.delta = random_delta();
g.z.assign(nl.nwire(), zero_b());
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const block lg = random_label();
switch (nl.kind_at(i))
{
case input::shared:
{
const block le = random_label();
g.z[i] = xor_b(lg, le);
g.direct.push_back(xor_b(lg, mask_bit(p0[i], g.delta)));
g.ot0.push_back(le);
g.ot1.push_back(xor_b(le, g.delta));
break;
}
case input::priv0:
g.z[i] = lg;
g.direct.push_back(xor_b(lg, mask_bit(p0[i], g.delta)));
break;
case input::priv1:
g.z[i] = lg;
g.ot0.push_back(lg);
g.ot1.push_back(xor_b(lg, g.delta));
break;
}
}
garble_gates(nl, g.delta, g.z, g.tables);
g.share.resize(nl.n_out());
const auto & outs = nl.outputs();
for (std::uint32_t i = 0; i < outs.size(); ++i)
g.share[i] = lsb(g.z[outs[i]]);
return g;
}
inline std::vector<std::uint8_t> eval_party1(const netlist & nl,
const std::vector<block> & tables, const std::vector<block> & direct,
const std::vector<block> & chosen)
{
if (tables.size() != static_cast<std::size_t>(nl.n_and()) * 2u)
throw std::runtime_error("yao: table size");
std::vector<block> w(nl.nwire(), zero_b());
std::size_t di = 0;
std::size_t oi = 0;
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
switch (nl.kind_at(i))
{
case input::shared:
if (di >= direct.size() || oi >= chosen.size())
throw std::runtime_error("yao: input label count");
w[i] = xor_b(direct[di++], chosen[oi++]);
break;
case input::priv0:
if (di >= direct.size())
throw std::runtime_error("yao: input label count");
w[i] = direct[di++];
break;
case input::priv1:
if (oi >= chosen.size())
throw std::runtime_error("yao: input label count");
w[i] = chosen[oi++];
break;
}
}
if (di != direct.size() || oi != chosen.size())
throw std::runtime_error("yao: input label count");
return eval_active(nl, tables.data(), w);
}
} // namespace detail
/// @brief Cleartext walk. `semantic[i]` is the bit on input wire `i`.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> eval_plain(const netlist & nl,
const std::uint8_t * semantic)
{
if (nl.n_in() != 0 && semantic == nullptr)
throw std::invalid_argument("yao: missing inputs");
std::vector<std::uint8_t> w(nl.nwire(), 0);
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
detail::require_bit(semantic[i], "yao: input bit");
w[i] = semantic[i];
}
const auto & gates = nl.gates();
const std::uint32_t base = nl.n_in();
for (std::uint32_t gi = 0; gi < gates.size(); ++gi)
{
const auto & g = gates[gi];
const std::uint32_t dst = base + gi;
switch (g.code)
{
case netlist::op::xor_:
w[dst] = static_cast<std::uint8_t>(w[g.a] ^ w[g.b]);
break;
case netlist::op::and_:
w[dst] = static_cast<std::uint8_t>(w[g.a] & w[g.b]);
break;
case netlist::op::not_:
w[dst] = static_cast<std::uint8_t>(w[g.a] ^ 1u);
break;
}
}
std::vector<std::uint8_t> out(nl.n_out());
const auto & outs = nl.outputs();
for (std::uint32_t i = 0; i < outs.size(); ++i)
out[i] = w[outs[i]];
return out;
}
/// @brief Garble and evaluate in one process. Returns both XOR shares.
/// @details `p0` and `p1` use the same layout as `session::eval`. The
/// reconstructed bit is `first[i] XOR second[i]`.
HEDLEY_WARN_UNUSED_RESULT
inline std::pair<std::vector<std::uint8_t>, std::vector<std::uint8_t>>
eval_pair(const netlist & nl, const std::uint8_t * p0, const std::uint8_t * p1)
{
detail::check_party_bits(nl, 0, p0);
detail::check_party_bits(nl, 1, p1);
auto g = detail::garble_party0(nl, p0);
std::vector<block> chosen;
chosen.reserve(g.ot0.size());
std::size_t oi = 0;
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
if (k != input::shared && k != input::priv1)
continue;
if (oi >= g.ot0.size())
throw std::logic_error("yao: ot count");
chosen.push_back((p1[i] & 1u) ? g.ot1[oi] : g.ot0[oi]);
++oi;
}
auto share1 = detail::eval_party1(nl, g.tables, g.direct, chosen);
return {std::move(g.share), std::move(share1)};
}
/// @brief One-process garble of a known input. Shared wires take the semantic
/// bit as party 0's share and zero as party 1's share.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> eval_local(const netlist & nl,
const std::uint8_t * semantic)
{
std::vector<std::uint8_t> p0(nl.n_in(), 0);
std::vector<std::uint8_t> p1(nl.n_in(), 0);
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
detail::require_bit(semantic[i], "yao: input bit");
if (nl.kind_at(i) == input::priv1)
p1[i] = semantic[i];
else
p0[i] = semantic[i];
}
auto [a, b] = eval_pair(nl, p0.data(), p1.data());
for (std::size_t i = 0; i < a.size(); ++i)
a[i] = static_cast<std::uint8_t>(a[i] ^ b[i]);
return a;
}
/// @brief One party's garbling session. Party 0 always garbles.
/// @details The first `eval` fixes `me`. A second call with the other party
/// id throws. Base OT state stays inside the session.
class session
{
public:
session() = default;
session(const session &) = delete;
session & operator=(const session &) = delete;
session(session &&) = default;
session & operator=(session &&) = default;
/// @brief XOR shares of each output wire.
/// @param me 0 (garbler) or 1 (evaluator)
/// @param nl circuit. Both parties pass equal netlists.
/// @param in_bits length `nl.n_in()`. Shared entries are this party's
/// XOR share. `priv0` is read on party 0. `priv1` is read on party 1.
/// @param link peer channel. No other traffic may run on it during `eval`.
HEDLEY_WARN_UNUSED_RESULT
std::vector<std::uint8_t> eval(int me, const netlist & nl,
const std::uint8_t * in_bits, net::channel & link)
{
if (!bound_)
{
if (me != 0 && me != 1)
throw std::invalid_argument("yao: party");
bound_ = true;
me_ = me;
}
else if (me != me_)
throw std::logic_error("yao: session party is fixed");
detail::check_party_bits(nl, me, in_bits);
if (me == 0)
{
auto g = detail::garble_party0(nl, in_bits);
std::vector<block> ignored;
iknp::transfer_labels(link, 0, true, ot_, g.ot0, g.ot1, {}, ignored);
std::vector<block> blob;
blob.reserve(g.tables.size() + g.direct.size());
blob.insert(blob.end(), g.tables.begin(), g.tables.end());
blob.insert(blob.end(), g.direct.begin(), g.direct.end());
link.send_vec(blob, net::msg::bytes);
return g.share;
}
std::vector<std::uint8_t> choices;
choices.reserve(nl.n_in());
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
if (k == input::shared || k == input::priv1)
choices.push_back(in_bits[i]);
}
std::vector<block> chosen;
iknp::transfer_labels(link, 1, false, ot_, {}, {}, choices, chosen);
const auto blob = link.recv_vec<block>(net::msg::bytes);
const std::size_t ntable = static_cast<std::size_t>(nl.n_and()) * 2u;
std::size_t ndirect = 0;
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
if (k == input::shared || k == input::priv0)
++ndirect;
}
if (blob.size() != ntable + ndirect)
throw std::runtime_error("yao: garbled payload size");
std::vector<block> tables(ntable);
std::vector<block> direct(ndirect);
if (ntable != 0)
std::memcpy(tables.data(), blob.data(), ntable * sizeof(block));
if (ndirect != 0)
std::memcpy(direct.data(), blob.data() + ntable, ndirect * sizeof(block));
return detail::eval_party1(nl, tables, direct, chosen);
}
private:
bool bound_ = false;
int me_ = 0;
iknp::detail::role_state ot_{};
};
} // namespace yao
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_YAO_HPP__