/// @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 #include #include #include #include #include #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 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(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(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(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(n); code_.back().aux2 = static_cast(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(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 & 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 code_; std::vector 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() & 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(words_[in.a] & 1u); const std::uint8_t q = static_cast(words_[in.b] & 1u); std::uint8_t mine[2] = { static_cast((p ^ t.a) & 1u), static_cast((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(mine[0] ^ peer[0]); const std::uint8_t e = static_cast(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 r_bits(width); for (unsigned i = 0; i < width; ++i) { auto d = prep.take_dabit(); r_bits[i] = static_cast(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((delta >> i) & 1u); const std::uint8_t ri = r_bits[i]; std::uint8_t mine[2] = { static_cast((ri ^ t.a) & 1u), static_cast((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(mine[0] ^ peer[0]); const std::uint8_t e = static_cast(mine[1] ^ peer[1]); const std::uint8_t rc = edabit::and_finish(t, d, e, id_); const std::uint8_t bi = static_cast( ((id_ == 0 ? (ri ^ di ^ carry) : (ri ^ carry))) & 1u); const std::uint8_t rd = static_cast(di & ri); const std::uint8_t dc = static_cast(di & carry); carry = static_cast((rd ^ dc ^ rc) & 1u); out |= (static_cast(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(mine[0] ^ peer[0]); const std::uint8_t e = static_cast(mine[1] ^ peer[1]); pred = static_cast( 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(t.a), static_cast(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(mine[0] ^ peer[0]); const std::uint8_t e = static_cast(mine[1] ^ peer[1]); wrap = static_cast( 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 words_; std::vector priv_; std::vector priv_set_; std::chrono::milliseconds wait_{30000}; }; } // namespace mpc } // namespace dpf #endif