/// @file dpf/share_expr.hpp /// @brief Imperative expression recorder over composer + beaver sessions. #ifndef LIBDPF_INCLUDE_DPF_SHARE_EXPR_HPP__ #define LIBDPF_INCLUDE_DPF_SHARE_EXPR_HPP__ #include #include #include #include #include #include #include #include #include #include "hedley/hedley.h" #include "dpf/beaver.hpp" #include "dpf/compose.hpp" #include "dpf/net/edge_mesh.hpp" #include "dpf/net/memory_sink.hpp" #include "dpf/share_cmp.hpp" #include "dpf/trunc.hpp" #include "dpf/xor_wrapper.hpp" namespace dpf { namespace expr { struct handle { std::uint32_t id = 0; }; class recorder { public: explicit recorder(std::size_t party, std::size_t parties = 2) : party_(party), parties_(parties), composer_(party) { if (parties_ != 2 && parties_ != 3) throw std::invalid_argument("recorder parties must be 2 or 3"); if (party_ >= parties_) throw std::invalid_argument("recorder party index"); } std::size_t party() const noexcept { return party_; } std::size_t parties() const noexcept { return parties_; } protocol::composer & composer() noexcept { return composer_; } const protocol::composer & composer() const noexcept { return composer_; } void bind_mesh(net::edge_mesh mesh) { mesh_ = std::move(mesh); } bool has_mesh() const noexcept { return mesh_.size() != 0; } template handle input(Ring clear, unsigned owner) { Ring mine{}; if (parties_ == 2) { auto [s0, s1] = share_cmp::share_input(clear, owner); mine = (party_ == 0) ? s0 : s1; } else mine = (party_ == owner) ? clear : Ring{}; values_[next_] = encode(mine); kinds_[next_] = kind::value; return handle{next_++}; } template handle bind(Ring mine) { values_[next_] = encode(mine); kinds_[next_] = kind::value; return handle{next_++}; } template handle add(handle a, handle b) { check(a); check(b); Ring va = decode(values_.at(a.id)); Ring vb = decode(values_.at(b.id)); values_[next_] = encode(static_cast(va + vb)); kinds_[next_] = kind::value; auto na = composer_.input(protocol::domain::a, sizeof(Ring)); auto nb = composer_.input(protocol::domain::a, sizeof(Ring)); composer_.compute(protocol::opcodes::user_base, {na, nb}, protocol::domain::a, sizeof(Ring)); return handle{next_++}; } /// @brief Record a product; local share filled by `complete_mul` with peer. template handle mul(handle a, handle b) { check(a); check(b); pending_mul pm; pm.a = a.id; pm.b = b.id; pm.out = next_; pending_.push_back(pm); values_[next_] = encode(Ring{}); // filled by complete_mul kinds_[next_] = kind::pending_mul; auto na = composer_.input(protocol::domain::a, sizeof(Ring)); auto nb = composer_.input(protocol::domain::a, sizeof(Ring)); composer_.compute(protocol::opcodes::user_base + 2, {na, nb}, protocol::domain::a, sizeof(Ring)); pending_muls_++; return handle{next_++}; } /// @brief Finish one pending mul with the peer's operand shares (Beaver). template static void complete_mul(recorder & r0, recorder & r1, handle out0, handle out1) { if (r0.kinds_.at(out0.id) != kind::pending_mul || r1.kinds_.at(out1.id) != kind::pending_mul) throw std::invalid_argument("complete_mul: not pending"); pending_mul pm0{}, pm1{}; for (const auto & p : r0.pending_) if (p.out == out0.id) pm0 = p; for (const auto & p : r1.pending_) if (p.out == out1.id) pm1 = p; const Ring x0 = r0.local_value(handle{pm0.a}); const Ring y0 = r0.local_value(handle{pm0.b}); const Ring x1 = r1.local_value(handle{pm1.a}); const Ring y1 = r1.local_value(handle{pm1.b}); // Honest-dealer triple; each party opens only its masked limb. const Ring a0 = dpf::uniform_sample(); const Ring b0 = dpf::uniform_sample(); const Ring a1 = dpf::uniform_sample(); const Ring b1 = dpf::uniform_sample(); const Ring a = static_cast(a0 + a1); const Ring b = static_cast(b0 + b1); const Ring c = static_cast(a * b); const Ring c0 = dpf::uniform_sample(); const Ring c1 = static_cast(c - c0); const Ring d = static_cast((x0 - a0) + (x1 - a1)); const Ring e = static_cast((y0 - b0) + (y1 - b1)); const Ring z0 = trunc::mul_trunc_party(x0, y0, a0, b0, c0, d, e, 0, 0); const Ring z1 = trunc::mul_trunc_party(x1, y1, a1, b1, c1, d, e, 0, 1); r0.values_[out0.id] = encode(z0); r1.values_[out1.id] = encode(z1); r0.kinds_[out0.id] = kind::value; r1.kinds_[out1.id] = kind::value; } /// @brief Multi-input AND via a bit-wire beaver monomial (2PC). handle band(const std::vector & bits) { if (bits.empty()) throw std::invalid_argument("band empty"); and_width_ = bits.size(); pending_and pa; pa.ins.reserve(bits.size()); for (auto h : bits) { check(h); pa.ins.push_back(h.id); } pa.out = next_; pending_ands_.push_back(pa); values_[next_] = {0}; kinds_[next_] = kind::pending_and; return handle{next_++}; } static void complete_band(recorder & r0, recorder & r1, handle out0, handle out1) { pending_and pa0{}, pa1{}; for (const auto & p : r0.pending_ands_) if (p.out == out0.id) pa0 = p; for (const auto & p : r1.pending_ands_) if (p.out == out1.id) pa1 = p; if (pa0.ins.size() != pa1.ins.size() || pa0.ins.empty()) throw std::invalid_argument("complete_band"); using bit_ring = dpf::xor_wrapper; using btraits = beavers::ring_traits; beavers::session s; std::vector> ws; ws.reserve(pa0.ins.size()); auto to_xw = [](std::uint8_t b) { return (b & 1u) ? btraits::one() : btraits::zero(); }; for (std::size_t i = 0; i < pa0.ins.size(); ++i) { auto w = s.bit(); ws.push_back(w); const std::uint8_t b0 = r0.values_.at(pa0.ins[i]).empty() ? 0 : (r0.values_.at(pa0.ins[i])[0] & 1u); const std::uint8_t b1 = r1.values_.at(pa1.ins[i]).empty() ? 0 : (r1.values_.at(pa1.ins[i])[0] & 1u); s.bind_shares(w, to_xw(b0), to_xw(b1)); } beavers::wire acc = ws[0]; for (std::size_t i = 1; i < ws.size(); ++i) acc = s.product(acc, ws[i]); s.pin(acc); s.sample(); s.evaluate(); const auto v = s.value(acc); r0.values_[out0.id] = { static_cast(static_cast(v.p0) & 1u)}; r1.values_[out1.id] = { static_cast(static_cast(v.p1) & 1u)}; r0.kinds_[out0.id] = kind::value; r1.kinds_[out1.id] = kind::value; } std::size_t last_and_arity() const noexcept { return and_width_; } std::size_t pending_muls() const noexcept { return pending_muls_; } template Ring local_value(handle h) const { check(h); if (kinds_.at(h.id) != kind::value) throw std::logic_error("local_value: pending op"); return decode(values_.at(h.id)); } /// @brief Declassify after flushing any pending ops; sums opened shares. template static Ring declassify(recorder & a, recorder & b, handle ha, handle hb) { if (a.kinds_.at(ha.id) == kind::pending_mul) complete_mul(a, b, ha, hb); if (a.kinds_.at(ha.id) == kind::pending_and) complete_band(a, b, ha, hb); if (!a.has_mesh() || !b.has_mesh()) throw std::logic_error("expr declassify requires bind_mesh"); protocol::composer c0(a.party_); protocol::composer c1(b.party_); auto in0 = c0.input(protocol::domain::a, sizeof(Ring)); auto ex0 = c0.exchange(in0); auto in1 = c1.input(protocol::domain::a, sizeof(Ring)); auto ex1 = c1.exchange(in1); auto p0 = c0.default_plan(); auto p1 = c1.default_plan(); std::vector> v0(p0.nodes().size()); std::vector> v1(p1.nodes().size()); const Ring mine_a = a.local_value(ha); const Ring mine_b = b.local_value(hb); v0[in0.id].assign(sizeof(Ring), 0); v1[in1.id].assign(sizeof(Ring), 0); std::memcpy(v0[in0.id].data(), &mine_a, sizeof(Ring)); std::memcpy(v1[in1.id].data(), &mine_b, sizeof(Ring)); std::map kernels; std::exception_ptr err; std::thread t0([&] { try { protocol::drive_via_schedule(p0, a.mesh_.at(net::edge_peer), v0, kernels, a.party_); } catch (...) { err = std::current_exception(); } }); std::thread t1([&] { try { protocol::drive_via_schedule(p1, b.mesh_.at(net::edge_peer), v1, kernels, b.party_); } catch (...) { err = std::current_exception(); } }); t0.join(); t1.join(); if (err) std::rethrow_exception(err); Ring opened{}; std::memcpy(&opened, v0[ex0.id].data(), sizeof(Ring)); (void)ex1; return opened; } template void install_tape(const beavers::party_tape & tape) { session_.install_party(static_cast(party_), tape); } beavers::session & session_u64() noexcept { return session_; } protocol::plan plan() const { return composer_.default_plan(); } private: enum class kind : unsigned char { value, pending_mul, pending_and }; struct pending_mul { std::uint32_t a = 0; std::uint32_t b = 0; std::uint32_t out = 0; }; struct pending_and { std::vector ins; std::uint32_t out = 0; }; void check(handle h) const { if (values_.find(h.id) == values_.end()) throw std::out_of_range("expr handle"); } template static std::vector encode(Ring v) { std::vector out(sizeof(Ring)); std::memcpy(out.data(), &v, sizeof(Ring)); return out; } template static Ring decode(const std::vector & b) { if (b.size() < sizeof(Ring)) throw std::invalid_argument("expr decode"); Ring v{}; std::memcpy(&v, b.data(), sizeof(Ring)); return v; } std::size_t party_ = 0; std::size_t parties_ = 2; protocol::composer composer_; net::edge_mesh mesh_{}; std::uint32_t next_ = 1; std::map> values_; std::map kinds_; beavers::session session_; std::vector pending_; std::vector pending_ands_; std::size_t and_width_ = 0; std::size_t pending_muls_ = 0; }; template HEDLEY_WARN_UNUSED_RESULT Ring eval_mul_add(Ring x, Ring y, Ring w) { return static_cast(x * y + w); } } // namespace expr } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_SHARE_EXPR_HPP__