/// @file dpf/factory_gadgets.hpp /// @brief Online MPC gadgets as `net::stream_array` round sequences. /// @details These mirror the limb logic in `mpc::circuit` but speak only /// `factory::detail::read_pod` / `exchange_pod` on dealer and peer /// arrays. Higher-level circuit ops compose from them: /// - `gt` — two `a2b_online` passes plus compare ANDs /// - `trunc_exact` — open `x-r` then a carry AND chain (as in A2B) /// - `mux` — Beaver mul on `(sel, a-b)` plus local add of `b` #ifndef LIBDPF_INCLUDE_DPF_FACTORY_GADGETS_HPP__ #define LIBDPF_INCLUDE_DPF_FACTORY_GADGETS_HPP__ #include #include #include #include #include #include #include "hedley/hedley.h" #include "dpf/circuit.hpp" #include "dpf/edabit.hpp" #include "dpf/gilboa.hpp" #include "dpf/net/round_sink.hpp" #include "dpf/net/stream_array.hpp" #include "dpf/net/stream_mesh.hpp" #include "dpf/ot_pack.hpp" #include "dpf/protocol_factory.hpp" #include "dpf/bit_inject.hpp" namespace dpf { namespace factory { namespace detail { inline std::uint64_t limb_mask(std::uint16_t limb) { if (limb == 0 || limb > 8) throw std::invalid_argument("factory_gadgets limb"); return limb >= 8 ? ~std::uint64_t{0} : ((std::uint64_t{1} << (8u * limb)) - 1u); } inline void ring_view_to_limbs(const ring_triple_view & v, std::uint16_t limb, std::uint64_t & a, std::uint64_t & b, std::uint64_t & c) { a = b = c = 0; const std::uint16_t use = v.limb != 0 ? v.limb : limb; std::memcpy(&a, v.a, use); std::memcpy(&b, v.b, use); std::memcpy(&c, v.c, use); const auto mask = limb_mask(limb); a &= mask; b &= mask; c &= mask; } } // namespace detail /// @brief Sum-open one ring share on `peer[stream]` (additive declassify). inline std::uint64_t open_sum_online(unsigned /*party*/, std::uint64_t mine, net::stream_array & peer, std::size_t stream, std::uint16_t limb) { const auto mask = detail::limb_mask(limb); std::uint64_t theirs = 0; detail::exchange_pod(peer, stream, mine & mask, theirs); return (mine + theirs) & mask; } /// @brief Beaver multiplication: triple on `dealer[dealer_stream]`, opens on /// `peer[peer_d_stream]` and `peer[peer_e_stream]`. inline std::uint64_t beaver_mul_online(unsigned party, std::uint64_t x, std::uint64_t y, net::stream_array & peer, net::stream_array & dealer, std::size_t dealer_stream = 0, std::size_t peer_d_stream = 0, std::size_t peer_e_stream = 1, std::uint16_t limb = 8) { const auto mask = detail::limb_mask(limb); const auto view = detail::read_pod(dealer, dealer_stream); std::uint64_t a = 0, b = 0, c = 0; detail::ring_view_to_limbs(view, limb, a, b, c); const auto d = open_sum_online(party, (x - a) & mask, peer, peer_d_stream, limb); const auto e = open_sum_online(party, (y - b) & mask, peer, peer_e_stream, limb); std::uint64_t z = (c + d * b + e * a) & mask; if (party == 0) z = (z + d * e) & mask; return z; } /// @brief GMW AND on XOR bit shares; blind from `dealer[dealer_stream]`. inline std::uint8_t gmw_and_online(unsigned party, std::uint8_t p, std::uint8_t q, net::stream_array & peer, net::stream_array & dealer, std::size_t dealer_stream = 0, std::size_t peer_stream = 0) { const auto blind = detail::read_pod(dealer, dealer_stream); gmw_and_msg mine{}; mine.d = static_cast((p ^ blind.a) & 1u); mine.e = static_cast((q ^ blind.b) & 1u); gmw_and_msg theirs{}; detail::exchange_pod(peer, peer_stream, mine, theirs); const std::uint8_t d_open = static_cast(mine.d ^ theirs.d); const std::uint8_t e_open = static_cast(mine.e ^ theirs.e); ot::bit_triple t{blind.a, blind.b, blind.c}; return edabit::and_finish(t, d_open, e_open, party); } /// @brief Mux: `b + sel·(a-b)` via one Beaver mul of `(sel, a-b)`. inline std::uint64_t mux_online(unsigned party, std::uint8_t sel, std::uint64_t a, std::uint64_t b, net::stream_array & peer, net::stream_array & dealer, std::size_t dealer_stream = 0, std::size_t peer_d_stream = 0, std::size_t peer_e_stream = 1, std::uint16_t limb = 8) { const auto mask = detail::limb_mask(limb); const auto diff = (a - b) & mask; const auto prod = beaver_mul_online(party, sel & 1u, diff, peer, dealer, dealer_stream, peer_d_stream, peer_e_stream, limb); return (b + prod) & mask; } /// @brief Exact trunc: open `x-r` (r from `s` dabits), then GMW wrap via /// low-bit AND chain on the same shape as `a2b_online`'s carry. inline std::uint64_t trunc_exact_online(unsigned party, std::uint64_t x_share, unsigned n, unsigned s, net::stream_array & peer, net::stream_array & dealer, std::uint16_t limb = 8, std::size_t dabit_stream = 0, std::size_t and_blind_stream = 1, std::size_t open_stream = 0, std::size_t and_peer_stream = 1) { if (s >= n || n > 64) throw std::invalid_argument("trunc_exact_online"); const auto mask = detail::limb_mask(limb); std::uint64_t r = 0; std::vector r_bits(s); for (unsigned i = 0; i < s; ++i) { const auto d = detail::read_pod(dealer, dabit_stream); r_bits[i] = static_cast(d.bit & 1u); std::uint64_t ar = 0; const std::uint16_t use = d.limb != 0 ? d.limb : limb; std::memcpy(&ar, d.arith, use); r = (r + (ar << i)) & mask; } const auto delta = open_sum_online(party, (x_share - r) & mask, peer, open_stream, limb); // Wrap = carry out of the low `s` bits of delta + r (GMW chain). std::uint8_t carry = 0; for (unsigned i = 0; i < s; ++i) { const std::uint8_t di = static_cast((delta >> i) & 1u); const std::uint8_t ri = r_bits[i]; const std::uint8_t rc = gmw_and_online(party, ri, carry, peer, dealer, and_blind_stream, and_peer_stream); 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); } const auto d_wrap = detail::read_pod(dealer, dabit_stream); const std::uint8_t mine_mask = static_cast(carry ^ (d_wrap.bit & 1u)); std::uint8_t peer_mask = 0; detail::exchange_pod(peer, open_stream, mine_mask, peer_mask); const std::uint8_t mask_open = static_cast(mine_mask ^ peer_mask); ot::dabit dr{}; dr.bit = d_wrap.bit; const std::uint16_t use_wrap = d_wrap.limb != 0 ? d_wrap.limb : limb; std::memcpy(&dr.arith, d_wrap.arith, use_wrap); const auto wrap = edabit::b2a_party_bit( carry, dr, mask_open, party, 0); const std::uint64_t high_m = ((n - s) >= 64) ? ~std::uint64_t{0} : ((std::uint64_t{1} << (n - s)) - 1u); std::uint64_t out = wrap & mask; if (party == 0) out = (out + ((delta >> s) & high_m)) & mask; return out; } /// @brief A2B: `width` dabits on `dealer[dabit_stream]` (sequential reads), /// `width` AND blinds on `dealer[and_blind_stream]`, open delta on /// `peer[open_stream]`, AND masks on `peer[and_peer_stream]` (sequential). inline std::uint64_t a2b_online(unsigned party, std::uint64_t x_share, unsigned width, net::stream_array & peer, net::stream_array & dealer, std::uint16_t limb = 8, std::size_t dabit_stream = 0, std::size_t and_blind_stream = 1, std::size_t open_stream = 0, std::size_t and_peer_stream = 1) { if (width == 0 || width > 64) throw std::invalid_argument("a2b_online width"); const auto mask = detail::limb_mask(limb); std::uint64_t r_arith = 0; std::vector r_bits(width); for (unsigned i = 0; i < width; ++i) { const auto d = detail::read_pod(dealer, dabit_stream); r_bits[i] = static_cast(d.bit & 1u); std::uint64_t ar = 0; const std::uint16_t use = d.limb != 0 ? d.limb : limb; std::memcpy(&ar, d.arith, use); r_arith = (r_arith + (ar << i)) & mask; } const auto delta = open_sum_online(party, (x_share - r_arith) & mask, peer, open_stream, limb); std::uint8_t carry = 0; std::uint64_t out = 0; for (unsigned i = 0; i < width; ++i) { const auto blind = detail::read_pod(dealer, and_blind_stream); ot::bit_triple t{blind.a, blind.b, blind.c}; const std::uint8_t di = static_cast((delta >> i) & 1u); const std::uint8_t ri = r_bits[i]; gmw_and_msg mine{}; mine.d = static_cast((ri ^ t.a) & 1u); mine.e = static_cast((carry ^ t.b) & 1u); gmw_and_msg peer_msg{}; detail::exchange_pod(peer, and_peer_stream, mine, peer_msg); const std::uint8_t d = static_cast(mine.d ^ peer_msg.d); const std::uint8_t e = static_cast(mine.e ^ peer_msg.e); const std::uint8_t rc = edabit::and_finish(t, d, e, party); const std::uint8_t bi = static_cast( ((party == 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; } /// @brief Unsigned GT via two A2B passes plus MSB-first compare ANDs. inline std::uint8_t gt_online(unsigned party, std::uint64_t x, std::uint64_t y, unsigned width, net::stream_array & peer, net::stream_array & dealer, std::uint16_t limb = 8) { // Dealer layout: streams 0..1 dabits+ANDs for x, 2..3 for y, 4 for compare. const auto bx = a2b_online(party, x, width, peer, dealer, limb, 0, 1, 0, 1); const auto by = a2b_online(party, y, width, peer, dealer, limb, 2, 3, 2, 3); std::uint8_t gt = 0; std::uint8_t eq = static_cast(party == 0 ? 1 : 0); for (unsigned k = width; k-- > 0; ) { const std::uint8_t xk = static_cast((bx >> k) & 1u); const std::uint8_t yk = static_cast((by >> k) & 1u); const std::uint8_t xny = gmw_and_online(party, xk, static_cast(yk ^ (party == 0 ? 1u : 0u)), peer, dealer, 4, 4); const std::uint8_t both = gmw_and_online(party, eq, xny, peer, dealer, 4, 4); gt = static_cast(gt ^ both); const std::uint8_t same = static_cast( (xk ^ yk ^ (party == 0 ? 1u : 0u)) & 1u); eq = gmw_and_online(party, eq, same, peer, dealer, 4, 4); } return gt; } // --------------------------------------------------------------------------- // RoundSink adapter: one peer stream index per circuit exchange round // --------------------------------------------------------------------------- /// @brief Drive `mpc::party::run` with one `stream_array` index per slot. class stream_array_circuit_sink : public net::RoundSink { public: stream_array_circuit_sink(net::stream_array & peer, std::vector slot_bytes) : peer_(&peer), slot_bytes_(std::move(slot_bytes)) { windows_.reserve(slot_bytes_.size()); for (std::size_t sb : slot_bytes_) windows_.emplace_back(1, sb); if (peer_->size() < slot_bytes_.size()) throw std::invalid_argument( "stream_array_circuit_sink: peer too few streams"); } std::size_t count() const noexcept override { return 1; } std::size_t rounds() const noexcept override { return slot_bytes_.size(); } std::size_t slot_bytes(std::uint16_t round) const override { if (round >= slot_bytes_.size()) throw std::out_of_range("stream_array_circuit_sink round"); return slot_bytes_[round]; } void submit(std::uint16_t round, std::size_t index, const std::uint8_t * bytes, std::size_t n) override { if (index != 0) throw std::out_of_range("stream_array_circuit_sink index"); window(round).submit(0, bytes, n); } bool peer_ready(std::uint16_t round, std::size_t index) const override { if (index != 0) return false; return window(round).peer_ready(0); } void read_peer(std::uint16_t round, std::size_t index, std::uint8_t * out, std::size_t n) const override { if (index != 0) throw std::out_of_range("stream_array_circuit_sink read"); window(round).read_peer(0, out, n); } void flush() override { for (std::uint16_t r = 0; r < slot_bytes_.size(); ++r) { std::size_t begin = 0; std::size_t n = 0; window(r).pending_out(begin, n); if (n != 0) flush_round(r); } } void flush_round(std::uint16_t round) override { if (round >= slot_bytes_.size()) throw std::out_of_range("stream_array_circuit_sink flush_round"); auto & w = window(round); std::size_t begin = 0; std::size_t nslots = 0; const auto * pend = w.pending_out(begin, nslots); if (nslots == 0) return; const std::size_t nbyte = nslots * slot_bytes_[round]; peer_->write(round, pend, nbyte); peer_->flush(round); w.mark_flushed(nslots); std::vector peer_buf(nbyte); peer_->read(round, peer_buf.data(), nbyte); w.accept_peer_at(0, peer_buf.data(), nbyte); } void poll() override {} private: net::round_window & window(std::uint16_t round) { return windows_.at(round); } const net::round_window & window(std::uint16_t round) const { return windows_.at(round); } net::stream_array * peer_ = nullptr; std::vector slot_bytes_; std::vector windows_; }; /// @brief Run a compiled circuit using stream indices for each online slot. inline void run_circuit_stream_array(const mpc::circuit & circ, mpc::party & me, net::stream_array & peer, net::stream_array & dealer, std::size_t prep_bytes) { stream_array_circuit_sink sink(peer, circ.slot_bytes()); me.run(sink, dealer, prep_bytes); } // --------------------------------------------------------------------------- // Gilboa OT pads on stream arrays // --------------------------------------------------------------------------- inline void write_ot_pack_wire(net::stream_array & dealer, std::size_t stream, const ot::pack::wire & w) { const std::uint64_t nb = static_cast(w.b2a.size()); const std::uint64_t nt = static_cast(w.bits.size()); const std::uint64_t nr = w.bit_ring_elem != 0 ? static_cast(w.bit_ring.size() / w.bit_ring_elem) : 0u; const std::uint64_t hdr[4] = {static_cast(w.me), nb, nt, nr}; dealer.write(stream, hdr, sizeof(hdr)); if (nb != 0) dealer.write(stream, w.b2a.data(), nb * sizeof(ot::b2a_slot)); if (nt != 0) dealer.write(stream, w.bits.data(), nt * sizeof(ot::bit_triple)); if (nr != 0) dealer.write(stream, w.bit_ring.data(), w.bit_ring.size()); dealer.flush(stream); } inline ot::pack read_ot_pack_wire(net::stream_array & dealer, std::size_t stream) { std::uint64_t hdr[4]{}; dealer.read(stream, hdr, sizeof(hdr)); ot::pack::wire w; w.me = static_cast(hdr[0]); const std::size_t nb = static_cast(hdr[1]); const std::size_t nt = static_cast(hdr[2]); const std::size_t nr = static_cast(hdr[3]); w.b2a.resize(nb); w.bits.resize(nt); if (nb != 0) dealer.read(stream, w.b2a.data(), nb * sizeof(ot::b2a_slot)); if (nt != 0) dealer.read(stream, w.bits.data(), nt * sizeof(ot::bit_triple)); if (nr != 0) { w.bit_ring_elem = sizeof(ot::bit_ring_triple); w.bit_ring.resize(nr * w.bit_ring_elem); dealer.read(stream, w.bit_ring.data(), w.bit_ring.size()); } return ot::pack::from_wire(std::move(w)); } /// @brief Sample correlated Gilboa/OT pads onto each party's dealer stream. inline void deal_gilboa_mul_tape(net::stream_array & dealer0, net::stream_array & dealer1, std::size_t bits) { auto pr = ot::sample_dealer_pair(bits, 0, bits); write_ot_pack_wire(dealer0, 0, pr.first.export_wire()); write_ot_pack_wire(dealer1, 0, pr.second.export_wire()); } /// @brief Online Gilboa mul: OT pads on `dealer`, `d`/`e` vectors on `peer`. inline std::uint64_t gilboa_mul_online(unsigned party, std::uint64_t x, std::uint64_t y, net::stream_array & peer, net::stream_array & dealer, std::size_t bits = 16, std::size_t peer_stream = 0, std::size_t dealer_stream = 0) { auto pack = read_ot_pack_wire(dealer, dealer_stream); auto r = gilboa::mul_from_ot_begin(pack, x, y, party, static_cast(bits)); std::vector peer_d(r.d_share.size()); std::vector peer_e(r.e_share.size()); for (std::size_t i = 0; i < r.d_share.size(); ++i) { detail::exchange_pod(peer, peer_stream, r.d_share[i], peer_d[i]); detail::exchange_pod(peer, peer_stream, r.e_share[i], peer_e[i]); } return gilboa::mul_from_ot_finish(r, peer_d, peer_e); } // --------------------------------------------------------------------------- // RSS ring refresh on a 3-party stream clique // --------------------------------------------------------------------------- /// @brief Send `mine` to the next party; receive predecessor's payload. /// @details Writes `(own, from_prev)` into `own_out` / `next_out` (RSS pair). inline void rss_refresh_ring_online(net::memory_stream_clique & clique, unsigned me, const std::uint8_t * mine, std::size_t n, std::uint8_t * own_out, std::uint8_t * next_out, std::size_t stream = 0) { if (clique.parties != 3 || me > 2) throw std::invalid_argument("rss_refresh_ring_online"); const unsigned next = static_cast((me + 1) % 3); const unsigned prev = static_cast((me + 2) % 3); clique.end(me, next).write(stream, mine, n); clique.end(me, next).flush(stream); clique.end(me, prev).read(stream, next_out, n); std::memcpy(own_out, mine, n); } } // namespace factory } // namespace dpf #endif