libdpf/include/dpf/beaver.hpp

2671 lines
83 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. An inner product is that
/// sum: `dot({x0,x1}, {y0,y1})` is `x0*y0 + x1*y1`, and the pair
/// products share one preprocessing value. 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
{
/// @brief Ring operations used to build and consume triples.
/// @details Specialize for a ring whose multiplicative identity is not `Ring{1}`.
/// @tparam Ring payload ring
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>(); }
};
/// @brief Bitwise AND uses the all-ones word as its multiplicative identity.
/// @tparam T value type
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(); }
};
/// @brief Additive (2,2) split. `open()` is `p0 + p1`.
/// @tparam Ring payload ring
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;
/// @brief 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.
/// @tparam Ring payload ring
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;
};
/// @brief Unevaluated sum of monomials. `*` distributes over `+`. A public
/// coefficient scales a term. Nothing is sampled until `session::operator()`.
/// @tparam Ring payload ring
template <typename Ring>
struct expr
{
struct term
{
Ring coeff{};
/// @brief 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);
namespace detail
{
template <typename Ring, typename ContX, typename ContY>
expr<Ring> dot_expr(const ContX & xs, const ContY & ys);
} // namespace detail
/// @brief One PRG lane per blind role, plus the share-mask stream for that role.
/// @details `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`.
/// @tparam Ring payload ring
/// @tparam PRG pseudorandom generator
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>;
/// @brief 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;
}
/// @brief 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);
}
/// @brief Additive split of `value`. The mask is the role's mask stream at `index`.
/// @param role the `role`
/// @param index the index
/// @param value the value to convert or store
/// @return Additive split of `value`
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_;
};
/// @brief Shares produced for one copy index of a recorded formula.
/// @tparam Ring payload ring
template <typename Ring>
struct prg_material
{
std::vector<split<Ring>> lambda;
std::vector<split<Ring>> monomial;
/// @brief Fused λ-combinations for polynomial gates, in bundle index order.
std::vector<split<Ring>> bundles;
/// @brief Parallel to the session's gates. Empty split when the gate is not a dot.
std::vector<split<Ring>> dot_cross;
};
/// @brief Dealer session: record formulae, `sample` blinds and monomials, `bind`
/// input secrets, `evaluate` every round.
/// @tparam Ring payload ring
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>;
/// @brief 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;
/// @brief Arithmetic input. Its blind is sampled once and then reused.
/// @return Arithmetic input
HEDLEY_WARN_UNUSED_RESULT
wire input()
{
return emplace_wire(0, true, false);
}
/// @brief Sample this wire's blind even if no recorded formula opens it.
/// @details One-shot triples use this for the product wire, so it can be reused.
/// @param w the `w`
void pin(wire w)
{
wires_[check(w)].pinned = true;
}
/// @brief 0/1 wire in this ring. `bind` accepts only `zero()` or `one()`
/// (`1` for integer rings, the all-ones word for `xor_wrapper`).
/// @return 0/1 wire in this ring
HEDLEY_WARN_UNUSED_RESULT
wire bit()
{
return emplace_wire(0, true, true);
}
/// @brief One-round product. Repeated wires share a blind.
/// @param factors the `factors`
/// @return One-round product
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)});
}
/// @brief Record a sum of monomials. Like terms share one blind product.
/// @param e the `e`
/// @return Record a sum of monomials
/// @throws std::invalid_argument if `beaver expression is from a different session`
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);
}
/// @brief `c[0] + c[1] x + c[2] x^2 + ...` in one round.
/// @param x the `x`
/// @param coeffs the public coefficients
/// @return `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));
}
/// @brief Sign-corrected Horner: `sign * (c[0] + c[1] x + ...)`, still one round
/// when `sign` and `x` are inputs.
/// @param sign the sign bit or sign value
/// @param x the `x`
/// @param coeffs the public coefficients
/// @return 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);
}
/// @brief One-round `a * x * x` (one blind for `x`).
/// @param a the `a`
/// @param x the `x`
/// @return 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)
{
return finish_dot(detail::dot_expr<Ring>(xs, ys));
}
HEDLEY_WARN_UNUSED_RESULT
wire dot(std::initializer_list<wire> xs, std::initializer_list<wire> ys)
{
return finish_dot(detail::dot_expr<Ring>(xs, ys));
}
/// @brief `z_i = scalar * lanes[i]`, one output wire per lane. The scalar blind
/// is shared. Each lane gets its own `λ_s λ_i` share.
/// @tparam Cont cont
/// @param scalar the `scalar`
/// @param lanes the lane values
/// @return `z_i = scalar * lanes[i]`, one output wire per lane
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;
}
/// @brief `bit * scalar`. `bit` must come from `bit()`.
/// @param selector the `selector`
/// @param scalar the `scalar`
/// @return `bit * scalar`
/// @throws std::invalid_argument if `bit_mul selector 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)});
}
/// @brief One-round `selector ? when1 : when0`, i.e. `when0 + selector * (when1 - when0)`.
/// @param selector the `selector`
/// @param when1 the `when1`
/// @param when0 the `when0`
/// @return One-round `selector ? when1 : when0`, i.e
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;
}
/// @brief Sample every missing wire blind and every missing monomial.
/// @details Blinds already sampled are left alone.
/// @tparam Sample sample
/// @param sampler the randomness sampler
/// @throws std::logic_error if `beaver blind is missing`
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>{});
}
/// @brief Install missing blinds and product shares from copy `index` of `src`.
/// @details Already-sampled wires keep their λ. New product shares are built from
/// those stored blinds, then split with the oracle's share lane.
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param index the index
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;
}
}
/// @brief Every wire, monomial, and dot cross of this formula at copy `index`.
/// @details Does not change the session. Copies are independent lanes samples, so
/// `material_at(src, 5)` does not depend on having asked for 0..4.
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param index the index
/// @return Every wire, monomial, and dot cross of this formula at copy `index`
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;
}
/// @brief Split `secret` into fresh additive shares and bind them to an input.
/// @tparam Sample sample
/// @param w the `w`
/// @param secret the secret value
/// @param sampler the randomness sampler
/// @throws std::invalid_argument if `only input wires can be bound`
/// @throws std::logic_error if `input wire is already bound`
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>{});
}
/// @brief Bind shares the caller already holds. Their sum is the secret.
/// @param w the `w`
/// @param p0 the `p0`
/// @param p1 the `p1`
/// @throws std::invalid_argument if `only input wires can be bound`
/// @throws std::logic_error if `input wire is already bound`
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;
}
/// @brief 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`.
/// @throws std::logic_error if `beaver wire is not ready to open`
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;
}
/// @brief Share of `Π λ_i^{e_i}`. A lone `λ_w` is the wire blind itself.
/// @param spec the specification
/// @return Share of `Π λ_i^{e_i}`
/// @throws std::invalid_argument if `beaver power is zero`
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);
}
/// @brief Public ABY2.0 mask δ = x + λ, after `evaluate` has opened the wire.
/// @param w the `w`
/// @return Public ABY2.0 mask δ = x + λ, after `evaluate` has opened the wire
/// @throws std::logic_error if `beaver wire has not been opened`
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);
if (!wires_[id].dot_output)
throw std::invalid_argument("wire is not a dot output");
int g = wires_[id].gate;
if (g < 0)
throw std::invalid_argument("wire is not a dot output");
const auto & gate = gates_[static_cast<std::size_t>(g)];
if (gate.kind == gate_kind::dot)
{
if (!gate.cross_ready)
throw std::logic_error("call sample() before reading a dot cross term");
return gate.cross;
}
Ring s0 = traits::zero();
Ring s1 = traits::zero();
bool found = false;
for (const auto & step : gate.steps)
{
if (step.bundle < 0 || !step.delta.empty())
continue;
const auto & bundle = bundles_[static_cast<std::size_t>(step.bundle)];
if (!bundle.ready)
throw std::logic_error("call sample() before reading a dot cross term");
s0 = traits::add(s0, traits::mul(step.scale, bundle.share.p0));
s1 = traits::add(s1, traits::mul(step.scale, bundle.share.p1));
found = true;
}
if (!found)
throw std::logic_error("call sample() before reading a dot cross term");
return split<Ring>{s0, s1};
}
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(); }
/// @brief 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.
/// @return Product shares beyond the per-wire blinds: subset monomials from `product` gates,
/// plus one fused bundle per public-δ class in a polynomial (Appendix E)
HEDLEY_NO_THROW
std::size_t monomial_count() const noexcept
{
return monos_.size() + bundles_.size();
}
/// @brief Wire blinds that the recorded formulae actually open, plus product shares.
/// @return 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;
};
/// @brief 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;
};
/// @brief 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 dot_output = 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 finish_dot(const expr<Ring> & e)
{
auto out = (*this)(e);
wires_[out.id_].dot_output = true;
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;
}
/// @brief Group λ-monomials that share a public δ monomial into one share.
/// @details A bucket that is only `c · λ_i` reuses the wire blind.
/// @param terms the polynomial terms
/// @return Group λ-monomials that share a public δ monomial into one share
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;
}
Ring scale = parts[0].coeff;
for (const auto & part : parts)
{
if (!(part.coeff == scale))
scale = traits::one();
}
if (!(scale == traits::one()))
{
for (auto & part : parts)
part.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));
}
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};
}
template <typename Ring, typename ContX, typename ContY>
expr<Ring> dot_expr(const ContX & xs, const ContY & ys)
{
std::vector<wire<Ring>> x;
std::vector<wire<Ring>> y;
for (const auto & w : xs)
x.push_back(w);
for (const auto & w : ys)
y.push_back(w);
if (x.empty() || x.size() != y.size())
throw std::invalid_argument(
"beaver dot operands must have the same non-zero length");
expr<Ring> acc = wire_expr(x[0]) * wire_expr(y[0]);
for (std::size_t i = 1; i < x.size(); ++i)
acc = add_exprs(std::move(acc), wire_expr(x[i]) * wire_expr(y[i]));
return acc;
}
} // namespace detail
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> dot(std::initializer_list<wire<Ring>> xs, std::initializer_list<wire<Ring>> ys)
{
return detail::dot_expr<Ring>(xs, ys);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> dot(const std::vector<wire<Ring>> & xs, const std::vector<wire<Ring>> & ys)
{
return detail::dot_expr<Ring>(xs, ys);
}
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;
}
/// @brief `coeff * v0 * v1 * ...`, with repeated wires counting as a power.
/// @tparam Ring payload ring
/// @tparam Wires wires
/// @param coeff the public coefficient
/// @param first the first element of the range
/// @param rest the remaining arguments
/// @return `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.
// ---------------------------------------------------------------------------
/// @brief `subset[mask - 1]` is `Π λ_i` over the bits set in `mask` (bit i selects
/// `in[i]`). Singleton masks are the wire blinds themselves.
/// @tparam Arity arity
/// @tparam Ring payload ring
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{};
};
/// @brief One fresh Beaver pair from copy `index` of an oracle.
/// @details Roles match a session that records `input, input, product`: wires 0 and 1,
/// the product wire, and monomial 0.
/// @tparam Ring payload ring
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param index the index
/// @return One fresh Beaver pair from copy `index` of an oracle
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)};
}
/// @brief `n` copies starting at `begin`. Each role is one contiguous lane read.
/// @tparam Ring payload ring
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param begin the iterator to the first query
/// @param out the output buffer
/// @param n the `n`
/// @throws std::invalid_argument if `beaver2 output is null`
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__