/// @file dpf/ot_pack.hpp /// @brief Consuming cursor over correlated B2A / bit / bit×ring pads. /// @details Dealer mode builds one shared tape and splits it into party views. /// Two parties each calling `dealer()` independently is unsupported; /// use `sample_dealer_pair`. #ifndef LIBDPF_INCLUDE_DPF_OT_PACK_HPP__ #define LIBDPF_INCLUDE_DPF_OT_PACK_HPP__ #include #include #include #include #include #include #include "hedley/hedley.h" #include "simde/simde/x86/avx2.h" #include "dpf/beaver.hpp" #include "dpf/random.hpp" namespace dpf { namespace ot { template struct dabit { std::uint8_t bit = 0; Ring arith{}; }; struct bit_triple { std::uint8_t a = 0; std::uint8_t b = 0; std::uint8_t c = 0; }; /// @brief Bit × ring triple: `(a0⊕a1)·(b0+b1) = c0+c1`. template struct bit_ring_triple { std::uint8_t a = 0; ///< XOR share of the bit Ring b{}; ///< Additive share of the scalar Ring c{}; ///< Additive share of the product }; struct b2a_slot { std::uint8_t r = 0; std::uint64_t add = 0; }; template struct dabit_pair { dabit p0; dabit p1; }; struct bit_triple_pair { bit_triple p0; bit_triple p1; }; template struct bit_ring_triple_pair { bit_ring_triple p0; bit_ring_triple p1; }; template HEDLEY_WARN_UNUSED_RESULT dabit_pair sample_dabit_pair() { struct pad { std::uint8_t bit() { return static_cast( dpf::uniform_sample() & 1u); } simde__m128i block() { return dpf::uniform_sample(); } } p; auto ba = beavers::sample_bit_arith(p); dabit_pair out; out.p0.bit = ba.xor0; out.p0.arith = static_cast(ba.add0); out.p1.bit = ba.xor1; out.p1.arith = static_cast(ba.add1); return out; } HEDLEY_WARN_UNUSED_RESULT inline bit_triple_pair sample_bit_triple_pair() { const std::uint8_t a = static_cast( dpf::uniform_sample() & 1u); const std::uint8_t b = static_cast( dpf::uniform_sample() & 1u); const std::uint8_t c = static_cast(a & b); const std::uint8_t a0 = static_cast( dpf::uniform_sample() & 1u); const std::uint8_t b0 = static_cast( dpf::uniform_sample() & 1u); const std::uint8_t c0 = static_cast( dpf::uniform_sample() & 1u); bit_triple_pair out; out.p0 = {a0, b0, c0}; out.p1 = {static_cast(a ^ a0), static_cast(b ^ b0), static_cast(c ^ c0)}; return out; } template HEDLEY_WARN_UNUSED_RESULT bit_ring_triple_pair sample_bit_ring_triple_pair() { // a is XOR-shared with party 1 holding 0 so a0+a1 = a0⊕a1 in the ring. const std::uint8_t a = static_cast( dpf::uniform_sample() & 1u); const Ring b = dpf::uniform_sample(); const Ring c = static_cast(Ring{a} * b); const Ring b0 = dpf::uniform_sample(); const Ring c0 = dpf::uniform_sample(); bit_ring_triple_pair out; out.p0 = {a, b0, c0}; out.p1 = {0, static_cast(b - b0), static_cast(c - c0)}; return out; } class pack { public: pack() = default; void load_b2a(std::vector slots) { mode_ = mode::pads; b2a_ = std::move(slots); b2a_pos_ = 0; } void load_bits(std::vector bits) { mode_ = mode::pads; bits_ = std::move(bits); bit_pos_ = 0; } template void load_bit_ring(std::vector> triples) { mode_ = mode::pads; bit_ring_.clear(); bit_ring_.reserve(triples.size() * sizeof(bit_ring_triple)); for (const auto & t : triples) { const auto * p = reinterpret_cast(&t); bit_ring_.insert(bit_ring_.end(), p, p + sizeof(t)); } bit_ring_elem_ = sizeof(bit_ring_triple); bit_ring_pos_ = 0; } /// @brief One party's view of a pre-split correlated tape. static pack from_party_view(int me, std::vector b2a, std::vector bits, std::vector bit_ring_bytes = {}, std::size_t bit_ring_elem = 0) { pack p; p.mode_ = mode::pads; p.me_ = me; p.b2a_ = std::move(b2a); p.bits_ = std::move(bits); p.bit_ring_ = std::move(bit_ring_bytes); p.bit_ring_elem_ = bit_ring_elem; return p; } HEDLEY_WARN_UNUSED_RESULT std::size_t remaining_b2a() const noexcept { return b2a_.size() > b2a_pos_ ? b2a_.size() - b2a_pos_ : 0; } HEDLEY_WARN_UNUSED_RESULT std::size_t remaining_bits() const noexcept { return bits_.size() > bit_pos_ ? bits_.size() - bit_pos_ : 0; } HEDLEY_WARN_UNUSED_RESULT std::size_t remaining_bit_ring() const noexcept { if (bit_ring_elem_ == 0) return 0; const std::size_t n = bit_ring_.size() / bit_ring_elem_; return n > bit_ring_pos_ ? n - bit_ring_pos_ : 0; } template dabit take_dabit() { if (b2a_pos_ >= b2a_.size()) throw std::runtime_error("ot_pack: no B2A slots left"); const auto & s = b2a_[b2a_pos_++]; dabit out; out.bit = s.r; out.arith = static_cast(s.add); return out; } bit_triple take_bit_triple() { if (bit_pos_ >= bits_.size()) throw std::runtime_error("ot_pack: no bit triples left"); return bits_[bit_pos_++]; } template bit_ring_triple take_bit_ring() { if (bit_ring_elem_ != sizeof(bit_ring_triple)) throw std::runtime_error("ot_pack: bit_ring element size mismatch"); if (remaining_bit_ring() == 0) throw std::runtime_error("ot_pack: no bit×ring triples left"); bit_ring_triple t{}; std::memcpy(&t, bit_ring_.data() + bit_ring_pos_ * bit_ring_elem_, sizeof(t)); ++bit_ring_pos_; return t; } int party() const noexcept { return me_; } /// @brief Remember a party tape written by `fill_tape_ot_pair`. template void stash_party_tape(const beavers::party_tape & tape) { tape_elem_ = sizeof(Ring); tape_lambda_ready_ = tape.lambda_ready; tape_mono_ready_ = tape.monomial_ready; tape_bundle_ready_ = tape.bundles_ready; tape_dot_ready_ = tape.dot_ready; tape_lambda_ = bytes_of(tape.lambda); tape_mono_ = bytes_of(tape.monomial); tape_bundle_ = bytes_of(tape.bundles); tape_dot_ = bytes_of(tape.dot_cross); has_tape_ = true; } bool has_party_tape() const noexcept { return has_tape_; } /// @brief Serializable OT pad view for dealer / online streams. struct wire { int me = 0; std::vector b2a; std::vector bits; std::vector bit_ring; std::size_t bit_ring_elem = 0; }; HEDLEY_WARN_UNUSED_RESULT wire export_wire() const { wire w; w.me = me_; w.b2a = b2a_; w.bits = bits_; w.bit_ring = bit_ring_; w.bit_ring_elem = bit_ring_elem_; return w; } static pack from_wire(wire w) { pack p; p.mode_ = mode::pads; p.me_ = w.me; p.b2a_ = std::move(w.b2a); p.bits_ = std::move(w.bits); p.bit_ring_ = std::move(w.bit_ring); p.bit_ring_elem_ = w.bit_ring_elem; return p; } template beavers::party_tape load_party_tape() const { if (!has_tape_ || tape_elem_ != sizeof(Ring)) throw std::runtime_error( "ot_pack: no stashed party tape for this ring"); beavers::party_tape t; t.lambda_ready = tape_lambda_ready_; t.monomial_ready = tape_mono_ready_; t.bundles_ready = tape_bundle_ready_; t.dot_ready = tape_dot_ready_; t.lambda = vec_of(tape_lambda_); t.monomial = vec_of(tape_mono_); t.bundles = vec_of(tape_bundle_); t.dot_cross = vec_of(tape_dot_); return t; } private: template static std::vector bytes_of(const std::vector & v) { std::vector out(v.size() * sizeof(Ring)); if (!out.empty()) std::memcpy(out.data(), v.data(), out.size()); return out; } template static std::vector vec_of(const std::vector & b) { if (b.size() % sizeof(Ring) != 0) throw std::runtime_error("ot_pack: tape byte size"); std::vector out(b.size() / sizeof(Ring)); if (!out.empty()) std::memcpy(out.data(), b.data(), b.size()); return out; } enum class mode : unsigned char { empty, pads }; mode mode_ = mode::empty; int me_ = 0; std::vector b2a_; std::vector bits_; std::vector bit_ring_; std::size_t bit_ring_elem_ = 0; std::size_t b2a_pos_ = 0; std::size_t bit_pos_ = 0; std::size_t bit_ring_pos_ = 0; bool has_tape_ = false; std::size_t tape_elem_ = 0; std::vector tape_lambda_ready_; std::vector tape_mono_ready_; std::vector tape_bundle_ready_; std::vector tape_dot_ready_; std::vector tape_lambda_; std::vector tape_mono_; std::vector tape_bundle_; std::vector tape_dot_; }; /// @brief Sample correlated pads and return both party views. template HEDLEY_WARN_UNUSED_RESULT std::pair sample_dealer_pair(std::size_t n_dabit, std::size_t n_bit_triple = 0, std::size_t n_bit_ring = 0) { std::vector b0, b1; b0.reserve(n_dabit); b1.reserve(n_dabit); for (std::size_t i = 0; i < n_dabit; ++i) { auto d = sample_dabit_pair(); b0.push_back({d.p0.bit, static_cast(d.p0.arith)}); b1.push_back({d.p1.bit, static_cast(d.p1.arith)}); } std::vector t0, t1; t0.reserve(n_bit_triple); t1.reserve(n_bit_triple); for (std::size_t i = 0; i < n_bit_triple; ++i) { auto tp = sample_bit_triple_pair(); t0.push_back(tp.p0); t1.push_back(tp.p1); } std::vector> r0, r1; r0.reserve(n_bit_ring); r1.reserve(n_bit_ring); for (std::size_t i = 0; i < n_bit_ring; ++i) { auto rp = sample_bit_ring_triple_pair(); r0.push_back(rp.p0); r1.push_back(rp.p1); } auto pack0 = pack::from_party_view(0, std::move(b0), std::move(t0)); auto pack1 = pack::from_party_view(1, std::move(b1), std::move(t1)); if (n_bit_ring != 0) { pack0.load_bit_ring(std::move(r0)); pack1.load_bit_ring(std::move(r1)); } return {std::move(pack0), std::move(pack1)}; } } // namespace ot } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_OT_PACK_HPP__