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>
552 lines
18 KiB
C++
552 lines
18 KiB
C++
/// @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
|