/// @file dpf/beaver.hpp /// @brief Dealer sampling of Beaver and ABY2.0 multiplication material. /// @details One blind is bound to each wire. A list of product formulae is /// compiled into the monomials of those blinds that a single opening /// round needs, with repeated operands and repeated formulae sharing /// one blind and one product share. A later round that mentions an /// old wire keeps that blind and samples only the new monomials. /// /// Classic Beaver triples are the same objects when every factor is a /// fresh wire. `sample_beaver2` / `sample_beaver3` / `sample_fresh` /// do that. A session is the ABY2.0 form: δ = x + λ is opened once /// per wire, and λ stays with the wire. /// /// A polynomial is a sum of monomials in several wires. /// `2 + 3*x + 4*y + 5*x*y + 6*pow(x, 2) + pow(x, 2)*y + x*y*z` /// is one round. `λ_x²` is stored once whether it appears as `x²`, /// inside `x² y`, or in a second polynomial. Wires that occur with the /// same exponents in every term, as in `a3*(x*z)^3 + a2*(x*z)^2 + a1*(x*z) + a0`, /// are multiplied first and the univariate polynomial is a later round. /// A factor shared by every term, such as a sign or a piecewise scale, /// is applied after the quotient when that uses fewer preprocessing /// values. A lone secret summand is added from its value share. /// /// Doerner–Shelat's per-level AND is a `bit_mul` of a fresh bit and /// a fresh block. A wildcard leaf is a `scale` of one scalar by each /// lane of the unit vector. Those call sites still draw on their own /// pad and `uniform_fill` streams so matched key tapes stay put; /// new steps should record a formula here and `sample` it. /// The namespace is `dpf::beavers` because `dpf::beaver` is already /// the wildcard triple stored on a key. /// /// The session is a dealer: it keeps each full λ so a later round can /// multiply blinds without another multiplication protocol. Parties /// receive only the additive `split`s. /// /// `oracle` collapses those draws onto `dpf::randomness::lane_table`. /// The default PRG is `dpf::prg::aes128`; any PRG with the usual /// `eval` interface can be substituted. Each blind role is one lane /// (wire id, or `mono_role` / `dot_role` for a derived product), and /// the matching share mask is the same index on the tweaked master. /// Copy `index` along every lane is one fresh triple. A thousand /// copies are a thousand indices, not a thousand lanes. `seed()` /// replays every blind and every share mask. /// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors) /// @license Released under a GNU General Public v2.0 (GPLv2) license; /// see [LICENSE.md](@ref license) for details. #ifndef LIBDPF_INCLUDE_DPF_BEAVER_HPP__ #define LIBDPF_INCLUDE_DPF_BEAVER_HPP__ #include #include #include #include #include #include #include #include #include #include #include "hedley/hedley.h" #include "dpf/buffered_prg.hpp" #include "dpf/random.hpp" #include "dpf/xor_wrapper.hpp" namespace dpf { namespace beavers { /// Ring operations used to build and consume triples. /// Specialize for a ring whose multiplicative identity is not `Ring{1}`. template struct ring_traits { static Ring zero() { return Ring{}; } static Ring one() { return Ring{1}; } static Ring add(const Ring & a, const Ring & b) { return static_cast(a + b); } static Ring sub(const Ring & a, const Ring & b) { return static_cast(a - b); } static Ring mul(const Ring & a, const Ring & b) { return static_cast(a * b); } static Ring neg(const Ring & a) { return static_cast(-a); } static Ring sample() { return dpf::uniform_sample(); } }; /// Bitwise AND uses the all-ones word as its multiplicative identity. template struct ring_traits> { using ring = dpf::xor_wrapper; static ring zero() { return ring{}; } static ring one() { using u = typename ring::value_type; return ring{static_cast(~u{0})}; } static ring add(const ring & a, const ring & b) { return a + b; } static ring sub(const ring & a, const ring & b) { return a - b; } static ring mul(const ring & a, const ring & b) { return a * b; } static ring neg(const ring & a) { return -a; } static ring sample() { return dpf::uniform_sample(); } }; template struct default_sampler { Ring operator()() const { return ring_traits::sample(); } }; /// Additive (2,2) split. `open()` is `p0 + p1`. template struct split { Ring p0{}; Ring p1{}; Ring open() const { return ring_traits::add(p0, p1); } dpf::additive_share party0() const { return dpf::additive_share::from_raw(p0); } dpf::additive_share party1() const { return dpf::additive_share::from_raw(p1); } friend bool operator==(const split & a, const split & b) { return a.p0 == b.p0 && a.p1 == b.p1; } friend bool operator!=(const split & a, const split & b) { return !(a == b); } }; template class session; /// A value in a session. Copying a wire copies its id; it does not copy the /// blind. The ring argument is on the type so `a * x * x` can build an /// expression without naming the session. template class wire { friend class session; public: constexpr wire() noexcept = default; constexpr std::uint32_t id() const noexcept { return id_; } constexpr session * owner() const noexcept { return sess_; } friend bool operator==(wire a, wire b) noexcept { return a.sess_ == b.sess_ && a.id_ == b.id_; } friend bool operator!=(wire a, wire b) noexcept { return !(a == b); } private: session * sess_ = nullptr; std::uint32_t id_ = 0; }; /// Unevaluated sum of monomials. `*` distributes over `+`. A public /// coefficient scales a term. Nothing is sampled until `session::operator()`. template struct expr { struct term { Ring coeff{}; /// Positive exponents, sorted by wire id. std::vector> powers; }; session * sess = nullptr; std::vector terms; }; template struct is_beaver_wire : std::false_type {}; template struct is_beaver_wire> : std::true_type {}; template struct is_beaver_expr : std::false_type {}; template struct is_beaver_expr> : std::true_type {}; template expr wire_expr(wire w); template expr horner_expr(wire x, std::initializer_list coeffs); /// One PRG lane per blind role, plus the share-mask stream for that role. /// `blind(role, index)` and `share(role, index, value)` do not depend on /// call order. Walking `index` forward stays inside a refilled window. /// `PRG` defaults to `dpf::prg::aes128`. template class oracle { static_assert(std::is_trivially_copyable_v, "prg oracle lanes require a trivially copyable ring"); public: using prg_type = PRG; using seed_type = typename PRG::block_type; using traits = ring_traits; /// Monomial roles sit above wire ids. Dot-cross roles sit above those. static constexpr std::uint32_t mono_role_base = 0x40000000u; static constexpr std::uint32_t dot_role_base = 0x80000000u; static constexpr std::uint32_t wire_role(std::uint32_t id) noexcept { return id; } static constexpr std::uint32_t mono_role(std::uint32_t i) noexcept { return mono_role_base + i; } static constexpr std::uint32_t dot_role(std::uint32_t gate) noexcept { return dot_role_base + gate; } /// Fused within-polynomial λ combinations (Appendix E groupings). static constexpr std::uint32_t bundle_role_base = 0xC0000000u; static constexpr std::uint32_t bundle_role(std::uint32_t i) noexcept { return bundle_role_base + i; } explicit oracle(std::size_t window = 256u) : lanes_(window) { } explicit oracle(seed_type seed, std::size_t window = 256u) : lanes_(std::move(seed), window) { } const seed_type & seed() const noexcept { return lanes_.seed(); } Ring blind(std::uint32_t role, std::uint64_t index) const { return lanes_.value_at(role, index); } Ring mask(std::uint32_t role, std::uint64_t index) const { return lanes_.mask_at(role, index); } /// Additive split of `value`. The mask is the role's mask stream at `index`. split share(std::uint32_t role, std::uint64_t index, const Ring & value) const { Ring p0 = mask(role, index); return split{p0, traits::sub(value, p0)}; } void fill_blinds(std::uint32_t role, std::uint64_t index, Ring * out, std::size_t n) const { lanes_.fill_values(role, index, out, n); } void fill_masks(std::uint32_t role, std::uint64_t index, Ring * out, std::size_t n) const { lanes_.fill_masks(role, index, out, n); } private: dpf::randomness::lane_table lanes_; }; /// Shares produced for one copy index of a recorded formula. template struct prg_material { std::vector> lambda; std::vector> monomial; /// Fused λ-combinations for polynomial gates, in bundle index order. std::vector> bundles; /// Parallel to the session's gates. Empty split when the gate is not a dot. std::vector> dot_cross; }; /// Dealer session: record formulae, `sample` blinds and monomials, `bind` /// input secrets, `evaluate` every round. template class session { static_assert(!std::is_same_v, bool>, "beaver bit wires use bit(), not bool arithmetic"); static_assert(!(std::is_integral_v && std::is_signed_v), "beaver rings are unsigned mod 2^n; use uintN_t or dpf::modint"); public: using traits = ring_traits; using exp_list = std::vector>; using wire = ::dpf::beavers::wire; /// One factor of a monomial query: `{{x, 2}, {a, 1}}`. struct power { wire base{}; unsigned exp = 1; }; session() = default; session(const session &) = delete; session & operator=(const session &) = delete; session(session &&) = delete; session & operator=(session &&) = delete; /// Arithmetic input. Its blind is sampled once and then reused. HEDLEY_WARN_UNUSED_RESULT wire input() { return emplace_wire(0, true, false); } /// Sample this wire's blind even if no recorded formula opens it. /// One-shot triples use this for the product wire, so it can be reused. void pin(wire w) { wires_[check(w)].pinned = true; } /// 0/1 wire in this ring. `bind` accepts only `zero()` or `one()` /// (`1` for integer rings, the all-ones word for `xor_wrapper`). HEDLEY_WARN_UNUSED_RESULT wire bit() { return emplace_wire(0, true, true); } /// One-round product. Repeated wires share a blind. HEDLEY_WARN_UNUSED_RESULT wire product(std::initializer_list factors) { std::vector ids; ids.reserve(factors.size()); for (auto w : factors) ids.push_back(check(w)); return commit_product(std::move(ids)); } HEDLEY_WARN_UNUSED_RESULT wire product(wire a, wire b) { return commit_product({check(a), check(b)}); } HEDLEY_WARN_UNUSED_RESULT wire product(wire a, wire b, wire c) { return commit_product({check(a), check(b), check(c)}); } /// Record a sum of monomials. Like terms share one blind product. HEDLEY_WARN_UNUSED_RESULT wire operator()(const expr & e) { if (e.sess != this) throw std::invalid_argument("beaver expression is from a different session"); return commit_poly(e); } /// `c[0] + c[1] x + c[2] x^2 + ...` in one round. HEDLEY_WARN_UNUSED_RESULT wire horner(wire x, std::initializer_list coeffs) { return (*this)(horner_expr(x, coeffs)); } /// Sign-corrected Horner: `sign * (c[0] + c[1] x + ...)`, still one round /// when `sign` and `x` are inputs. HEDLEY_WARN_UNUSED_RESULT wire horner(wire sign, wire x, std::initializer_list coeffs) { return (*this)(wire_expr(sign) * horner_expr(x, coeffs)); } HEDLEY_WARN_UNUSED_RESULT wire square(wire x) { return product(x, x); } /// One-round `a * x * x` (one blind for `x`). HEDLEY_WARN_UNUSED_RESULT wire mul_square(wire a, wire x) { return product({a, x, x}); } template HEDLEY_WARN_UNUSED_RESULT wire dot(const ContX & xs, const ContY & ys) { std::vector x; std::vector y; for (const auto & w : xs) x.push_back(check(w)); for (const auto & w : ys) y.push_back(check(w)); return commit_dot(std::move(x), std::move(y)); } HEDLEY_WARN_UNUSED_RESULT wire dot(std::initializer_list xs, std::initializer_list ys) { std::vector x; std::vector y; x.reserve(xs.size()); y.reserve(ys.size()); for (auto w : xs) x.push_back(check(w)); for (auto w : ys) y.push_back(check(w)); return commit_dot(std::move(x), std::move(y)); } /// `z_i = scalar * lanes[i]`, one output wire per lane. The scalar blind /// is shared. Each lane gets its own `λ_s λ_i` share. template HEDLEY_WARN_UNUSED_RESULT std::vector scale(wire scalar, const Cont & lanes) { std::vector out; for (const auto & lane : lanes) out.push_back(product(scalar, lane)); return out; } HEDLEY_WARN_UNUSED_RESULT std::vector scale(wire scalar, std::initializer_list lanes) { std::vector out; out.reserve(lanes.size()); for (auto lane : lanes) out.push_back(product(scalar, lane)); return out; } /// `bit * scalar`. `bit` must come from `bit()`. HEDLEY_WARN_UNUSED_RESULT wire bit_mul(wire selector, wire scalar) { auto id = check(selector); if (!wires_[id].is_bit) throw std::invalid_argument("bit_mul selector must come from bit()"); return commit_product({id, check(scalar)}); } /// One-round `selector ? when1 : when0`, i.e. `when0 + selector * (when1 - when0)`. HEDLEY_WARN_UNUSED_RESULT wire mux(wire selector, wire when1, wire when0) { auto ib = check(selector); auto ix = check(when1); auto iy = check(when0); int round = 1; round = std::max(round, wires_[ib].ready_round + 1); round = std::max(round, wires_[ix].ready_round + 1); round = std::max(round, wires_[iy].ready_round + 1); auto out = emplace_wire(round, false, false); require_product_monomials({ib, ix}); require_product_monomials({ib, iy}); gate g; g.kind = gate_kind::mux; g.out = out.id_; g.lhs = {ib, ix, iy}; gates_.push_back(std::move(g)); wires_[out.id_].gate = static_cast(gates_.size() - 1); return out; } /// Sample every missing wire blind and every missing monomial. /// Blinds already sampled are left alone. template void sample(Sample && sampler) { auto & rng = sampler; for (std::uint32_t id = 0; id < wires_.size(); ++id) { auto & w = wires_[id]; if (w.lambda_ready || !needs_blind(id)) continue; w.lambda_full = rng(); w.lambda = share_of(w.lambda_full, rng); w.lambda_ready = true; } for (auto & m : monos_) { if (m.ready) continue; Ring full = traits::one(); for (auto [id, e] : m.key) { if (!wires_[id].lambda_ready) throw std::logic_error("beaver blind is missing"); full = traits::mul(full, pow_small(wires_[id].lambda_full, e)); } m.share = share_of(full, rng); m.ready = true; } for (auto & b : bundles_) { if (b.ready) continue; b.share = share_of(bundle_value(b), rng); b.ready = true; } for (auto & g : gates_) { if (g.kind != gate_kind::dot || g.cross_ready) continue; Ring sum = traits::zero(); for (std::size_t i = 0; i < g.lhs.size(); ++i) { sum = traits::add(sum, traits::mul( wires_[g.lhs[i]].lambda_full, wires_[g.rhs[i]].lambda_full)); } g.cross = share_of(sum, rng); g.cross_ready = true; } } void sample() { sample(default_sampler{}); } /// Install missing blinds and product shares from copy `index` of `src`. /// Already-sampled wires keep their λ. New product shares are built from /// those stored blinds, then split with the oracle's share lane. template void sample_from(const oracle & src, std::uint64_t index = 0) { for (std::uint32_t id = 0; id < wires_.size(); ++id) { auto & w = wires_[id]; if (w.lambda_ready || !needs_blind(id)) continue; w.lambda_full = src.blind(oracle::wire_role(id), index); w.lambda = src.share(oracle::wire_role(id), index, w.lambda_full); w.lambda_ready = true; } for (std::uint32_t i = 0; i < monos_.size(); ++i) { auto & m = monos_[i]; if (m.ready) continue; Ring full = mono_full(m.key); m.share = src.share(oracle::mono_role(i), index, full); m.ready = true; } for (std::uint32_t i = 0; i < bundles_.size(); ++i) { auto & b = bundles_[i]; if (b.ready) continue; b.share = src.share(oracle::bundle_role(i), index, bundle_value(b)); b.ready = true; } for (std::uint32_t g = 0; g < gates_.size(); ++g) { auto & gate = gates_[g]; if (gate.kind != gate_kind::dot || gate.cross_ready) continue; gate.cross = src.share(oracle::dot_role(g), index, dot_full(gate)); gate.cross_ready = true; } } /// Every wire, monomial, and dot cross of this formula at copy `index`. /// Does not change the session. Copies are independent lanes samples, so /// `material_at(src, 5)` does not depend on having asked for 0..4. template prg_material material_at(const oracle & src, std::uint64_t index) const { prg_material out; std::vector full(wires_.size()); out.lambda.resize(wires_.size()); for (std::uint32_t id = 0; id < wires_.size(); ++id) { full[id] = src.blind(oracle::wire_role(id), index); out.lambda[id] = src.share(oracle::wire_role(id), index, full[id]); } out.monomial.resize(monos_.size()); for (std::uint32_t i = 0; i < monos_.size(); ++i) { Ring prod = traits::one(); for (auto [id, e] : monos_[i].key) prod = traits::mul(prod, pow_small(full[id], e)); out.monomial[i] = src.share(oracle::mono_role(i), index, prod); } out.bundles.resize(bundles_.size()); for (std::uint32_t i = 0; i < bundles_.size(); ++i) { Ring value = traits::zero(); for (const auto & part : bundles_[i].parts) { Ring m = traits::one(); for (auto [id, e] : part.lam) m = traits::mul(m, pow_small(full[id], e)); value = traits::add(value, traits::mul(part.coeff, m)); } out.bundles[i] = src.share(oracle::bundle_role(i), index, value); } out.dot_cross.resize(gates_.size()); for (std::uint32_t g = 0; g < gates_.size(); ++g) { if (gates_[g].kind != gate_kind::dot) continue; Ring sum = traits::zero(); for (std::size_t i = 0; i < gates_[g].lhs.size(); ++i) { sum = traits::add(sum, traits::mul( full[gates_[g].lhs[i]], full[gates_[g].rhs[i]])); } out.dot_cross[g] = src.share(oracle::dot_role(g), index, sum); } return out; } /// Split `secret` into fresh additive shares and bind them to an input. template void bind(wire w, Ring secret, Sample && sampler) { auto id = check(w); if (wires_[id].ready_round != 0) throw std::invalid_argument("only input wires can be bound"); if (wires_[id].value_ready) throw std::logic_error("input wire is already bound"); if (wires_[id].is_bit && !is_bit_secret(secret)) throw std::invalid_argument("bit wire must be 0 or 1"); auto & rng = sampler; wires_[id].value = share_of(secret, rng); wires_[id].value_ready = true; } void bind(wire w, Ring secret) { bind(w, secret, default_sampler{}); } /// Bind shares the caller already holds. Their sum is the secret. void bind_shares(wire w, Ring p0, Ring p1) { auto id = check(w); if (wires_[id].ready_round != 0) throw std::invalid_argument("only input wires can be bound"); if (wires_[id].value_ready) throw std::logic_error("input wire is already bound"); if (wires_[id].is_bit && !is_bit_secret(traits::add(p0, p1))) throw std::invalid_argument("bit wire must be 0 or 1"); wires_[id].value = split{p0, p1}; wires_[id].value_ready = true; } /// Open every ready round. Inputs used by round-1 gates are opened /// together; a gate output is a later round's input and keeps the blind /// chosen in `sample`. void evaluate() { int max_round = 0; for (const auto & w : wires_) max_round = std::max(max_round, w.ready_round); for (int r = 1; r <= max_round; ++r) { for (const auto & g : gates_) { if (wires_[g.out].ready_round != r) continue; if (wires_[g.out].value_ready) continue; split val; if (g.kind == gate_kind::product) { for (auto id : g.lhs) ensure_delta(id); val = eval_product(g.lhs); } else if (g.kind == gate_kind::dot) { for (auto id : g.lhs) ensure_delta(id); for (auto id : g.rhs) ensure_delta(id); val = eval_dot(g); } else if (g.kind == gate_kind::poly) { for (const auto & step : g.steps) { if (step.value_wire >= 0) { const auto id = static_cast(step.value_wire); if (!wires_[id].value_ready) throw std::logic_error("beaver wire is not ready to open"); continue; } for (auto [id, exp] : step.delta) { (void)exp; ensure_delta(id); } if (step.mask_wire >= 0) ensure_delta(static_cast(step.mask_wire)); } val = eval_poly(g); } else { for (auto id : g.lhs) ensure_delta(id); val = eval_mux(g); } wires_[g.out].value = val; wires_[g.out].value_ready = true; if (wires_[g.out].lambda_ready) ensure_delta(g.out); } } } split lambda(wire w) const { auto id = check(w); if (!wires_[id].lambda_ready) throw std::logic_error("call sample() before reading a beaver blind"); return wires_[id].lambda; } /// Share of `Π λ_i^{e_i}`. A lone `λ_w` is the wire blind itself. split monomial(const std::vector & spec) const { std::vector> raw; raw.reserve(spec.size()); for (auto p : spec) { if (p.exp == 0) throw std::invalid_argument("beaver power is zero"); raw.emplace_back(check(p.base), p.exp); } return term_share(merge_exponents(raw)); } split monomial(std::initializer_list spec) const { return monomial(std::vector(spec)); } split value(wire w) const { auto id = check(w); if (!wires_[id].value_ready) throw std::logic_error("beaver wire has no value yet"); return wires_[id].value; } Ring open(wire w) const { auto v = value(w); return traits::add(v.p0, v.p1); } /// Public ABY2.0 mask δ = x + λ, after `evaluate` has opened the wire. Ring delta(wire w) const { auto id = check(w); if (!wires_[id].delta_ready) throw std::logic_error("beaver wire has not been opened"); return wires_[id].delta; } split dot_cross(wire w) const { auto id = check(w); int g = wires_[id].gate; if (g < 0 || gates_[static_cast(g)].kind != gate_kind::dot) throw std::invalid_argument("wire is not a dot output"); if (!gates_[static_cast(g)].cross_ready) throw std::logic_error("call sample() before reading a dot cross term"); return gates_[static_cast(g)].cross; } int round_of(wire w) const { return wires_[check(w)].ready_round; } std::size_t wire_count() const noexcept { return wires_.size(); } /// Product shares beyond the per-wire blinds: subset monomials from /// `product` gates, plus one fused bundle per public-δ class in a /// polynomial (Appendix E). A lone mask is not counted. std::size_t monomial_count() const noexcept { return monos_.size() + bundles_.size(); } /// Wire blinds that the recorded formulae actually open, plus product shares. std::size_t preprocessing_count() const { std::size_t n = monomial_count(); for (std::uint32_t id = 0; id < wires_.size(); ++id) { if (needs_blind(id)) ++n; } return n; } private: enum class gate_kind : unsigned char { product, dot, mux, poly }; struct poly_term { Ring coeff{}; std::vector factors; }; /// One λ-monomial in a fused preprocessing share. struct bundle_part { Ring coeff{}; exp_list lam; }; struct bundle { std::vector parts; split share{}; bool ready = false; }; /// Online: `public(δ) * scale * share`, where share is a wire mask, /// a raw monomial, or a fused sum of monomials. struct poly_step { exp_list delta; Ring scale{}; int bundle = -1; int mask_wire = -1; int value_wire = -1; bool public_only = false; }; struct wstate { bool is_input = false; bool is_bit = false; bool pinned = false; bool lambda_ready = false; bool value_ready = false; bool delta_ready = false; int ready_round = 0; int gate = -1; Ring lambda_full{}; split lambda{}; split value{}; Ring delta{}; }; struct gate { gate_kind kind = gate_kind::product; std::uint32_t out = 0; std::vector lhs; std::vector rhs; std::vector terms; std::vector steps; split cross{}; bool cross_ready = false; }; struct mono { exp_list key; split share{}; bool ready = false; }; std::vector wires_; std::vector gates_; std::vector monos_; std::vector bundles_; std::map mono_index_; static bool is_bit_secret(const Ring & secret) { return secret == traits::zero() || secret == traits::one(); } static Ring pow_small(Ring base, unsigned e) { Ring acc = traits::one(); while (e > 0) { if ((e & 1u) != 0) acc = traits::mul(acc, base); base = traits::mul(base, base); e >>= 1u; } return acc; } static Ring scale_int(Ring v, long long k) { if (k == 0) return traits::zero(); if (k < 0) { v = traits::neg(v); k = -k; } Ring acc = traits::zero(); while (k > 0) { if ((k & 1ll) != 0) acc = traits::add(acc, v); v = traits::add(v, v); k >>= 1ll; } return acc; } static long long binom(unsigned n, unsigned k) { if (k > n) return 0; if (k > n - k) k = n - k; long long c = 1; for (unsigned i = 1; i <= k; ++i) c = c * static_cast(n - k + i) / static_cast(i); return c; } template static split share_of(const Ring & secret, Sample & rng) { Ring p0 = rng(); return split{p0, traits::sub(secret, p0)}; } Ring mono_full(const exp_list & key) const { Ring full = traits::one(); for (auto [id, e] : key) { if (!wires_[id].lambda_ready) throw std::logic_error("beaver blind is missing"); full = traits::mul(full, pow_small(wires_[id].lambda_full, e)); } return full; } Ring dot_full(const gate & g) const { Ring sum = traits::zero(); for (std::size_t i = 0; i < g.lhs.size(); ++i) { if (!wires_[g.lhs[i]].lambda_ready || !wires_[g.rhs[i]].lambda_ready) throw std::logic_error("beaver blind is missing"); sum = traits::add(sum, traits::mul( wires_[g.lhs[i]].lambda_full, wires_[g.rhs[i]].lambda_full)); } return sum; } wire emplace_wire(int round, bool is_input, bool is_bit) { wstate st; st.is_input = is_input; st.is_bit = is_bit; st.ready_round = round; wires_.push_back(std::move(st)); wire w; w.sess_ = this; w.id_ = static_cast(wires_.size() - 1); return w; } std::uint32_t check(wire w) const { if (w.sess_ != this || static_cast(w.id_) >= wires_.size()) throw std::invalid_argument("wire is not from this beaver session"); return w.id_; } static exp_list merge_exponents( const std::vector> & raw) { std::map acc; for (auto [id, e] : raw) { acc[id] += e; if (acc[id] > 16u) throw std::invalid_argument("beaver exponent is too large"); } exp_list key; key.reserve(acc.size()); for (auto [id, e] : acc) { if (e == 0) continue; key.emplace_back(id, static_cast(e)); } return key; } static exp_list group_exponents(const std::vector & factors) { std::vector> raw; raw.reserve(factors.size()); for (auto id : factors) raw.emplace_back(id, 1u); return merge_exponents(raw); } static std::size_t expansion_size(const exp_list & groups) { std::size_t terms = 1; for (auto [id, e] : groups) { (void)id; auto span = static_cast(e) + 1u; if (terms > 4096u / span) throw std::invalid_argument("beaver expansion is too large"); terms *= span; } return terms; } template static void for_each_term(const exp_list & groups, Fn && fn) { const auto n = groups.size(); std::vector exp(n, 0); while (true) { exp_list key; key.reserve(n); for (std::size_t i = 0; i < n; ++i) { if (exp[i] != 0) key.emplace_back(groups[i].first, exp[i]); } fn(key); std::size_t i = 0; for (; i < n; ++i) { ++exp[i]; if (exp[i] <= groups[i].second) break; exp[i] = 0; } if (i == n) break; } } void require_key(const exp_list & key) { if (key.empty() || (key.size() == 1 && key[0].second == 1)) return; if (mono_index_.find(key) != mono_index_.end()) return; auto idx = static_cast(monos_.size()); mono_index_.emplace(key, idx); monos_.push_back(mono{key, {}, false}); } void require_product_monomials(const std::vector & factors) { auto groups = group_exponents(factors); (void)expansion_size(groups); for_each_term(groups, [&](const exp_list & key) { require_key(key); }); } wire commit_product(std::vector factors) { if (factors.empty()) throw std::invalid_argument("beaver product needs at least one factor"); int round = 1; for (auto id : factors) round = std::max(round, wires_[id].ready_round + 1); require_product_monomials(factors); auto out = emplace_wire(round, false, false); gate g; g.kind = gate_kind::product; g.out = out.id_; g.lhs = std::move(factors); gates_.push_back(std::move(g)); wires_[out.id_].gate = static_cast(gates_.size() - 1); return out; } wire commit_poly(const expr & e) { int round = 1; std::vector terms; terms.reserve(e.terms.size()); for (const auto & src : e.terms) { poly_term term; term.coeff = src.coeff; for (auto [id, exp] : src.powers) { if (id >= wires_.size()) throw std::invalid_argument("beaver expression has a bad factor"); round = std::max(round, wires_[id].ready_round + 1); for (std::uint8_t k = 0; k < exp; ++k) term.factors.push_back(id); } terms.push_back(std::move(term)); } (void)round; return schedule_terms(std::move(terms)); } wire emit_terms(std::vector terms) { int round = 1; for (const auto & term : terms) { for (auto id : term.factors) { if (id >= wires_.size()) throw std::invalid_argument("beaver expression has a bad factor"); round = std::max(round, wires_[id].ready_round + 1); } } auto steps = compile_poly(terms); auto out = emplace_wire(round, false, false); gate g; g.kind = gate_kind::poly; g.out = out.id_; g.terms = std::move(terms); g.steps = std::move(steps); gates_.push_back(std::move(g)); wires_[out.id_].gate = static_cast(gates_.size() - 1); return out; } static std::uint8_t factor_exp(const poly_term & term, std::uint32_t id) { std::uint8_t n = 0; for (auto f : term.factors) { if (f == id) ++n; } return n; } std::size_t estimate_pieces(const std::vector> & pieces) { std::vector already(wires_.size(), 0); for (std::uint32_t id = 0; id < wires_.size(); ++id) already[id] = needs_blind(id) ? 1 : 0; auto saved = bundles_; std::vector> compiled; compiled.reserve(pieces.size()); for (const auto & piece : pieces) compiled.push_back(compile_poly(piece)); const auto bundle_base = saved.size(); std::size_t added = bundles_.size() - bundle_base; std::map blinds; auto note = [&](std::uint32_t id) { if (id < already.size() && already[id] != 0) return; blinds[id] = 1; }; for (std::size_t i = bundle_base; i < bundles_.size(); ++i) { for (const auto & part : bundles_[i].parts) { for (auto [wid, exp] : part.lam) { (void)exp; note(wid); } } } for (const auto & steps : compiled) { for (const auto & step : steps) { if (step.mask_wire >= 0) note(static_cast(step.mask_wire)); for (auto [wid, exp] : step.delta) { (void)exp; note(wid); } } } bundles_ = std::move(saved); return added + blinds.size(); } static std::vector drop_one(std::vector terms, std::uint32_t wire_id) { for (auto & term : terms) { auto it = std::find(term.factors.begin(), term.factors.end(), wire_id); if (it != term.factors.end()) term.factors.erase(it); } return terms; } static bool wire_in_every(const std::vector & terms, std::uint32_t wire_id) { for (const auto & term : terms) { if (factor_exp(term, wire_id) == 0) return false; } return !terms.empty(); } wire schedule_terms(std::vector terms) { if (terms.size() < 2) return emit_terms(std::move(terms)); std::vector> best{terms}; std::size_t best_cost = estimate_pieces(best); auto consider = [&](std::vector> seq) { const std::size_t cost = estimate_pieces(seq); if (cost < best_cost) { best = std::move(seq); best_cost = cost; } }; std::map seen; for (const auto & term : terms) for (auto id : term.factors) seen[id] = 1; std::vector common; for (auto [wire_id, present] : seen) { (void)present; if (wire_in_every(terms, wire_id)) common.push_back(wire_id); } const auto fresh = static_cast(wires_.size()); for (auto wire_id : common) { poly_term mul; mul.coeff = traits::one(); mul.factors = {fresh, wire_id}; consider({drop_one(terms, wire_id), {std::move(mul)}}); } std::map, std::vector> clusters; for (auto [wire_id, present] : seen) { (void)present; std::vector shape; shape.reserve(terms.size()); bool any = false; for (const auto & term : terms) { auto e = factor_exp(term, wire_id); shape.push_back(e); any = any || e != 0; } if (any) clusters[std::move(shape)].push_back(wire_id); } for (auto & [shape, group] : clusters) { (void)shape; if (group.size() < 2) continue; poly_term prod; prod.coeff = traits::one(); prod.factors = group; std::vector rewritten; rewritten.reserve(terms.size()); for (const auto & term : terms) { const auto e = factor_exp(term, group[0]); poly_term next; next.coeff = term.coeff; for (auto f : term.factors) { if (std::find(group.begin(), group.end(), f) == group.end()) next.factors.push_back(f); } for (std::uint8_t i = 0; i < e; ++i) next.factors.push_back(fresh); rewritten.push_back(std::move(next)); } consider({{prod}, rewritten}); std::map rewritten_seen; for (const auto & term : rewritten) for (auto id : term.factors) rewritten_seen[id] = 1; const auto later = fresh + 1; for (auto [wire_id, present] : rewritten_seen) { (void)present; if (!wire_in_every(rewritten, wire_id)) continue; poly_term mul; mul.coeff = traits::one(); mul.factors = {later, wire_id}; consider({{prod}, drop_one(rewritten, wire_id), {std::move(mul)}}); } } wire last{}; for (auto & piece : best) last = emit_terms(std::move(piece)); return last; } wire commit_dot(std::vector xs, std::vector ys) { if (xs.empty() || xs.size() != ys.size()) throw std::invalid_argument( "beaver dot operands must have the same non-zero length"); int round = 1; for (std::size_t i = 0; i < xs.size(); ++i) { round = std::max(round, wires_[xs[i]].ready_round + 1); round = std::max(round, wires_[ys[i]].ready_round + 1); } auto out = emplace_wire(round, false, false); gate g; g.kind = gate_kind::dot; g.out = out.id_; g.lhs = std::move(xs); g.rhs = std::move(ys); gates_.push_back(std::move(g)); wires_[out.id_].gate = static_cast(gates_.size() - 1); return out; } void ensure_delta(std::uint32_t id) { auto & w = wires_[id]; if (w.delta_ready) return; if (!w.lambda_ready || !w.value_ready) throw std::logic_error("beaver wire is not ready to open"); w.delta = traits::add( traits::add(w.value.p0, w.value.p1), traits::add(w.lambda.p0, w.lambda.p1)); w.delta_ready = true; } split term_share(const exp_list & key) const { if (key.empty()) throw std::invalid_argument("beaver monomial is empty"); if (key.size() == 1 && key[0].second == 1) { const auto & w = wires_[key[0].first]; if (!w.lambda_ready) throw std::logic_error("call sample() before reading a beaver blind"); return w.lambda; } auto it = mono_index_.find(key); if (it != mono_index_.end() && monos_[it->second].ready) return monos_[it->second].share; for (const auto & b : bundles_) { if (!b.ready || b.parts.size() != 1 || !(b.parts[0].coeff == traits::one())) continue; if (b.parts[0].lam == key) return b.share; } throw std::logic_error("beaver monomial was not prepared"); } unsigned exponent_of(const exp_list & key, std::uint32_t id) const { for (auto [kid, ke] : key) { if (kid == id) return ke; } return 0; } split eval_product(const std::vector & factors) const { auto groups = group_exponents(factors); Ring s0 = traits::zero(); Ring s1 = traits::zero(); for_each_term(groups, [&](const exp_list & key) { long long coeff = 1; bool any = false; Ring pub = traits::one(); for (auto [id, e] : groups) { unsigned ti = exponent_of(key, id); coeff *= binom(e, ti); if ((ti & 1u) != 0) coeff = -coeff; any = any || ti != 0; pub = traits::mul(pub, pow_small(wires_[id].delta, e - ti)); } if (coeff == 0) return; if (!any) { s0 = traits::add(s0, scale_int(pub, coeff)); return; } auto sh = term_share(key); s0 = traits::add(s0, scale_int(traits::mul(pub, sh.p0), coeff)); s1 = traits::add(s1, scale_int(traits::mul(pub, sh.p1), coeff)); }); return split{s0, s1}; } split eval_dot(const gate & g) const { if (!g.cross_ready) throw std::logic_error("beaver dot cross term is not sampled"); Ring s0 = traits::zero(); Ring s1 = traits::zero(); for (std::size_t i = 0; i < g.lhs.size(); ++i) { const auto & x = wires_[g.lhs[i]]; const auto & y = wires_[g.rhs[i]]; s0 = traits::add(s0, traits::mul(x.delta, y.delta)); s0 = traits::sub(s0, traits::mul(x.delta, y.lambda.p0)); s1 = traits::sub(s1, traits::mul(x.delta, y.lambda.p1)); s0 = traits::sub(s0, traits::mul(y.delta, x.lambda.p0)); s1 = traits::sub(s1, traits::mul(y.delta, x.lambda.p1)); } s0 = traits::add(s0, g.cross.p0); s1 = traits::add(s1, g.cross.p1); return split{s0, s1}; } exp_list pair_key(std::uint32_t a, std::uint32_t b) const { if (a == b) return exp_list{{a, 2}}; if (a > b) std::swap(a, b); return exp_list{{a, 1}, {b, 1}}; } bool same_parts(const std::vector & a, const std::vector & b) const { if (a.size() != b.size()) return false; for (std::size_t i = 0; i < a.size(); ++i) { if (!(a[i].coeff == b[i].coeff) || a[i].lam != b[i].lam) return false; } return true; } int require_bundle(std::vector parts) { for (std::size_t i = 0; i < bundles_.size(); ++i) { if (same_parts(bundles_[i].parts, parts)) return static_cast(i); } bundles_.push_back(bundle{std::move(parts), {}, false}); return static_cast(bundles_.size() - 1); } bool needs_blind(std::uint32_t id) const { if (wires_[id].pinned) return true; for (const auto & b : bundles_) { for (const auto & part : b.parts) { for (auto [wid, exp] : part.lam) { (void)exp; if (wid == id) return true; } } } for (const auto & g : gates_) { const auto & factors = g.kind == gate_kind::dot ? g.rhs : g.lhs; if (g.kind != gate_kind::poly) { for (auto f : g.lhs) { if (f == id) return true; } for (auto f : factors) { if (f == id) return true; } continue; } for (const auto & step : g.steps) { if (step.mask_wire == static_cast(id)) return true; for (auto [wid, exp] : step.delta) { (void)exp; if (wid == id) return true; } } } return false; } /// Group λ-monomials that share a public δ monomial into one share. /// A bucket that is only `c · λ_i` reuses the wire blind. std::vector compile_poly(const std::vector & terms) { struct bucket { Ring pub{}; std::map lams; }; std::map buckets; std::vector steps; for (const auto & term : terms) { if (term.factors.size() == 1) { poly_step step; step.scale = term.coeff; step.value_wire = static_cast(term.factors[0]); steps.push_back(std::move(step)); continue; } auto groups = group_exponents(term.factors); (void)expansion_size(groups); for_each_term(groups, [&](const exp_list & key) { long long ncoeff = 1; exp_list delta; exp_list lam; for (auto [id, e] : groups) { unsigned ti = exponent_of(key, id); ncoeff *= binom(e, ti); if ((ti & 1u) != 0) ncoeff = -ncoeff; if (ti != 0) lam.emplace_back(id, static_cast(ti)); if (e > ti) delta.emplace_back(id, static_cast(e - ti)); } if (ncoeff == 0) return; Ring coeff = scale_int(term.coeff, ncoeff); if (coeff == traits::zero()) return; auto & slot = buckets[delta]; if (lam.empty()) slot.pub = traits::add(slot.pub, coeff); else slot.lams[lam] = traits::add(slot.lams[lam], coeff); }); } for (auto & [delta, slot] : buckets) { if (!(slot.pub == traits::zero())) { poly_step step; step.delta = delta; step.scale = slot.pub; step.public_only = true; steps.push_back(std::move(step)); } std::vector parts; for (auto & [lam, coeff] : slot.lams) { if (coeff == traits::zero()) continue; parts.push_back(bundle_part{coeff, lam}); } if (parts.empty()) continue; const bool lone_mask = parts.size() == 1 && parts[0].lam.size() == 1 && parts[0].lam[0].second == 1; if (lone_mask) { poly_step step; step.delta = delta; step.scale = parts[0].coeff; step.mask_wire = static_cast(parts[0].lam[0].first); steps.push_back(std::move(step)); continue; } if (parts.size() == 1) { Ring scale = parts[0].coeff; parts[0].coeff = traits::one(); poly_step step; step.delta = delta; step.scale = scale; step.bundle = require_bundle(std::move(parts)); steps.push_back(std::move(step)); continue; } poly_step step; step.delta = delta; step.scale = traits::one(); step.bundle = require_bundle(std::move(parts)); steps.push_back(std::move(step)); } return steps; } Ring bundle_value(const bundle & b) const { Ring full = traits::zero(); for (const auto & part : b.parts) { Ring m = traits::one(); for (auto [id, e] : part.lam) { if (!wires_[id].lambda_ready) throw std::logic_error("beaver blind is missing"); m = traits::mul(m, pow_small(wires_[id].lambda_full, e)); } full = traits::add(full, traits::mul(part.coeff, m)); } return full; } Ring pow_delta(const exp_list & delta) const { Ring pub = traits::one(); for (auto [id, e] : delta) pub = traits::mul(pub, pow_small(wires_[id].delta, e)); return pub; } split eval_poly(const gate & g) const { Ring s0 = traits::zero(); Ring s1 = traits::zero(); for (const auto & step : g.steps) { if (step.value_wire >= 0) { const auto & val = wires_[static_cast(step.value_wire)].value; s0 = traits::add(s0, traits::mul(step.scale, val.p0)); s1 = traits::add(s1, traits::mul(step.scale, val.p1)); continue; } Ring pub = pow_delta(step.delta); if (step.public_only) { s0 = traits::add(s0, traits::mul(pub, step.scale)); continue; } if (step.mask_wire >= 0) { const auto & lam = wires_[static_cast(step.mask_wire)].lambda; s0 = traits::add(s0, traits::mul(pub, traits::mul(step.scale, lam.p0))); s1 = traits::add(s1, traits::mul(pub, traits::mul(step.scale, lam.p1))); continue; } const auto & share = bundles_[static_cast(step.bundle)].share; if (!bundles_[static_cast(step.bundle)].ready) throw std::logic_error("beaver bundle is not sampled"); s0 = traits::add(s0, traits::mul(pub, traits::mul(step.scale, share.p0))); s1 = traits::add(s1, traits::mul(pub, traits::mul(step.scale, share.p1))); } return split{s0, s1}; } split eval_mux(const gate & g) const { const auto & b = wires_[g.lhs[0]]; const auto & x = wires_[g.lhs[1]]; const auto & y = wires_[g.lhs[2]]; Ring dd = traits::sub(x.delta, y.delta); split ld{ traits::sub(x.lambda.p0, y.lambda.p0), traits::sub(x.lambda.p1, y.lambda.p1)}; auto bx = term_share(pair_key(g.lhs[0], g.lhs[1])); auto by = term_share(pair_key(g.lhs[0], g.lhs[2])); split cross{ traits::sub(bx.p0, by.p0), traits::sub(bx.p1, by.p1)}; Ring s0 = traits::mul(b.delta, dd); Ring s1 = traits::zero(); s0 = traits::sub(s0, traits::mul(b.delta, ld.p0)); s1 = traits::sub(s1, traits::mul(b.delta, ld.p1)); s0 = traits::sub(s0, traits::mul(dd, b.lambda.p0)); s1 = traits::sub(s1, traits::mul(dd, b.lambda.p1)); s0 = traits::add(s0, cross.p0); s1 = traits::add(s1, cross.p1); s0 = traits::add(s0, y.value.p0); s1 = traits::add(s1, y.value.p1); return split{s0, s1}; } }; namespace detail { template void add_power(typename expr::term & term, std::uint32_t id, unsigned exp) { if (exp == 0) return; for (auto & power : term.powers) { if (power.first != id) continue; unsigned sum = static_cast(power.second) + exp; if (sum > 16u) throw std::invalid_argument("beaver exponent is too large"); power.second = static_cast(sum); return; } if (exp > 16u) throw std::invalid_argument("beaver exponent is too large"); term.powers.emplace_back(id, static_cast(exp)); std::sort(term.powers.begin(), term.powers.end()); } template expr merge_terms(expr e) { using traits = ring_traits; std::map>, Ring> acc; for (const auto & term : e.terms) acc[term.powers] = traits::add(acc[term.powers], term.coeff); e.terms.clear(); for (auto & [powers, coeff] : acc) { if (coeff == traits::zero()) continue; typename expr::term term; term.coeff = coeff; term.powers = powers; e.terms.push_back(std::move(term)); } return e; } template expr mul_exprs(const expr & a, const expr & b) { if (a.sess == nullptr || a.sess != b.sess) throw std::invalid_argument("beaver factors are from different sessions"); expr out; out.sess = a.sess; out.terms.reserve(a.terms.size() * b.terms.size()); for (const auto & left : a.terms) { for (const auto & right : b.terms) { typename expr::term term = left; term.coeff = ring_traits::mul(left.coeff, right.coeff); for (auto [id, exp] : right.powers) add_power(term, id, exp); out.terms.push_back(std::move(term)); } } return merge_terms(std::move(out)); } template expr add_exprs(expr a, const expr & b) { if (a.sess == nullptr || a.sess != b.sess) throw std::invalid_argument("beaver factors are from different sessions"); a.terms.insert(a.terms.end(), b.terms.begin(), b.terms.end()); return merge_terms(std::move(a)); } template expr scale_expr(expr e, Ring coeff) { using traits = ring_traits; for (auto & term : e.terms) term.coeff = traits::mul(term.coeff, coeff); return merge_terms(std::move(e)); } template Ring coeff_of(Coeff value) { return Ring{value}; } } // namespace detail template HEDLEY_WARN_UNUSED_RESULT expr wire_expr(wire w) { if (w.owner() == nullptr) throw std::invalid_argument("wire is not from a beaver session"); expr e; e.sess = w.owner(); typename expr::term term; term.coeff = ring_traits::one(); term.powers.push_back({w.id(), 1}); e.terms.push_back(std::move(term)); return e; } template HEDLEY_WARN_UNUSED_RESULT expr horner_expr(wire x, std::initializer_list coeffs) { if (x.owner() == nullptr) throw std::invalid_argument("wire is not from a beaver session"); expr e; e.sess = x.owner(); unsigned power = 0; for (const Ring & coeff : coeffs) { if (power > 16u) throw std::invalid_argument("beaver exponent is too large"); if (!(coeff == ring_traits::zero())) { typename expr::term term; term.coeff = coeff; if (power > 0) term.powers.push_back({x.id(), static_cast(power)}); e.terms.push_back(std::move(term)); } ++power; } return e; } /// `coeff * v0 * v1 * ...`, with repeated wires counting as a power. template HEDLEY_WARN_UNUSED_RESULT expr monomial(Ring coeff, wire first, Wires... rest) { static_assert((std::is_same_v> && ...), "monomial factors are wires"); expr e = wire_expr(first); ((e = e * rest), ...); return detail::scale_expr(std::move(e), coeff); } template HEDLEY_WARN_UNUSED_RESULT expr pow(wire w, unsigned exp) { if (exp == 0) { expr e = wire_expr(w); e.terms.clear(); typename expr::term term; term.coeff = ring_traits::one(); e.terms.push_back(std::move(term)); return e; } if (exp == 1) return wire_expr(w); if (exp > 16u) throw std::invalid_argument("beaver exponent is too large"); expr e = wire_expr(w); e.terms[0].powers[0].second = static_cast(exp); return e; } template HEDLEY_WARN_UNUSED_RESULT expr operator*(wire a, wire b) { return detail::mul_exprs(wire_expr(a), wire_expr(b)); } template HEDLEY_WARN_UNUSED_RESULT expr operator*(expr a, wire b) { return detail::mul_exprs(a, wire_expr(b)); } template HEDLEY_WARN_UNUSED_RESULT expr operator*(wire a, expr b) { return detail::mul_exprs(wire_expr(a), b); } template HEDLEY_WARN_UNUSED_RESULT expr operator*(expr a, expr b) { return detail::mul_exprs(a, b); } template HEDLEY_WARN_UNUSED_RESULT expr operator+(wire a, wire b) { return detail::add_exprs(wire_expr(a), wire_expr(b)); } template HEDLEY_WARN_UNUSED_RESULT expr operator+(expr a, wire b) { return detail::add_exprs(std::move(a), wire_expr(b)); } template HEDLEY_WARN_UNUSED_RESULT expr operator+(wire a, expr b) { return detail::add_exprs(wire_expr(a), b); } template HEDLEY_WARN_UNUSED_RESULT expr operator+(expr a, expr b) { return detail::add_exprs(std::move(a), b); } template HEDLEY_WARN_UNUSED_RESULT expr operator-(expr e) { return detail::scale_expr(std::move(e), ring_traits::neg(ring_traits::one())); } template HEDLEY_WARN_UNUSED_RESULT expr operator-(wire w) { return -wire_expr(w); } template HEDLEY_WARN_UNUSED_RESULT expr operator-(expr a, expr b) { return std::move(a) + (-std::move(b)); } template HEDLEY_WARN_UNUSED_RESULT expr operator-(expr a, wire b) { return std::move(a) + (-b); } template HEDLEY_WARN_UNUSED_RESULT expr operator-(wire a, expr b) { return wire_expr(a) + (-std::move(b)); } template HEDLEY_WARN_UNUSED_RESULT expr operator-(wire a, wire b) { return wire_expr(a) + (-b); } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator*(Coeff coeff, wire w) { return detail::scale_expr(wire_expr(w), detail::coeff_of(coeff)); } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator*(wire w, Coeff coeff) { return detail::coeff_of(coeff) * w; } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator*(Coeff coeff, expr e) { return detail::scale_expr(std::move(e), detail::coeff_of(coeff)); } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator*(expr e, Coeff coeff) { return detail::coeff_of(coeff) * std::move(e); } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator+(Coeff coeff, wire w) { expr constant = wire_expr(w); constant.terms[0].powers.clear(); constant.terms[0].coeff = detail::coeff_of(coeff); return detail::add_exprs(std::move(constant), wire_expr(w)); } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator+(wire w, Coeff coeff) { return coeff + w; } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator+(Coeff coeff, expr e) { expr constant; constant.sess = e.sess; typename expr::term term; term.coeff = detail::coeff_of(coeff); constant.terms.push_back(std::move(term)); return detail::add_exprs(std::move(constant), e); } template && !is_beaver_wire>::value && !is_beaver_expr>::value>> HEDLEY_WARN_UNUSED_RESULT expr operator+(expr e, Coeff coeff) { return coeff + std::move(e); } // --------------------------------------------------------------------------- // One-shot samplers. Fresh wires, so the blinds are a classic Beaver tuple. // `out` is the ABY2.0 blind of the product wire, for a later round. // --------------------------------------------------------------------------- /// `subset[mask - 1]` is `Π λ_i` over the bits set in `mask` (bit i selects /// `in[i]`). Singleton masks are the wire blinds themselves. template struct fresh_beaver { static_assert(Arity >= 2 && Arity <= 8, "fresh beaver arity is 2..8"); std::array, Arity> in{}; std::array, (std::size_t{1} << Arity) - 1> subset{}; split out{}; }; template HEDLEY_WARN_UNUSED_RESULT fresh_beaver sample_fresh(Sample && rng) { static_assert(Arity >= 2 && Arity <= 8, "fresh beaver arity is 2..8"); session s; std::array::wire, Arity> w{}; for (std::size_t i = 0; i < Arity; ++i) w[i] = s.input(); expr e = w[0] * w[1]; for (std::size_t i = 2; i < Arity; ++i) e = e * w[i]; auto out = s(e); s.pin(out); s.sample(std::forward(rng)); fresh_beaver t; for (std::size_t i = 0; i < Arity; ++i) t.in[i] = s.lambda(w[i]); for (unsigned mask = 1; mask < (1u << Arity); ++mask) { std::vector::power> spec; for (std::size_t i = 0; i < Arity; ++i) { if ((mask & (1u << i)) != 0) spec.push_back(typename session::power{w[i], 1u}); } t.subset[mask - 1] = s.monomial(spec); } t.out = s.lambda(out); return t; } template HEDLEY_WARN_UNUSED_RESULT fresh_beaver sample_fresh() { return sample_fresh(default_sampler{}); } template struct beaver2 { split a{}; split b{}; split ab{}; split out{}; }; template struct beaver3 { split a{}; split b{}; split c{}; split ab{}; split ac{}; split bc{}; split abc{}; split out{}; }; template struct square_beaver { split x{}; split x2{}; split out{}; }; template struct mul_square_beaver { split a{}; split x{}; split x2{}; split ax{}; split ax2{}; split out{}; }; template struct dot_beaver { std::vector> x; std::vector> y; split cross{}; split out{}; }; template struct scale_beaver { split scalar{}; std::vector> lanes; std::vector> cross; std::vector> out; }; template struct bit_mul_beaver { split bit{}; split scalar{}; split product{}; split out{}; }; template struct mux_beaver { split bit{}; split when1{}; split when0{}; split bit_when1{}; split bit_when0{}; split out{}; }; /// One fresh Beaver pair from copy `index` of an oracle. /// Roles match a session that records `input, input, product`: wires 0 and 1, /// the product wire, and monomial 0. template HEDLEY_WARN_UNUSED_RESULT beaver2 beaver2_at(const oracle & src, std::uint64_t index) { using traits = ring_traits; Ring a = src.blind(oracle::wire_role(0), index); Ring b = src.blind(oracle::wire_role(1), index); Ring out = src.blind(oracle::wire_role(2), index); Ring ab = traits::mul(a, b); return beaver2{ src.share(oracle::wire_role(0), index, a), src.share(oracle::wire_role(1), index, b), src.share(oracle::mono_role(0), index, ab), src.share(oracle::wire_role(2), index, out)}; } /// `n` copies starting at `begin`. Each role is one contiguous lane read. template void fill_beaver2(const oracle & src, std::uint64_t begin, beaver2 * out, std::size_t n) { if (n == 0) return; if (out == nullptr) throw std::invalid_argument("beaver2 output is null"); using traits = ring_traits; std::vector a(n), b(n), z(n); std::vector ma(n), mb(n), mab(n), mz(n); src.fill_blinds(oracle::wire_role(0), begin, a.data(), n); src.fill_blinds(oracle::wire_role(1), begin, b.data(), n); src.fill_blinds(oracle::wire_role(2), begin, z.data(), n); src.fill_masks(oracle::wire_role(0), begin, ma.data(), n); src.fill_masks(oracle::wire_role(1), begin, mb.data(), n); src.fill_masks(oracle::mono_role(0), begin, mab.data(), n); src.fill_masks(oracle::wire_role(2), begin, mz.data(), n); for (std::size_t i = 0; i < n; ++i) { Ring ab = traits::mul(a[i], b[i]); out[i] = beaver2{ split{ma[i], traits::sub(a[i], ma[i])}, split{mb[i], traits::sub(b[i], mb[i])}, split{mab[i], traits::sub(ab, mab[i])}, split{mz[i], traits::sub(z[i], mz[i])}}; } } template HEDLEY_WARN_UNUSED_RESULT beaver2 sample_beaver2(Sample && rng) { auto t = sample_fresh<2, Ring>(std::forward(rng)); // mask 0b11 = 3, index 2, is λ_a λ_b return beaver2{t.in[0], t.in[1], t.subset[2], t.out}; } template HEDLEY_WARN_UNUSED_RESULT beaver2 sample_beaver2() { return sample_beaver2(default_sampler{}); } template HEDLEY_WARN_UNUSED_RESULT beaver3 sample_beaver3(Sample && rng) { auto t = sample_fresh<3, Ring>(std::forward(rng)); return beaver3{ t.in[0], t.in[1], t.in[2], t.subset[0b011u - 1], t.subset[0b101u - 1], t.subset[0b110u - 1], t.subset[0b111u - 1], t.out}; } template HEDLEY_WARN_UNUSED_RESULT beaver3 sample_beaver3() { return sample_beaver3(default_sampler{}); } template HEDLEY_WARN_UNUSED_RESULT square_beaver sample_square(Sample && rng) { session s; auto x = s.input(); auto out = s(x * x); s.pin(out); s.sample(std::forward(rng)); return square_beaver{ s.lambda(x), s.monomial({{x, 2u}}), s.lambda(out)}; } template HEDLEY_WARN_UNUSED_RESULT square_beaver sample_square() { return sample_square(default_sampler{}); } template HEDLEY_WARN_UNUSED_RESULT mul_square_beaver sample_mul_square(Sample && rng) { session s; auto a = s.input(); auto x = s.input(); auto out = s(a * x * x); s.pin(out); s.sample(std::forward(rng)); return mul_square_beaver{ s.lambda(a), s.lambda(x), s.monomial({{x, 2u}}), s.monomial({{a, 1u}, {x, 1u}}), s.monomial({{a, 1u}, {x, 2u}}), s.lambda(out)}; } template HEDLEY_WARN_UNUSED_RESULT mul_square_beaver sample_mul_square() { return sample_mul_square(default_sampler{}); } template HEDLEY_WARN_UNUSED_RESULT dot_beaver sample_dot(std::size_t n, Sample && rng) { if (n == 0) throw std::invalid_argument("beaver dot length is zero"); session s; std::vector::wire> x; std::vector::wire> y; x.reserve(n); y.reserve(n); for (std::size_t i = 0; i < n; ++i) { x.push_back(s.input()); y.push_back(s.input()); } auto out = s.dot(x, y); s.pin(out); s.sample(std::forward(rng)); dot_beaver t; t.cross = s.dot_cross(out); t.out = s.lambda(out); t.x.reserve(n); t.y.reserve(n); for (std::size_t i = 0; i < n; ++i) { t.x.push_back(s.lambda(x[i])); t.y.push_back(s.lambda(y[i])); } return t; } template HEDLEY_WARN_UNUSED_RESULT dot_beaver sample_dot(std::size_t n) { return sample_dot(n, default_sampler{}); } template HEDLEY_WARN_UNUSED_RESULT scale_beaver sample_scale(std::size_t n, Sample && rng) { if (n == 0) throw std::invalid_argument("beaver scale length is zero"); session s; auto scalar = s.input(); std::vector::wire> lanes; lanes.reserve(n); for (std::size_t i = 0; i < n; ++i) lanes.push_back(s.input()); auto outs = s.scale(scalar, lanes); for (const auto & out : outs) s.pin(out); s.sample(std::forward(rng)); scale_beaver t; t.scalar = s.lambda(scalar); t.lanes.reserve(n); t.cross.reserve(n); t.out.reserve(n); for (std::size_t i = 0; i < n; ++i) { t.lanes.push_back(s.lambda(lanes[i])); t.cross.push_back(s.monomial({{scalar, 1u}, {lanes[i], 1u}})); t.out.push_back(s.lambda(outs[i])); } return t; } template HEDLEY_WARN_UNUSED_RESULT scale_beaver sample_scale(std::size_t n) { return sample_scale(n, default_sampler{}); } template HEDLEY_WARN_UNUSED_RESULT bit_mul_beaver sample_bit_mul(Sample && rng) { session s; auto b = s.bit(); auto x = s.input(); auto out = s.bit_mul(b, x); s.pin(out); s.sample(std::forward(rng)); return bit_mul_beaver{ s.lambda(b), s.lambda(x), s.monomial({{b, 1u}, {x, 1u}}), s.lambda(out)}; } template HEDLEY_WARN_UNUSED_RESULT bit_mul_beaver sample_bit_mul() { return sample_bit_mul(default_sampler{}); } template HEDLEY_WARN_UNUSED_RESULT mux_beaver sample_mux(Sample && rng) { session s; auto b = s.bit(); auto when1 = s.input(); auto when0 = s.input(); auto out = s.mux(b, when1, when0); s.pin(out); s.sample(std::forward(rng)); return mux_beaver{ s.lambda(b), s.lambda(when1), s.lambda(when0), s.monomial({{b, 1u}, {when1, 1u}}), s.monomial({{b, 1u}, {when0, 1u}}), s.lambda(out)}; } template HEDLEY_WARN_UNUSED_RESULT mux_beaver sample_mux() { return sample_mux(default_sampler{}); } } // namespace beavers } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_BEAVER_HPP__