/// @file dpf/protocol_factory.hpp /// @brief Dealer and online factories over `net::stream_array`. /// @details A dealer functor writes one view per party on stream `i`. An online /// factory reads blinds from a dealer array and swaps messages on a /// peer array. Both arrays are the same type so the backend can change /// without rewriting the protocol. #ifndef LIBDPF_INCLUDE_DPF_PROTOCOL_FACTORY_HPP__ #define LIBDPF_INCLUDE_DPF_PROTOCOL_FACTORY_HPP__ #include #include #include #include #include #include #include #include #include #include #include "hedley/hedley.h" #include "dpf/net/asio_ns.hpp" #include "dpf/edabit.hpp" #include "dpf/net/stream_array.hpp" #include "dpf/ot_pack.hpp" #include "dpf/random.hpp" namespace dpf { namespace factory { namespace detail { template void write_pod(net::stream_array & a, std::size_t i, const T & v) { static_assert(std::is_trivially_copyable_v, "factory view must be trivially copyable"); a.write(i, &v, sizeof(T)); } template T read_pod(net::stream_array & a, std::size_t i) { static_assert(std::is_trivially_copyable_v, "factory view must be trivially copyable"); T v{}; a.read(i, &v, sizeof(T)); return v; } template void exchange_pod(net::stream_array & peer, std::size_t i, const T & mine, T & theirs) { static_assert(std::is_trivially_copyable_v, "factory message must be trivially copyable"); peer.write(i, &mine, sizeof(T)); peer.flush(i); peer.read(i, &theirs, sizeof(T)); } } // namespace detail // --------------------------------------------------------------------------- // Dealer functors (prep atoms) // --------------------------------------------------------------------------- struct ring_triple_view { std::uint8_t a[8]{}; std::uint8_t b[8]{}; std::uint8_t c[8]{}; std::uint16_t limb = 8; }; struct dabit_view { std::uint8_t bit = 0; std::uint8_t arith[8]{}; std::uint16_t limb = 8; }; struct bit_triple_view { std::uint8_t a = 0; std::uint8_t b = 0; std::uint8_t c = 0; }; inline std::pair deal_ring_triple(std::uint16_t limb) { if (limb == 0 || limb > 8) throw std::invalid_argument("deal_ring_triple limb"); const std::uint64_t mask = limb >= 8 ? ~std::uint64_t{0} : ((std::uint64_t{1} << (8u * limb)) - 1u); const auto a = dpf::uniform_sample() & mask; const auto b = dpf::uniform_sample() & mask; const auto c = (a * b) & mask; const auto a0 = dpf::uniform_sample() & mask; const auto b0 = dpf::uniform_sample() & mask; const auto c0 = dpf::uniform_sample() & mask; ring_triple_view v0{}, v1{}; v0.limb = limb; v1.limb = limb; std::memcpy(v0.a, &a0, limb); std::memcpy(v0.b, &b0, limb); std::memcpy(v0.c, &c0, limb); const auto a1 = (a - a0) & mask; const auto b1 = (b - b0) & mask; const auto c1 = (c - c0) & mask; std::memcpy(v1.a, &a1, limb); std::memcpy(v1.b, &b1, limb); std::memcpy(v1.c, &c1, limb); return {v0, v1}; } inline std::pair deal_dabit(std::uint16_t limb) { if (limb == 0 || limb > 8) throw std::invalid_argument("deal_dabit limb"); auto pair = ot::sample_dabit_pair(); dabit_view v0{}, v1{}; v0.limb = limb; v1.limb = limb; v0.bit = pair.p0.bit; v1.bit = pair.p1.bit; std::memcpy(v0.arith, &pair.p0.arith, limb); std::memcpy(v1.arith, &pair.p1.arith, limb); return {v0, v1}; } inline std::pair deal_bit_triple() { auto t = ot::sample_bit_triple_pair(); return {{t.p0.a, t.p0.b, t.p0.c}, {t.p1.a, t.p1.b, t.p1.c}}; } // --------------------------------------------------------------------------- // make_dealer // --------------------------------------------------------------------------- /// @brief Run each functor `count` times; write view0/view1 on stream `i`. template void make_dealer(net::stream_array & party0, net::stream_array & party1, std::size_t count, Fs &&... fs) { if (party0.size() < sizeof...(Fs) || party1.size() < sizeof...(Fs)) throw std::invalid_argument("make_dealer: not enough streams"); std::size_t stream = 0; auto run_one = [&](auto && f) { for (std::size_t j = 0; j < count; ++j) { auto views = f(); using V0 = std::decay_t; using V1 = std::decay_t; static_assert(std::is_trivially_copyable_v && std::is_trivially_copyable_v, "make_dealer views must be trivially copyable"); if (sizeof(V0) != sizeof(V1)) throw std::logic_error("make_dealer: unequal view sizes"); detail::write_pod(party0, stream, views.first); detail::write_pod(party1, stream, views.second); } party0.flush(stream); party1.flush(stream); ++stream; }; (run_one(std::forward(fs)), ...); } /// @brief Like `make_dealer`, but the functor receives PRG lanes and the index. template void make_dealer_prg(net::stream_array & party0, net::stream_array & party1, std::size_t count, Prg0 & prg0, Prg1 & prg1, F && f, std::size_t stream = 0) { if (party0.size() <= stream || party1.size() <= stream) throw std::invalid_argument("make_dealer_prg: stream index"); for (std::size_t j = 0; j < count; ++j) { auto views = f(prg0, prg1, j); using V0 = std::decay_t; using V1 = std::decay_t; static_assert(std::is_trivially_copyable_v && std::is_trivially_copyable_v, "make_dealer_prg views must be trivially copyable"); if (sizeof(V0) != sizeof(V1)) throw std::logic_error("make_dealer_prg: unequal view sizes"); detail::write_pod(party0, stream, views.first); detail::write_pod(party1, stream, views.second); } party0.flush(stream); party1.flush(stream); } /// @brief Three-party dealer. Functors return `std::array`. template void make_dealer(net::stream_array & party0, net::stream_array & party1, net::stream_array & party2, std::size_t count, Fs &&... fs) { if (party0.size() < sizeof...(Fs) || party1.size() < sizeof...(Fs) || party2.size() < sizeof...(Fs)) throw std::invalid_argument("make_dealer 3p: not enough streams"); std::size_t stream = 0; auto run_one = [&](auto && f) { for (std::size_t j = 0; j < count; ++j) { auto views = f(); using Arr = std::decay_t; static_assert(std::tuple_size_v == 3, "3-party dealer functor must return array/tuple of 3"); using V = std::decay_t(views))>; static_assert(std::is_trivially_copyable_v, "make_dealer views must be trivially copyable"); detail::write_pod(party0, stream, std::get<0>(views)); detail::write_pod(party1, stream, std::get<1>(views)); detail::write_pod(party2, stream, std::get<2>(views)); } party0.flush(stream); party1.flush(stream); party2.flush(stream); ++stream; }; (run_one(std::forward(fs)), ...); } // --------------------------------------------------------------------------- // Online protocol (blocking round list) // --------------------------------------------------------------------------- struct gmw_and_blind { std::uint8_t a = 0; std::uint8_t b = 0; std::uint8_t c = 0; }; struct gmw_and_msg { std::uint8_t d = 0; std::uint8_t e = 0; }; struct gmw_and_keep { gmw_and_blind blind{}; gmw_and_msg mine{}; }; /// @brief Round 0: `(p,q), blind -> {msg, keep}`. struct gmw_and_round0 { std::pair operator()( std::pair pq, gmw_and_blind blind) const { gmw_and_msg m{}; m.d = static_cast((pq.first ^ blind.a) & 1u); m.e = static_cast((pq.second ^ blind.b) & 1u); return {m, gmw_and_keep{blind, m}}; } }; /// @brief Round 1: finish AND from peer masks. struct gmw_and_round1 { unsigned party = 0; std::uint8_t operator()(gmw_and_keep keep, gmw_and_msg peer, gmw_and_blind /*unused_blind*/, gmw_and_blind /*corr*/) const { const std::uint8_t d_open = static_cast(keep.mine.d ^ peer.d); const std::uint8_t e_open = static_cast(keep.mine.e ^ peer.e); ot::bit_triple t{keep.blind.a, keep.blind.b, keep.blind.c}; return edabit::and_finish(t, d_open, e_open, party); } }; /// @brief Two-round online object: round0 produces a peer message; round1 /// consumes the peer reply. Blinds come from `dealer[0]` (round0 only /// for GMW AND). template class two_round_protocol { public: two_round_protocol(std::size_t count, net::stream_array & peer, net::stream_array & dealer, Round0 r0, Round1 r1) : count_(count), peer_(&peer), dealer_(&dealer), r0_(std::move(r0)), r1_(std::move(r1)) { if (peer_->size() < 1 || dealer_->size() < 1) throw std::invalid_argument("two_round_protocol: empty arrays"); } template auto operator()(std::size_t index, const Input & input) { if (index >= count_) throw std::out_of_range("protocol index"); using Blind = gmw_and_blind; Blind blind = detail::read_pod(*dealer_, 0); auto step0 = r0_(input, blind); using Msg = std::decay_t; Msg peer_msg{}; detail::exchange_pod(*peer_, 0, step0.first, peer_msg); Blind unused{}; return r1_(step0.second, peer_msg, unused, unused); } private: std::size_t count_ = 0; net::stream_array * peer_ = nullptr; net::stream_array * dealer_ = nullptr; Round0 r0_; Round1 r1_; }; template class protocol_factory { public: protocol_factory(Round0 r0, Round1 r1) : r0_(std::move(r0)), r1_(std::move(r1)) { } two_round_protocol create(std::size_t count, net::stream_array & peer, net::stream_array & dealer) const { return two_round_protocol(count, peer, dealer, r0_, r1_); } private: Round0 r0_; Round1 r1_; }; template HEDLEY_WARN_UNUSED_RESULT protocol_factory, std::decay_t> make_protocol_factory(Round0 && r0, Round1 && r1) { return protocol_factory, std::decay_t>( std::forward(r0), std::forward(r1)); } /// @brief Dealer functor writing `gmw_and_blind` shares on stream 0. inline auto make_gmw_and_dealer_functor() { return [] { auto t = deal_bit_triple(); gmw_and_blind v0{t.first.a, t.first.b, t.first.c}; gmw_and_blind v1{t.second.a, t.second.b, t.second.c}; return std::pair{v0, v1}; }; } // --------------------------------------------------------------------------- // Byte-oriented N-round online runner // --------------------------------------------------------------------------- struct round_spec { std::size_t blind_bytes = 0; std::size_t msg_bytes = 0; }; struct byte_round { std::size_t blind_bytes = 0; std::size_t msg_bytes = 0; std::function(std::vector & state, const std::uint8_t * blind, std::size_t nb)> produce; std::function & state, const std::uint8_t * peer, std::size_t nm, const std::uint8_t * blind, std::size_t nb)> finish; }; /// @brief Run `rounds` in order: blind from `dealer[r]`, message on `peer[r]`. class byte_protocol { public: byte_protocol(std::size_t count, net::stream_array & peer, net::stream_array & dealer, std::vector rounds) : count_(count), peer_(&peer), dealer_(&dealer), rounds_(std::move(rounds)) { if (peer_->size() < rounds_.size() || dealer_->size() < rounds_.size()) throw std::invalid_argument("byte_protocol: not enough streams"); for (const auto & r : rounds_) { if (!r.produce || !r.finish) throw std::invalid_argument("byte_protocol: missing callbacks"); } } std::vector operator()(std::size_t index, std::vector state) { if (index >= count_) throw std::out_of_range("byte_protocol index"); (void)index; for (std::size_t r = 0; r < rounds_.size(); ++r) { const auto & spec = rounds_[r]; std::vector blind(spec.blind_bytes); if (spec.blind_bytes != 0) dealer_->read(r, blind.data(), spec.blind_bytes); auto outbound = spec.produce(state, blind.data(), blind.size()); if (outbound.size() != spec.msg_bytes) throw std::logic_error("byte_protocol: produce size"); peer_->write(r, outbound.data(), outbound.size()); peer_->flush(r); std::vector inbound(spec.msg_bytes); peer_->read(r, inbound.data(), inbound.size()); spec.finish(state, inbound.data(), inbound.size(), blind.data(), blind.size()); } return state; } std::size_t count() const noexcept { return count_; } std::size_t rounds() const noexcept { return rounds_.size(); } private: std::size_t count_ = 0; net::stream_array * peer_ = nullptr; net::stream_array * dealer_ = nullptr; std::vector rounds_; }; class byte_protocol_factory { public: explicit byte_protocol_factory(std::vector rounds) : rounds_(std::move(rounds)) { } byte_protocol create(std::size_t count, net::stream_array & peer, net::stream_array & dealer) const { return byte_protocol(count, peer, dealer, rounds_); } private: std::vector rounds_; }; HEDLEY_WARN_UNUSED_RESULT inline byte_protocol_factory make_byte_protocol_factory( std::vector rounds) { return byte_protocol_factory(std::move(rounds)); } // --------------------------------------------------------------------------- // Async-shaped byte protocol (blocking today; completion hook for later I/O) // --------------------------------------------------------------------------- /// @brief Same round shape as `byte_round`; used by `async_byte_protocol`. struct async_byte_round { std::size_t blind_bytes = 0; std::size_t msg_bytes = 0; std::function(std::vector & state, const std::uint8_t * blind, std::size_t nb)> produce; std::function & state, const std::uint8_t * peer, std::size_t nm, const std::uint8_t * blind, std::size_t nb)> finish; }; /// @brief Blocking runner with an optional per-round completion callback. /// @details This keeps the legacy synchronous shape: rounds run on the calling /// thread and the callback is a progress hook. For TRUE overlapped, /// event-driven I/O with no busy-waiting and no blocking protocol /// calls, use `dpf::async::overlapped_byte_protocol` in /// `dpf/async_protocol.hpp`, which drives the same `async_byte_round` /// list over an `async_stream_array` and completes via an /// `asio::io_context`. class async_byte_protocol { public: using on_round_complete_fn = std::function; async_byte_protocol(std::size_t count, net::stream_array & peer, net::stream_array & dealer, std::vector rounds, on_round_complete_fn on_round_complete = {}) : count_(count), peer_(&peer), dealer_(&dealer), rounds_(std::move(rounds)), on_round_complete_(std::move(on_round_complete)) { if (peer_->size() < rounds_.size() || dealer_->size() < rounds_.size()) throw std::invalid_argument("async_byte_protocol: not enough streams"); for (const auto & r : rounds_) { if (!r.produce || !r.finish) throw std::invalid_argument("async_byte_protocol: missing callbacks"); } } std::vector operator()(std::size_t index, std::vector state) { if (index >= count_) throw std::out_of_range("async_byte_protocol index"); (void)index; for (std::size_t r = 0; r < rounds_.size(); ++r) { const auto & spec = rounds_[r]; std::vector blind(spec.blind_bytes); if (spec.blind_bytes != 0) dealer_->read(r, blind.data(), spec.blind_bytes); auto outbound = spec.produce(state, blind.data(), blind.size()); if (outbound.size() != spec.msg_bytes) throw std::logic_error("async_byte_protocol: produce size"); peer_->write(r, outbound.data(), outbound.size()); peer_->flush(r); std::vector inbound(spec.msg_bytes); peer_->read(r, inbound.data(), inbound.size()); spec.finish(state, inbound.data(), inbound.size(), blind.data(), blind.size()); if (on_round_complete_) on_round_complete_(r); } return state; } std::size_t count() const noexcept { return count_; } std::size_t rounds() const noexcept { return rounds_.size(); } private: std::size_t count_ = 0; net::stream_array * peer_ = nullptr; net::stream_array * dealer_ = nullptr; std::vector rounds_; on_round_complete_fn on_round_complete_; }; class async_byte_protocol_factory { public: explicit async_byte_protocol_factory(std::vector rounds) : rounds_(std::move(rounds)) { } async_byte_protocol create(std::size_t count, net::stream_array & peer, net::stream_array & dealer, async_byte_protocol::on_round_complete_fn on_round_complete = {}) const { return async_byte_protocol(count, peer, dealer, rounds_, std::move(on_round_complete)); } private: std::vector rounds_; }; HEDLEY_WARN_UNUSED_RESULT inline async_byte_protocol_factory make_async_byte_protocol_factory( std::vector rounds) { return async_byte_protocol_factory(std::move(rounds)); } // --------------------------------------------------------------------------- // Beaver ring multiply (open d, e; finish locally) // --------------------------------------------------------------------------- struct beaver_mul_input { std::uint64_t x = 0; std::uint64_t y = 0; }; struct beaver_mul_keep { std::uint64_t x = 0; std::uint64_t y = 0; std::uint64_t a = 0; std::uint64_t b = 0; std::uint64_t c = 0; std::uint64_t d_open = 0; std::uint64_t d_mine = 0; std::uint64_t e_mine = 0; }; /// @brief Round 0: open `d = x - a` on peer stream 0. struct beaver_mul_round0 { std::uint16_t limb = 8; std::pair, beaver_mul_keep> operator()( beaver_mul_input in, ring_triple_view blind) const { if (blind.limb != limb) throw std::invalid_argument("beaver_mul_round0 limb"); std::uint64_t a = 0, b = 0, c = 0; std::memcpy(&a, blind.a, limb); std::memcpy(&b, blind.b, limb); std::memcpy(&c, blind.c, limb); const std::uint64_t mask = limb >= 8 ? ~std::uint64_t{0} : ((std::uint64_t{1} << (8u * limb)) - 1u); const std::uint64_t d = (in.x - a) & mask; std::vector msg(limb); std::memcpy(msg.data(), &d, limb); beaver_mul_keep keep{}; keep.x = in.x & mask; keep.y = in.y & mask; keep.a = a; keep.b = b; keep.c = c; keep.d_mine = d; return {std::move(msg), keep}; } }; /// @brief Round 1: open `e = y - b`, finish product share. struct beaver_mul_round1 { std::uint16_t limb = 8; unsigned party = 0; std::uint64_t operator()(beaver_mul_keep keep, std::vector peer_e, ring_triple_view /*unused*/) const { if (peer_e.size() != limb) throw std::invalid_argument("beaver_mul_round1 msg"); const std::uint64_t mask = limb >= 8 ? ~std::uint64_t{0} : ((std::uint64_t{1} << (8u * limb)) - 1u); std::uint64_t e_peer = 0; std::memcpy(&e_peer, peer_e.data(), limb); const std::uint64_t e_open = (keep.e_mine + e_peer) & mask; std::uint64_t z = (keep.c + keep.d_open * keep.b + e_open * keep.a) & mask; if (party == 0) z = (z + keep.d_open * e_open) & mask; return z; } }; /// @brief Two-round beaver multiply: blind triple on dealer stream 0 only. template class beaver_mul_protocol { public: beaver_mul_protocol(std::size_t count, net::stream_array & peer, net::stream_array & dealer, Round0 r0, Round1 r1, std::uint16_t limb) : count_(count), peer_(&peer), dealer_(&dealer), r0_(std::move(r0)), r1_(std::move(r1)), limb_(limb) { if (peer_->size() < 2 || dealer_->size() < 1) throw std::invalid_argument("beaver_mul_protocol streams"); } std::uint64_t operator()(std::size_t index, beaver_mul_input input) { if (index >= count_) throw std::out_of_range("beaver_mul_protocol index"); ring_triple_view blind{}; blind.limb = limb_; dealer_->read(0, &blind, sizeof(blind)); auto step0 = r0_(input, blind); peer_->write(0, step0.first.data(), step0.first.size()); peer_->flush(0); std::vector peer_d(limb_); peer_->read(0, peer_d.data(), limb_); std::uint64_t d_peer = 0; std::memcpy(&d_peer, peer_d.data(), limb_); const std::uint64_t mask = limb_ >= 8 ? ~std::uint64_t{0} : ((std::uint64_t{1} << (8u * limb_)) - 1u); auto keep = step0.second; keep.d_open = (keep.d_mine + d_peer) & mask; const std::uint64_t e = (keep.y - keep.b) & mask; keep.e_mine = e; std::vector msg_e(limb_); std::memcpy(msg_e.data(), &e, limb_); peer_->write(1, msg_e.data(), limb_); peer_->flush(1); std::vector peer_e(limb_); peer_->read(1, peer_e.data(), limb_); ring_triple_view unused{}; return r1_(keep, std::move(peer_e), unused); } private: std::size_t count_ = 0; net::stream_array * peer_ = nullptr; net::stream_array * dealer_ = nullptr; Round0 r0_; Round1 r1_; std::uint16_t limb_ = 8; }; class beaver_mul_factory { public: beaver_mul_factory(beaver_mul_round0 r0, beaver_mul_round1 r1) : r0_(std::move(r0)), r1_(std::move(r1)) { } beaver_mul_protocol create( std::size_t count, net::stream_array & peer, net::stream_array & dealer) const { return beaver_mul_protocol( count, peer, dealer, r0_, r1_, r0_.limb); } private: beaver_mul_round0 r0_; beaver_mul_round1 r1_; }; HEDLEY_WARN_UNUSED_RESULT inline beaver_mul_factory make_beaver_mul_factory(beaver_mul_round0 r0, beaver_mul_round1 r1) { return beaver_mul_factory(std::move(r0), std::move(r1)); } inline auto make_beaver_mul_dealer_functor() { return [limb = std::uint16_t{8}] { return deal_ring_triple(limb); }; } } // namespace factory } // namespace dpf #endif