libdpf/include/dpf/beaver.hpp

2527 lines
77 KiB
C++
Raw Normal View History

/// @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:
HEDLEY_NO_THROW
constexpr wire() noexcept = default;
HEDLEY_NO_THROW
constexpr std::uint32_t id() const noexcept { return id_; }
HEDLEY_NO_THROW
constexpr session<Ring> * owner() const noexcept { return sess_; }
HEDLEY_NO_THROW
friend bool operator==(wire a, wire b) noexcept
{
return a.sess_ == b.sess_ && a.id_ == b.id_;
}
HEDLEY_NO_THROW
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;
HEDLEY_NO_THROW
static constexpr std::uint32_t wire_role(std::uint32_t id) noexcept { return id; }
HEDLEY_NO_THROW
static constexpr std::uint32_t mono_role(std::uint32_t i) noexcept
{
return mono_role_base + i;
}
HEDLEY_NO_THROW
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;
HEDLEY_NO_THROW
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)
{ }
HEDLEY_NO_THROW
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;
}
HEDLEY_NO_THROW
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.
HEDLEY_NO_THROW
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__