libdpf/include/dpf/beaver.hpp
Ryan Henry 875f09fec1 Record Grotto half-ulp tables and comparison geneval, and factor shared beaver terms before the quotient.
Horner and window evaluation need those tables in the tree. Comparison geneval opens the same value words as a Doerner–Shelat key. A factor common to every polynomial term is multiplied first so that preprocessing stays smaller.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-24 15:16:21 -06:00

2514 lines
77 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @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 <algorithm>
#include <array>
#include <cstddef>
#include <cstdint>
#include <initializer_list>
#include <map>
#include <stdexcept>
#include <type_traits>
#include <utility>
#include <vector>
#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 <typename Ring>
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<Ring>(a + b); }
static Ring sub(const Ring & a, const Ring & b) { return static_cast<Ring>(a - b); }
static Ring mul(const Ring & a, const Ring & b) { return static_cast<Ring>(a * b); }
static Ring neg(const Ring & a) { return static_cast<Ring>(-a); }
static Ring sample() { return dpf::uniform_sample<Ring>(); }
};
/// Bitwise AND uses the all-ones word as its multiplicative identity.
template <typename T>
struct ring_traits<dpf::xor_wrapper<T>>
{
using ring = dpf::xor_wrapper<T>;
static ring zero() { return ring{}; }
static ring one()
{
using u = typename ring::value_type;
return ring{static_cast<u>(~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<ring>(); }
};
template <typename Ring>
struct default_sampler
{
Ring operator()() const { return ring_traits<Ring>::sample(); }
};
/// Additive (2,2) split. `open()` is `p0 + p1`.
template <typename Ring>
struct split
{
Ring p0{};
Ring p1{};
Ring open() const
{
return ring_traits<Ring>::add(p0, p1);
}
dpf::additive_share<Ring, 0> party0() const
{
return dpf::additive_share<Ring, 0>::from_raw(p0);
}
dpf::additive_share<Ring, 1> party1() const
{
return dpf::additive_share<Ring, 1>::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 <typename Ring>
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 <typename Ring>
class wire
{
friend class session<Ring>;
public:
constexpr wire() noexcept = default;
constexpr std::uint32_t id() const noexcept { return id_; }
constexpr session<Ring> * 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<Ring> * 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 <typename Ring>
struct expr
{
struct term
{
Ring coeff{};
/// Positive exponents, sorted by wire id.
std::vector<std::pair<std::uint32_t, std::uint8_t>> powers;
};
session<Ring> * sess = nullptr;
std::vector<term> terms;
};
template <typename T>
struct is_beaver_wire : std::false_type {};
template <typename Ring>
struct is_beaver_wire<wire<Ring>> : std::true_type {};
template <typename T>
struct is_beaver_expr : std::false_type {};
template <typename Ring>
struct is_beaver_expr<expr<Ring>> : std::true_type {};
template <typename Ring>
expr<Ring> wire_expr(wire<Ring> w);
template <typename Ring>
expr<Ring> horner_expr(wire<Ring> x, std::initializer_list<Ring> 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 <typename Ring, typename PRG = dpf::prg::aes128>
class oracle
{
static_assert(std::is_trivially_copyable_v<Ring>,
"prg oracle lanes require a trivially copyable ring");
public:
using prg_type = PRG;
using seed_type = typename PRG::block_type;
using traits = ring_traits<Ring>;
/// 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<Ring> share(std::uint32_t role, std::uint64_t index, const Ring & value) const
{
Ring p0 = mask(role, index);
return split<Ring>{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<Ring, PRG> lanes_;
};
/// Shares produced for one copy index of a recorded formula.
template <typename Ring>
struct prg_material
{
std::vector<split<Ring>> lambda;
std::vector<split<Ring>> monomial;
/// Fused λ-combinations for polynomial gates, in bundle index order.
std::vector<split<Ring>> bundles;
/// Parallel to the session's gates. Empty split when the gate is not a dot.
std::vector<split<Ring>> dot_cross;
};
/// Dealer session: record formulae, `sample` blinds and monomials, `bind`
/// input secrets, `evaluate` every round.
template <typename Ring>
class session
{
static_assert(!std::is_same_v<std::remove_cv_t<Ring>, bool>,
"beaver bit wires use bit(), not bool arithmetic");
static_assert(!(std::is_integral_v<Ring> && std::is_signed_v<Ring>),
"beaver rings are unsigned mod 2^n; use uintN_t or dpf::modint");
public:
using traits = ring_traits<Ring>;
using exp_list = std::vector<std::pair<std::uint32_t, std::uint8_t>>;
using wire = ::dpf::beavers::wire<Ring>;
/// 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<wire> factors)
{
std::vector<std::uint32_t> 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<Ring> & 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<Ring> 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<Ring> 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 <typename ContX, typename ContY>
HEDLEY_WARN_UNUSED_RESULT
wire dot(const ContX & xs, const ContY & ys)
{
std::vector<std::uint32_t> x;
std::vector<std::uint32_t> 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<wire> xs, std::initializer_list<wire> ys)
{
std::vector<std::uint32_t> x;
std::vector<std::uint32_t> 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 <typename Cont>
HEDLEY_WARN_UNUSED_RESULT
std::vector<wire> scale(wire scalar, const Cont & lanes)
{
std::vector<wire> out;
for (const auto & lane : lanes)
out.push_back(product(scalar, lane));
return out;
}
HEDLEY_WARN_UNUSED_RESULT
std::vector<wire> scale(wire scalar, std::initializer_list<wire> lanes)
{
std::vector<wire> 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<int>(gates_.size() - 1);
return out;
}
/// Sample every missing wire blind and every missing monomial.
/// Blinds already sampled are left alone.
template <typename Sample>
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<Ring>{});
}
/// 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 <typename PRG>
void sample_from(const oracle<Ring, PRG> & 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<Ring>::wire_role(id), index);
w.lambda = src.share(oracle<Ring>::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<Ring>::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<Ring>::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<Ring>::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 <typename PRG>
prg_material<Ring> material_at(const oracle<Ring, PRG> & src, std::uint64_t index) const
{
prg_material<Ring> out;
std::vector<Ring> full(wires_.size());
out.lambda.resize(wires_.size());
for (std::uint32_t id = 0; id < wires_.size(); ++id)
{
full[id] = src.blind(oracle<Ring>::wire_role(id), index);
out.lambda[id] = src.share(oracle<Ring>::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<Ring>::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<Ring>::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<Ring>::dot_role(g), index, sum);
}
return out;
}
/// Split `secret` into fresh additive shares and bind them to an input.
template <typename Sample>
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<Ring>{});
}
/// 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<Ring>{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<Ring> 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<std::uint32_t>(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<std::uint32_t>(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<Ring> 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<Ring> monomial(const std::vector<power> & spec) const
{
std::vector<std::pair<std::uint32_t, unsigned>> 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<Ring> monomial(std::initializer_list<power> spec) const
{
return monomial(std::vector<power>(spec));
}
split<Ring> 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<Ring> dot_cross(wire w) const
{
auto id = check(w);
int g = wires_[id].gate;
if (g < 0 || gates_[static_cast<std::size_t>(g)].kind != gate_kind::dot)
throw std::invalid_argument("wire is not a dot output");
if (!gates_[static_cast<std::size_t>(g)].cross_ready)
throw std::logic_error("call sample() before reading a dot cross term");
return gates_[static_cast<std::size_t>(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<std::uint32_t> factors;
};
/// One λ-monomial in a fused preprocessing share.
struct bundle_part
{
Ring coeff{};
exp_list lam;
};
struct bundle
{
std::vector<bundle_part> parts;
split<Ring> 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<Ring> lambda{};
split<Ring> value{};
Ring delta{};
};
struct gate
{
gate_kind kind = gate_kind::product;
std::uint32_t out = 0;
std::vector<std::uint32_t> lhs;
std::vector<std::uint32_t> rhs;
std::vector<poly_term> terms;
std::vector<poly_step> steps;
split<Ring> cross{};
bool cross_ready = false;
};
struct mono
{
exp_list key;
split<Ring> share{};
bool ready = false;
};
std::vector<wstate> wires_;
std::vector<gate> gates_;
std::vector<mono> monos_;
std::vector<bundle> bundles_;
std::map<exp_list, std::uint32_t> 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<long long>(n - k + i) / static_cast<long long>(i);
return c;
}
template <typename Sample>
static split<Ring> share_of(const Ring & secret, Sample & rng)
{
Ring p0 = rng();
return split<Ring>{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<std::uint32_t>(wires_.size() - 1);
return w;
}
std::uint32_t check(wire w) const
{
if (w.sess_ != this || static_cast<std::size_t>(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<std::pair<std::uint32_t, unsigned>> & raw)
{
std::map<std::uint32_t, unsigned> 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<std::uint8_t>(e));
}
return key;
}
static exp_list group_exponents(const std::vector<std::uint32_t> & factors)
{
std::vector<std::pair<std::uint32_t, unsigned>> 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<std::size_t>(e) + 1u;
if (terms > 4096u / span)
throw std::invalid_argument("beaver expansion is too large");
terms *= span;
}
return terms;
}
template <typename Fn>
static void for_each_term(const exp_list & groups, Fn && fn)
{
const auto n = groups.size();
std::vector<std::uint8_t> 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<std::uint32_t>(monos_.size());
mono_index_.emplace(key, idx);
monos_.push_back(mono{key, {}, false});
}
void require_product_monomials(const std::vector<std::uint32_t> & 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<std::uint32_t> 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<int>(gates_.size() - 1);
return out;
}
wire commit_poly(const expr<Ring> & e)
{
int round = 1;
std::vector<poly_term> 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<poly_term> 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<int>(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<std::vector<poly_term>> & pieces)
{
std::vector<char> 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<std::vector<poly_step>> 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<std::uint32_t, char> 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<std::uint32_t>(step.mask_wire));
for (auto [wid, exp] : step.delta)
{
(void)exp;
note(wid);
}
}
}
bundles_ = std::move(saved);
return added + blinds.size();
}
static std::vector<poly_term> drop_one(std::vector<poly_term> 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<poly_term> & 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<poly_term> terms)
{
if (terms.size() < 2)
return emit_terms(std::move(terms));
std::vector<std::vector<poly_term>> best{terms};
std::size_t best_cost = estimate_pieces(best);
auto consider = [&](std::vector<std::vector<poly_term>> seq) {
const std::size_t cost = estimate_pieces(seq);
if (cost < best_cost)
{
best = std::move(seq);
best_cost = cost;
}
};
std::map<std::uint32_t, char> seen;
for (const auto & term : terms)
for (auto id : term.factors)
seen[id] = 1;
std::vector<std::uint32_t> 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<std::uint32_t>(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<std::uint8_t>, std::vector<std::uint32_t>> clusters;
for (auto [wire_id, present] : seen)
{
(void)present;
std::vector<std::uint8_t> 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<poly_term> 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<std::uint32_t, char> 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<std::uint32_t> xs, std::vector<std::uint32_t> 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<int>(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<Ring> 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<Ring> eval_product(const std::vector<std::uint32_t> & 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<Ring>{s0, s1};
}
split<Ring> 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<Ring>{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<bundle_part> & a,
const std::vector<bundle_part> & 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<bundle_part> parts)
{
for (std::size_t i = 0; i < bundles_.size(); ++i)
{
if (same_parts(bundles_[i].parts, parts))
return static_cast<int>(i);
}
bundles_.push_back(bundle{std::move(parts), {}, false});
return static_cast<int>(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<int>(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<poly_step> compile_poly(const std::vector<poly_term> & terms)
{
struct bucket
{
Ring pub{};
std::map<exp_list, Ring> lams;
};
std::map<exp_list, bucket> buckets;
std::vector<poly_step> steps;
for (const auto & term : terms)
{
if (term.factors.size() == 1)
{
poly_step step;
step.scale = term.coeff;
step.value_wire = static_cast<int>(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<std::uint8_t>(ti));
if (e > ti)
delta.emplace_back(id, static_cast<std::uint8_t>(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<bundle_part> 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<int>(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<Ring> 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<std::size_t>(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<std::size_t>(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<std::size_t>(step.bundle)].share;
if (!bundles_[static_cast<std::size_t>(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<Ring>{s0, s1};
}
split<Ring> 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<Ring> 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<Ring> 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<Ring>{s0, s1};
}
};
namespace detail
{
template <typename Ring>
void add_power(typename expr<Ring>::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<unsigned>(power.second) + exp;
if (sum > 16u)
throw std::invalid_argument("beaver exponent is too large");
power.second = static_cast<std::uint8_t>(sum);
return;
}
if (exp > 16u)
throw std::invalid_argument("beaver exponent is too large");
term.powers.emplace_back(id, static_cast<std::uint8_t>(exp));
std::sort(term.powers.begin(), term.powers.end());
}
template <typename Ring>
expr<Ring> merge_terms(expr<Ring> e)
{
using traits = ring_traits<Ring>;
std::map<std::vector<std::pair<std::uint32_t, std::uint8_t>>, 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<Ring>::term term;
term.coeff = coeff;
term.powers = powers;
e.terms.push_back(std::move(term));
}
return e;
}
template <typename Ring>
expr<Ring> mul_exprs(const expr<Ring> & a, const expr<Ring> & b)
{
if (a.sess == nullptr || a.sess != b.sess)
throw std::invalid_argument("beaver factors are from different sessions");
expr<Ring> 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<Ring>::term term = left;
term.coeff = ring_traits<Ring>::mul(left.coeff, right.coeff);
for (auto [id, exp] : right.powers)
add_power<Ring>(term, id, exp);
out.terms.push_back(std::move(term));
}
}
return merge_terms(std::move(out));
}
template <typename Ring>
expr<Ring> add_exprs(expr<Ring> a, const expr<Ring> & 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 <typename Ring>
expr<Ring> scale_expr(expr<Ring> e, Ring coeff)
{
using traits = ring_traits<Ring>;
for (auto & term : e.terms)
term.coeff = traits::mul(term.coeff, coeff);
return merge_terms(std::move(e));
}
template <typename Coeff, typename Ring>
Ring coeff_of(Coeff value)
{
return Ring{value};
}
} // namespace detail
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> wire_expr(wire<Ring> w)
{
if (w.owner() == nullptr)
throw std::invalid_argument("wire is not from a beaver session");
expr<Ring> e;
e.sess = w.owner();
typename expr<Ring>::term term;
term.coeff = ring_traits<Ring>::one();
term.powers.push_back({w.id(), 1});
e.terms.push_back(std::move(term));
return e;
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> horner_expr(wire<Ring> x, std::initializer_list<Ring> coeffs)
{
if (x.owner() == nullptr)
throw std::invalid_argument("wire is not from a beaver session");
expr<Ring> 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<Ring>::zero()))
{
typename expr<Ring>::term term;
term.coeff = coeff;
if (power > 0)
term.powers.push_back({x.id(), static_cast<std::uint8_t>(power)});
e.terms.push_back(std::move(term));
}
++power;
}
return e;
}
/// `coeff * v0 * v1 * ...`, with repeated wires counting as a power.
template <typename Ring, typename... Wires>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> monomial(Ring coeff, wire<Ring> first, Wires... rest)
{
static_assert((std::is_same_v<Wires, wire<Ring>> && ...),
"monomial factors are wires");
expr<Ring> e = wire_expr(first);
((e = e * rest), ...);
return detail::scale_expr(std::move(e), coeff);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> pow(wire<Ring> w, unsigned exp)
{
if (exp == 0)
{
expr<Ring> e = wire_expr(w);
e.terms.clear();
typename expr<Ring>::term term;
term.coeff = ring_traits<Ring>::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<Ring> e = wire_expr(w);
e.terms[0].powers[0].second = static_cast<std::uint8_t>(exp);
return e;
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(wire<Ring> a, wire<Ring> b)
{
return detail::mul_exprs(wire_expr(a), wire_expr(b));
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(expr<Ring> a, wire<Ring> b)
{
return detail::mul_exprs(a, wire_expr(b));
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(wire<Ring> a, expr<Ring> b)
{
return detail::mul_exprs(wire_expr(a), b);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(expr<Ring> a, expr<Ring> b)
{
return detail::mul_exprs(a, b);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(wire<Ring> a, wire<Ring> b)
{
return detail::add_exprs(wire_expr(a), wire_expr(b));
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(expr<Ring> a, wire<Ring> b)
{
return detail::add_exprs(std::move(a), wire_expr(b));
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(wire<Ring> a, expr<Ring> b)
{
return detail::add_exprs(wire_expr(a), b);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(expr<Ring> a, expr<Ring> b)
{
return detail::add_exprs(std::move(a), b);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator-(expr<Ring> e)
{
return detail::scale_expr(std::move(e), ring_traits<Ring>::neg(ring_traits<Ring>::one()));
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator-(wire<Ring> w)
{
return -wire_expr(w);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator-(expr<Ring> a, expr<Ring> b)
{
return std::move(a) + (-std::move(b));
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator-(expr<Ring> a, wire<Ring> b)
{
return std::move(a) + (-b);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator-(wire<Ring> a, expr<Ring> b)
{
return wire_expr(a) + (-std::move(b));
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator-(wire<Ring> a, wire<Ring> b)
{
return wire_expr(a) + (-b);
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(Coeff coeff, wire<Ring> w)
{
return detail::scale_expr(wire_expr(w), detail::coeff_of<Coeff, Ring>(coeff));
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(wire<Ring> w, Coeff coeff)
{
return detail::coeff_of<Coeff, Ring>(coeff) * w;
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(Coeff coeff, expr<Ring> e)
{
return detail::scale_expr(std::move(e), detail::coeff_of<Coeff, Ring>(coeff));
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator*(expr<Ring> e, Coeff coeff)
{
return detail::coeff_of<Coeff, Ring>(coeff) * std::move(e);
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(Coeff coeff, wire<Ring> w)
{
expr<Ring> constant = wire_expr(w);
constant.terms[0].powers.clear();
constant.terms[0].coeff = detail::coeff_of<Coeff, Ring>(coeff);
return detail::add_exprs(std::move(constant), wire_expr(w));
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(wire<Ring> w, Coeff coeff)
{
return coeff + w;
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(Coeff coeff, expr<Ring> e)
{
expr<Ring> constant;
constant.sess = e.sess;
typename expr<Ring>::term term;
term.coeff = detail::coeff_of<Coeff, Ring>(coeff);
constant.terms.push_back(std::move(term));
return detail::add_exprs(std::move(constant), e);
}
template <typename Ring, typename Coeff,
typename = std::enable_if_t<
std::is_constructible_v<Ring, Coeff>
&& !is_beaver_wire<std::decay_t<Coeff>>::value
&& !is_beaver_expr<std::decay_t<Coeff>>::value>>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> operator+(expr<Ring> 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 <std::size_t Arity, typename Ring>
struct fresh_beaver
{
static_assert(Arity >= 2 && Arity <= 8, "fresh beaver arity is 2..8");
std::array<split<Ring>, Arity> in{};
std::array<split<Ring>, (std::size_t{1} << Arity) - 1> subset{};
split<Ring> out{};
};
template <std::size_t Arity, typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
fresh_beaver<Arity, Ring> sample_fresh(Sample && rng)
{
static_assert(Arity >= 2 && Arity <= 8, "fresh beaver arity is 2..8");
session<Ring> s;
std::array<typename session<Ring>::wire, Arity> w{};
for (std::size_t i = 0; i < Arity; ++i)
w[i] = s.input();
expr<Ring> 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<Sample>(rng));
fresh_beaver<Arity, Ring> 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<typename session<Ring>::power> spec;
for (std::size_t i = 0; i < Arity; ++i)
{
if ((mask & (1u << i)) != 0)
spec.push_back(typename session<Ring>::power{w[i], 1u});
}
t.subset[mask - 1] = s.monomial(spec);
}
t.out = s.lambda(out);
return t;
}
template <std::size_t Arity, typename Ring>
HEDLEY_WARN_UNUSED_RESULT
fresh_beaver<Arity, Ring> sample_fresh()
{
return sample_fresh<Arity, Ring>(default_sampler<Ring>{});
}
template <typename Ring>
struct beaver2
{
split<Ring> a{};
split<Ring> b{};
split<Ring> ab{};
split<Ring> out{};
};
template <typename Ring>
struct beaver3
{
split<Ring> a{};
split<Ring> b{};
split<Ring> c{};
split<Ring> ab{};
split<Ring> ac{};
split<Ring> bc{};
split<Ring> abc{};
split<Ring> out{};
};
template <typename Ring>
struct square_beaver
{
split<Ring> x{};
split<Ring> x2{};
split<Ring> out{};
};
template <typename Ring>
struct mul_square_beaver
{
split<Ring> a{};
split<Ring> x{};
split<Ring> x2{};
split<Ring> ax{};
split<Ring> ax2{};
split<Ring> out{};
};
template <typename Ring>
struct dot_beaver
{
std::vector<split<Ring>> x;
std::vector<split<Ring>> y;
split<Ring> cross{};
split<Ring> out{};
};
template <typename Ring>
struct scale_beaver
{
split<Ring> scalar{};
std::vector<split<Ring>> lanes;
std::vector<split<Ring>> cross;
std::vector<split<Ring>> out;
};
template <typename Ring>
struct bit_mul_beaver
{
split<Ring> bit{};
split<Ring> scalar{};
split<Ring> product{};
split<Ring> out{};
};
template <typename Ring>
struct mux_beaver
{
split<Ring> bit{};
split<Ring> when1{};
split<Ring> when0{};
split<Ring> bit_when1{};
split<Ring> bit_when0{};
split<Ring> 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 <typename Ring, typename PRG = dpf::prg::aes128>
HEDLEY_WARN_UNUSED_RESULT
beaver2<Ring> beaver2_at(const oracle<Ring, PRG> & src, std::uint64_t index)
{
using traits = ring_traits<Ring>;
Ring a = src.blind(oracle<Ring>::wire_role(0), index);
Ring b = src.blind(oracle<Ring>::wire_role(1), index);
Ring out = src.blind(oracle<Ring>::wire_role(2), index);
Ring ab = traits::mul(a, b);
return beaver2<Ring>{
src.share(oracle<Ring>::wire_role(0), index, a),
src.share(oracle<Ring>::wire_role(1), index, b),
src.share(oracle<Ring>::mono_role(0), index, ab),
src.share(oracle<Ring>::wire_role(2), index, out)};
}
/// `n` copies starting at `begin`. Each role is one contiguous lane read.
template <typename Ring, typename PRG = dpf::prg::aes128>
void fill_beaver2(const oracle<Ring, PRG> & src, std::uint64_t begin,
beaver2<Ring> * out, std::size_t n)
{
if (n == 0)
return;
if (out == nullptr)
throw std::invalid_argument("beaver2 output is null");
using traits = ring_traits<Ring>;
std::vector<Ring> a(n), b(n), z(n);
std::vector<Ring> ma(n), mb(n), mab(n), mz(n);
src.fill_blinds(oracle<Ring>::wire_role(0), begin, a.data(), n);
src.fill_blinds(oracle<Ring>::wire_role(1), begin, b.data(), n);
src.fill_blinds(oracle<Ring>::wire_role(2), begin, z.data(), n);
src.fill_masks(oracle<Ring>::wire_role(0), begin, ma.data(), n);
src.fill_masks(oracle<Ring>::wire_role(1), begin, mb.data(), n);
src.fill_masks(oracle<Ring>::mono_role(0), begin, mab.data(), n);
src.fill_masks(oracle<Ring>::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<Ring>{
split<Ring>{ma[i], traits::sub(a[i], ma[i])},
split<Ring>{mb[i], traits::sub(b[i], mb[i])},
split<Ring>{mab[i], traits::sub(ab, mab[i])},
split<Ring>{mz[i], traits::sub(z[i], mz[i])}};
}
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
beaver2<Ring> sample_beaver2(Sample && rng)
{
auto t = sample_fresh<2, Ring>(std::forward<Sample>(rng));
// mask 0b11 = 3, index 2, is λ_a λ_b
return beaver2<Ring>{t.in[0], t.in[1], t.subset[2], t.out};
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
beaver2<Ring> sample_beaver2()
{
return sample_beaver2<Ring>(default_sampler<Ring>{});
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
beaver3<Ring> sample_beaver3(Sample && rng)
{
auto t = sample_fresh<3, Ring>(std::forward<Sample>(rng));
return beaver3<Ring>{
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 <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
beaver3<Ring> sample_beaver3()
{
return sample_beaver3<Ring>(default_sampler<Ring>{});
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
square_beaver<Ring> sample_square(Sample && rng)
{
session<Ring> s;
auto x = s.input();
auto out = s(x * x);
s.pin(out);
s.sample(std::forward<Sample>(rng));
return square_beaver<Ring>{
s.lambda(x),
s.monomial({{x, 2u}}),
s.lambda(out)};
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
square_beaver<Ring> sample_square()
{
return sample_square<Ring>(default_sampler<Ring>{});
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
mul_square_beaver<Ring> sample_mul_square(Sample && rng)
{
session<Ring> s;
auto a = s.input();
auto x = s.input();
auto out = s(a * x * x);
s.pin(out);
s.sample(std::forward<Sample>(rng));
return mul_square_beaver<Ring>{
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 <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
mul_square_beaver<Ring> sample_mul_square()
{
return sample_mul_square<Ring>(default_sampler<Ring>{});
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
dot_beaver<Ring> sample_dot(std::size_t n, Sample && rng)
{
if (n == 0)
throw std::invalid_argument("beaver dot length is zero");
session<Ring> s;
std::vector<typename session<Ring>::wire> x;
std::vector<typename session<Ring>::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<Sample>(rng));
dot_beaver<Ring> 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 <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
dot_beaver<Ring> sample_dot(std::size_t n)
{
return sample_dot<Ring>(n, default_sampler<Ring>{});
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
scale_beaver<Ring> sample_scale(std::size_t n, Sample && rng)
{
if (n == 0)
throw std::invalid_argument("beaver scale length is zero");
session<Ring> s;
auto scalar = s.input();
std::vector<typename session<Ring>::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<Sample>(rng));
scale_beaver<Ring> 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 <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
scale_beaver<Ring> sample_scale(std::size_t n)
{
return sample_scale<Ring>(n, default_sampler<Ring>{});
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
bit_mul_beaver<Ring> sample_bit_mul(Sample && rng)
{
session<Ring> s;
auto b = s.bit();
auto x = s.input();
auto out = s.bit_mul(b, x);
s.pin(out);
s.sample(std::forward<Sample>(rng));
return bit_mul_beaver<Ring>{
s.lambda(b),
s.lambda(x),
s.monomial({{b, 1u}, {x, 1u}}),
s.lambda(out)};
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
bit_mul_beaver<Ring> sample_bit_mul()
{
return sample_bit_mul<Ring>(default_sampler<Ring>{});
}
template <typename Ring, typename Sample>
HEDLEY_WARN_UNUSED_RESULT
mux_beaver<Ring> sample_mux(Sample && rng)
{
session<Ring> 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<Sample>(rng));
return mux_beaver<Ring>{
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 <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
mux_beaver<Ring> sample_mux()
{
return sample_mux<Ring>(default_sampler<Ring>{});
}
} // namespace beavers
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_BEAVER_HPP__