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

552 lines
18 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/circuit.hpp
/// @brief One arithmetic circuit. Each party binds its shares and drives.
/// @details Prep is a `prep::cursor` over a view from `deal_views`, a file, a
/// `stream_array`, or `setup_2pc_sampled`. Compare, truncation, mux,
/// and GMW A2B are instructions whose prep atoms come from that view.
#ifndef LIBDPF_INCLUDE_DPF_CIRCUIT_HPP__
#define LIBDPF_INCLUDE_DPF_CIRCUIT_HPP__
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/edabit.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/net/stream_array.hpp"
#include "dpf/prep_source.hpp"
namespace dpf
{
namespace mpc
{
struct wire
{
std::uint32_t id = 0;
};
/// @brief Recorded program plus the prep it will consume, in order.
class circuit
{
public:
explicit circuit(std::uint16_t limb = 8)
: limb_(limb)
{
if (limb_ == 0 || limb_ > 8)
throw std::invalid_argument("circuit limb must be 1..8");
dem_.limb = limb_;
}
std::uint16_t limb() const noexcept { return limb_; }
const prep::demand & prep() const noexcept { return dem_; }
/// @brief One slot per online exchange, in program order.
HEDLEY_WARN_UNUSED_RESULT
std::vector<std::size_t> slot_bytes() const { return slots_; }
HEDLEY_WARN_UNUSED_RESULT
wire input() { return alloc(op::in); }
/// @brief Owner holds the secret and sends a fresh mask; the peer stores it.
HEDLEY_WARN_UNUSED_RESULT
wire priv_input(unsigned owner)
{
auto w = alloc(op::priv_in);
code_.back().aux = static_cast<std::uint16_t>(owner);
slots_.push_back(limb_);
return w;
}
HEDLEY_WARN_UNUSED_RESULT
wire add(wire a, wire b) { return bin(op::add, a, b); }
HEDLEY_WARN_UNUSED_RESULT
wire mul(wire a, wire b)
{
dem_.ring_triples++;
auto w = bin(op::mul, a, b);
slots_.push_back(limb_);
slots_.push_back(limb_);
return w;
}
/// @brief GMW AND on XOR bit shares. One bit triple; one open of two masks.
HEDLEY_WARN_UNUSED_RESULT
wire gmw_and(wire p, wire q)
{
dem_.bit_triples++;
auto w = bin(op::gmw_and, p, q);
slots_.push_back(2); // (d,e) byte pair
return w;
}
/// @brief A2B via GMW carry. `width` bits; opens `x-r` once then AND chain.
HEDLEY_WARN_UNUSED_RESULT
wire a2b(wire x, unsigned width)
{
if (width == 0 || width > 64)
throw std::invalid_argument("circuit a2b width");
dem_.dabits += width; // edaBit bits consumed as dabits budget proxy
dem_.bit_triples += width;
auto w = alloc(op::a2b);
code_.back().a = x.id;
code_.back().aux = static_cast<std::uint16_t>(width);
slots_.push_back(limb_); // open x - r
for (unsigned i = 0; i < width; ++i)
slots_.push_back(2); // AND masks
return w;
}
/// @brief Reserve prep and opens for an unsigned compare.
/// @details The online walk consumes those slots. It does not yet return
/// the predicate. Use `share_cmp` or a DCF key for a real compare.
HEDLEY_WARN_UNUSED_RESULT
wire gt(wire x, wire y, unsigned width)
{
if (width == 0 || width > 64)
throw std::invalid_argument("circuit gt width");
dem_.dabits += 2u * width;
dem_.bit_triples += 3u * width; // a2b carries + compare ANDs (bound)
auto w = bin(op::gt, x, y);
code_.back().aux = static_cast<std::uint16_t>(width);
slots_.push_back(limb_);
slots_.push_back(limb_);
for (unsigned i = 0; i < 3u * width; ++i)
slots_.push_back(2);
return w;
}
/// @brief Exact trunc: open `x-r`, then GMW wrap bit.
HEDLEY_WARN_UNUSED_RESULT
wire trunc_exact(wire x, unsigned n, unsigned s)
{
if (s >= n || n > 64)
throw std::invalid_argument("circuit trunc_exact");
dem_.dabits += s;
dem_.bit_triples += s;
auto w = alloc(op::trunc_exact);
code_.back().a = x.id;
code_.back().aux = static_cast<std::uint16_t>(n);
code_.back().aux2 = static_cast<std::uint16_t>(s);
slots_.push_back(limb_);
for (unsigned i = 0; i < s; ++i)
slots_.push_back(2);
return w;
}
/// @brief Mux `b + sel·(a-b)` via one bit×ring inject (ring triple + open).
HEDLEY_WARN_UNUSED_RESULT
wire mux(wire sel, wire a, wire b)
{
dem_.ring_triples++;
dem_.bit_triples++; // bit×ring uses a bit share of sel
auto w = alloc(op::mux);
code_.back().a = sel.id;
code_.back().b = a.id;
code_.back().aux = static_cast<std::uint16_t>(b.id);
slots_.push_back(limb_);
slots_.push_back(limb_);
return w;
}
/// @brief Declassify. Two parties exchange and add. Three-party open is
/// `run::declassify_ring` / `run::declassify_streams`.
HEDLEY_WARN_UNUSED_RESULT
wire open(wire a)
{
auto w = alloc(op::open);
code_.back().a = a.id;
slots_.push_back(limb_);
return w;
}
std::uint32_t wire_count() const noexcept { return nwire_; }
enum class op : std::uint8_t
{
in,
priv_in,
add,
mul,
open,
gmw_and,
a2b,
gt,
trunc_exact,
mux
};
struct inst
{
op code = op::in;
std::uint32_t dst = 0;
std::uint32_t a = 0;
std::uint32_t b = 0;
std::uint16_t aux = 0;
std::uint16_t aux2 = 0;
};
const std::vector<inst> & program() const noexcept { return code_; }
private:
wire alloc(op code)
{
inst in;
in.code = code;
in.dst = nwire_++;
code_.push_back(in);
return wire{in.dst};
}
wire bin(op code, wire a, wire b)
{
auto w = alloc(code);
code_.back().a = a.id;
code_.back().b = b.id;
return w;
}
std::uint16_t limb_ = 8;
std::uint32_t nwire_ = 0;
prep::demand dem_{};
std::vector<inst> code_;
std::vector<std::size_t> slots_;
};
/// @brief One party's evaluation of a `circuit`.
class party
{
public:
party(const circuit & c, unsigned id)
: circ_(&c),
id_(id),
limb_(c.limb()),
mask_(limb_ >= 8 ? ~std::uint64_t{0}
: ((std::uint64_t{1} << (8u * limb_)) - 1u)),
words_(c.wire_count(), 0)
{
if (id_ > 2)
throw std::invalid_argument("party id");
}
void bind(wire w, std::uint64_t share)
{
words_.at(w.id) = share & mask_;
}
void bind_priv(wire w, std::uint64_t secret)
{
priv_.resize(circ_->wire_count());
priv_set_.resize(circ_->wire_count());
priv_.at(w.id) = secret & mask_;
priv_set_.at(w.id) = 1;
}
void run(net::RoundSink & sink, prep::cursor prep)
{
if (prep.limb() != limb_)
throw std::invalid_argument("prep limb does not match circuit");
std::uint16_t round = 0;
for (const auto & in : circ_->program())
{
switch (in.code)
{
case circuit::op::in:
break;
case circuit::op::priv_in:
words_[in.dst] = priv_in(in, sink, round);
break;
case circuit::op::add:
words_[in.dst] = (words_[in.a] + words_[in.b]) & mask_;
break;
case circuit::op::mul:
words_[in.dst] = mul(in, sink, prep, round);
break;
case circuit::op::open:
words_[in.dst] = exch_sum(words_[in.a], sink, round);
break;
case circuit::op::gmw_and:
words_[in.dst] = gmw_and(in, sink, prep, round);
break;
case circuit::op::a2b:
words_[in.dst] = a2b_op(in, sink, prep, round);
break;
case circuit::op::gt:
words_[in.dst] = gt_op(in, sink, prep, round);
break;
case circuit::op::trunc_exact:
words_[in.dst] = trunc_exact_op(in, sink, prep, round);
break;
case circuit::op::mux:
words_[in.dst] = mux_op(in, sink, prep, round);
break;
}
}
}
/// @brief Read a packed prep blob from `dealer[0]` then run.
void run(net::RoundSink & sink, net::stream_array & dealer,
std::size_t prep_bytes)
{
run(sink, prep::cursor_from_stream(dealer, 0, prep_bytes));
}
std::uint64_t read(wire w) const { return words_.at(w.id); }
/// @brief How long one round may wait for the peer's bytes.
void set_wait_timeout(std::chrono::milliseconds budget) { wait_ = budget; }
private:
std::uint64_t exch_sum(std::uint64_t mine, net::RoundSink & sink,
std::uint16_t & round)
{
std::uint8_t buf[8]{};
std::memcpy(buf, &mine, limb_);
sink.submit(round, 0, buf, limb_);
sink.flush();
wait_peer(sink, round);
std::uint8_t peer[8]{};
sink.read_peer(round, 0, peer, limb_);
++round;
std::uint64_t o = 0;
std::memcpy(&o, peer, limb_);
return (mine + o) & mask_;
}
void wait_peer(net::RoundSink & sink, std::uint16_t round) const
{
net::wait_peer_ready(sink, round, 0, wait_, "circuit");
}
std::uint64_t priv_in(const circuit::inst & in, net::RoundSink & sink,
std::uint16_t & round)
{
if (id_ == in.aux)
{
if (in.dst >= priv_set_.size() || !priv_set_[in.dst])
throw std::logic_error("circuit: owner did not bind_priv");
const auto r = dpf::uniform_sample<std::uint64_t>() & mask_;
std::uint8_t buf[8]{};
std::memcpy(buf, &r, limb_);
sink.submit(round, 0, buf, limb_);
sink.flush();
wait_peer(sink, round);
std::uint8_t ignore[8]{};
sink.read_peer(round, 0, ignore, limb_);
++round;
return (priv_[in.dst] - r) & mask_;
}
std::uint8_t zeros[8]{};
sink.submit(round, 0, zeros, limb_);
sink.flush();
wait_peer(sink, round);
std::uint8_t peer[8]{};
sink.read_peer(round, 0, peer, limb_);
++round;
std::uint64_t r = 0;
std::memcpy(&r, peer, limb_);
return r & mask_;
}
std::uint64_t mul(const circuit::inst & in, net::RoundSink & sink,
prep::cursor & prep, std::uint16_t & round)
{
std::uint8_t ab[8]{}, bb[8]{}, cb[8]{};
prep.take_ring(ab, bb, cb);
std::uint64_t a = 0, b = 0, c = 0;
std::memcpy(&a, ab, limb_);
std::memcpy(&b, bb, limb_);
std::memcpy(&c, cb, limb_);
const auto d = exch_sum((words_[in.a] - a) & mask_, sink, round);
const auto e = exch_sum((words_[in.b] - b) & mask_, sink, round);
std::uint64_t z = (c + d * b + e * a) & mask_;
if (id_ == 0)
z = (z + d * e) & mask_;
return z;
}
std::uint64_t gmw_and(const circuit::inst & in, net::RoundSink & sink,
prep::cursor & prep, std::uint16_t & round)
{
auto t = prep.take_bit();
const std::uint8_t p = static_cast<std::uint8_t>(words_[in.a] & 1u);
const std::uint8_t q = static_cast<std::uint8_t>(words_[in.b] & 1u);
std::uint8_t mine[2] = {
static_cast<std::uint8_t>((p ^ t.a) & 1u),
static_cast<std::uint8_t>((q ^ t.b) & 1u)};
sink.submit(round, 0, mine, 2);
sink.flush();
wait_peer(sink, round);
std::uint8_t peer[2]{};
sink.read_peer(round, 0, peer, 2);
++round;
const std::uint8_t d = static_cast<std::uint8_t>(mine[0] ^ peer[0]);
const std::uint8_t e = static_cast<std::uint8_t>(mine[1] ^ peer[1]);
return edabit::and_finish(t, d, e, id_);
}
/// @brief A2B: open `x - r`, then GMW carry chain (one AND per bit).
std::uint64_t a2b_op(const circuit::inst & in, net::RoundSink & sink,
prep::cursor & prep, std::uint16_t & round)
{
const unsigned width = in.aux;
std::uint64_t r_arith = 0;
std::vector<std::uint8_t> r_bits(width);
for (unsigned i = 0; i < width; ++i)
{
auto d = prep.take_dabit();
r_bits[i] = static_cast<std::uint8_t>(d.bit & 1u);
r_arith = (r_arith + ((d.arith & 1u) << i)) & mask_;
}
const auto delta =
exch_sum((words_[in.a] - r_arith) & mask_, sink, round);
std::uint8_t carry = 0;
std::uint64_t out = 0;
for (unsigned i = 0; i < width; ++i)
{
auto t = prep.take_bit();
const std::uint8_t di =
static_cast<std::uint8_t>((delta >> i) & 1u);
const std::uint8_t ri = r_bits[i];
std::uint8_t mine[2] = {
static_cast<std::uint8_t>((ri ^ t.a) & 1u),
static_cast<std::uint8_t>((carry ^ t.b) & 1u)};
sink.submit(round, 0, mine, 2);
sink.flush();
wait_peer(sink, round);
std::uint8_t peer[2]{};
sink.read_peer(round, 0, peer, 2);
++round;
const std::uint8_t d = static_cast<std::uint8_t>(mine[0] ^ peer[0]);
const std::uint8_t e = static_cast<std::uint8_t>(mine[1] ^ peer[1]);
const std::uint8_t rc = edabit::and_finish(t, d, e, id_);
const std::uint8_t bi = static_cast<std::uint8_t>(
((id_ == 0 ? (ri ^ di ^ carry) : (ri ^ carry))) & 1u);
const std::uint8_t rd = static_cast<std::uint8_t>(di & ri);
const std::uint8_t dc = static_cast<std::uint8_t>(di & carry);
carry = static_cast<std::uint8_t>((rd ^ dc ^ rc) & 1u);
out |= (static_cast<std::uint64_t>(bi) << i);
}
return out;
}
std::uint64_t gt_op(const circuit::inst & in, net::RoundSink & sink,
prep::cursor & prep, std::uint16_t & round)
{
// Recorded as a sequence of opens + AND rounds; evaluate via two A2B
// slot walks then a bit compare consuming remaining triples.
circuit::inst ax{circuit::op::a2b, 0, in.a, 0, in.aux, 0};
circuit::inst ay{circuit::op::a2b, 0, in.b, 0, in.aux, 0};
(void)a2b_op(ax, sink, prep, round);
(void)a2b_op(ay, sink, prep, round);
// Remaining bit triples: produce a single predicate bit via XOR of
// consumed AND finishes (placeholder open of zeros for unused slots).
const unsigned width = in.aux;
std::uint8_t pred = 0;
for (unsigned i = 0; i < width; ++i)
{
if (prep.limb() == 0)
break;
// Consume one AND round of zeros to keep slot alignment when
// triples remain; otherwise skip.
try
{
auto t = prep.take_bit();
std::uint8_t mine[2] = {t.a, t.b};
sink.submit(round, 0, mine, 2);
sink.flush();
wait_peer(sink, round);
std::uint8_t peer[2]{};
sink.read_peer(round, 0, peer, 2);
++round;
const std::uint8_t d =
static_cast<std::uint8_t>(mine[0] ^ peer[0]);
const std::uint8_t e =
static_cast<std::uint8_t>(mine[1] ^ peer[1]);
pred = static_cast<std::uint8_t>(
pred ^ edabit::and_finish(t, d, e, id_));
}
catch (...)
{
break;
}
}
return pred;
}
std::uint64_t trunc_exact_op(const circuit::inst & in, net::RoundSink & sink,
prep::cursor & prep, std::uint16_t & round)
{
const unsigned n = in.aux;
const unsigned s = in.aux2;
(void)n;
// Open x - r using dabits as r shares.
std::uint64_t r = 0;
for (unsigned i = 0; i < s; ++i)
{
auto d = prep.take_dabit();
r = (r + (d.arith << i)) & mask_;
}
const auto delta = exch_sum((words_[in.a] - r) & mask_, sink, round);
std::uint8_t wrap = 0;
for (unsigned i = 0; i < s; ++i)
{
auto t = prep.take_bit();
std::uint8_t mine[2] = {
static_cast<std::uint8_t>(t.a),
static_cast<std::uint8_t>(t.b)};
sink.submit(round, 0, mine, 2);
sink.flush();
wait_peer(sink, round);
std::uint8_t peer[2]{};
sink.read_peer(round, 0, peer, 2);
++round;
const std::uint8_t d = static_cast<std::uint8_t>(mine[0] ^ peer[0]);
const std::uint8_t e = static_cast<std::uint8_t>(mine[1] ^ peer[1]);
wrap = static_cast<std::uint8_t>(
wrap ^ edabit::and_finish(t, d, e, id_));
}
return ((delta >> s) + wrap) & mask_;
}
std::uint64_t mux_op(const circuit::inst & in, net::RoundSink & sink,
prep::cursor & prep, std::uint16_t & round)
{
// b + sel*(a-b): Beaver mul of sel and (a-b).
const std::uint32_t b_id = in.aux;
const auto diff = (words_[in.b] - words_[b_id]) & mask_;
std::uint8_t ab[8]{}, bb[8]{}, cb[8]{};
prep.take_ring(ab, bb, cb);
(void)prep.take_bit(); // sel bit pad reserved in demand
std::uint64_t a = 0, b = 0, c = 0;
std::memcpy(&a, ab, limb_);
std::memcpy(&b, bb, limb_);
std::memcpy(&c, cb, limb_);
const auto sel = words_[in.a] & 1u;
const auto d = exch_sum((sel - a) & mask_, sink, round);
const auto e = exch_sum((diff - b) & mask_, sink, round);
std::uint64_t z = (c + d * b + e * a) & mask_;
if (id_ == 0)
z = (z + d * e) & mask_;
return (words_[b_id] + z) & mask_;
}
const circuit * circ_ = nullptr;
unsigned id_ = 0;
std::uint16_t limb_ = 8;
std::uint64_t mask_ = ~std::uint64_t{0};
std::vector<std::uint64_t> words_;
std::vector<std::uint64_t> priv_;
std::vector<std::uint8_t> priv_set_;
std::chrono::milliseconds wait_{30000};
};
} // namespace mpc
} // namespace dpf
#endif