libdpf/include/dpf/circuit.hpp

553 lines
18 KiB
C++
Raw Normal View History

/// @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