2441 lines
74 KiB
C++
2441 lines
74 KiB
C++
/// @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. `sgn * (x*y + pow(x, 2))`
|
||
/// is the sign-corrected form and reuses those powers. A product that
|
||
/// uses an output of an earlier polynomial is a later round.
|
||
///
|
||
/// Doerner–Shelat's per-level AND is a `bit_mul` of a fresh bit and
|
||
/// a fresh block. A wildcard leaf is a `scale` of one scalar by each
|
||
/// lane of the unit vector. Those call sites still draw on their own
|
||
/// pad and `uniform_fill` streams so matched key tapes stay put;
|
||
/// new steps should record a formula here and `sample` it.
|
||
/// The namespace is `dpf::beavers` because `dpf::beaver` is already
|
||
/// the wildcard triple stored on a key.
|
||
///
|
||
/// The session is a dealer: it keeps each full λ so a later round can
|
||
/// multiply blinds without another multiplication protocol. Parties
|
||
/// receive only the additive `split`s.
|
||
///
|
||
/// `oracle` collapses those draws onto `dpf::randomness::lane_table`.
|
||
/// The default PRG is `dpf::prg::aes128`; any PRG with the usual
|
||
/// `eval` interface can be substituted. Each blind role is one lane
|
||
/// (wire id, or `mono_role` / `dot_role` for a derived product), and
|
||
/// the matching share mask is the same index on the tweaked master.
|
||
/// Copy `index` along every lane is one fresh triple. A thousand
|
||
/// copies are a thousand indices, not a thousand lanes. `seed()`
|
||
/// replays every blind and every share mask.
|
||
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
|
||
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
|
||
/// see [LICENSE.md](@ref license) for details.
|
||
|
||
#ifndef LIBDPF_INCLUDE_DPF_BEAVER_HPP__
|
||
#define LIBDPF_INCLUDE_DPF_BEAVER_HPP__
|
||
|
||
#include <algorithm>
|
||
#include <array>
|
||
#include <cstddef>
|
||
#include <cstdint>
|
||
#include <initializer_list>
|
||
#include <map>
|
||
#include <stdexcept>
|
||
#include <type_traits>
|
||
#include <utility>
|
||
#include <vector>
|
||
|
||
#include "hedley/hedley.h"
|
||
|
||
#include "dpf/buffered_prg.hpp"
|
||
#include "dpf/random.hpp"
|
||
#include "dpf/xor_wrapper.hpp"
|
||
|
||
namespace dpf
|
||
{
|
||
namespace beavers
|
||
{
|
||
|
||
/// Ring operations used to build and consume triples.
|
||
/// Specialize for a ring whose multiplicative identity is not `Ring{1}`.
|
||
template <typename Ring>
|
||
struct ring_traits
|
||
{
|
||
static Ring zero() { return Ring{}; }
|
||
|
||
static Ring one() { return Ring{1}; }
|
||
|
||
static Ring add(const Ring & a, const Ring & b) { return static_cast<Ring>(a + b); }
|
||
|
||
static Ring sub(const Ring & a, const Ring & b) { return static_cast<Ring>(a - b); }
|
||
|
||
static Ring mul(const Ring & a, const Ring & b) { return static_cast<Ring>(a * b); }
|
||
|
||
static Ring neg(const Ring & a) { return static_cast<Ring>(-a); }
|
||
|
||
static Ring sample() { return dpf::uniform_sample<Ring>(); }
|
||
};
|
||
|
||
/// Bitwise AND uses the all-ones word as its multiplicative identity.
|
||
template <typename T>
|
||
struct ring_traits<dpf::xor_wrapper<T>>
|
||
{
|
||
using ring = dpf::xor_wrapper<T>;
|
||
|
||
static ring zero() { return ring{}; }
|
||
|
||
static ring one()
|
||
{
|
||
using u = typename ring::value_type;
|
||
return ring{static_cast<u>(~u{0})};
|
||
}
|
||
|
||
static ring add(const ring & a, const ring & b) { return a + b; }
|
||
|
||
static ring sub(const ring & a, const ring & b) { return a - b; }
|
||
|
||
static ring mul(const ring & a, const ring & b) { return a * b; }
|
||
|
||
static ring neg(const ring & a) { return -a; }
|
||
|
||
static ring sample() { return dpf::uniform_sample<ring>(); }
|
||
};
|
||
|
||
template <typename Ring>
|
||
struct default_sampler
|
||
{
|
||
Ring operator()() const { return ring_traits<Ring>::sample(); }
|
||
};
|
||
|
||
/// Additive (2,2) split. `open()` is `p0 + p1`.
|
||
template <typename Ring>
|
||
struct split
|
||
{
|
||
Ring p0{};
|
||
Ring p1{};
|
||
|
||
Ring open() const
|
||
{
|
||
return ring_traits<Ring>::add(p0, p1);
|
||
}
|
||
|
||
dpf::additive_share<Ring, 0> party0() const
|
||
{
|
||
return dpf::additive_share<Ring, 0>::from_raw(p0);
|
||
}
|
||
|
||
dpf::additive_share<Ring, 1> party1() const
|
||
{
|
||
return dpf::additive_share<Ring, 1>::from_raw(p1);
|
||
}
|
||
|
||
friend bool operator==(const split & a, const split & b)
|
||
{
|
||
return a.p0 == b.p0 && a.p1 == b.p1;
|
||
}
|
||
|
||
friend bool operator!=(const split & a, const split & b)
|
||
{
|
||
return !(a == b);
|
||
}
|
||
};
|
||
|
||
template <typename Ring>
|
||
class session;
|
||
|
||
/// A value in a session. Copying a wire copies its id; it does not copy the
|
||
/// blind. The ring argument is on the type so `a * x * x` can build an
|
||
/// expression without naming the session.
|
||
template <typename Ring>
|
||
class wire
|
||
{
|
||
friend class session<Ring>;
|
||
|
||
public:
|
||
constexpr wire() noexcept = default;
|
||
|
||
constexpr std::uint32_t id() const noexcept { return id_; }
|
||
|
||
constexpr session<Ring> * owner() const noexcept { return sess_; }
|
||
|
||
friend bool operator==(wire a, wire b) noexcept
|
||
{
|
||
return a.sess_ == b.sess_ && a.id_ == b.id_;
|
||
}
|
||
|
||
friend bool operator!=(wire a, wire b) noexcept
|
||
{
|
||
return !(a == b);
|
||
}
|
||
|
||
private:
|
||
session<Ring> * sess_ = nullptr;
|
||
std::uint32_t id_ = 0;
|
||
};
|
||
|
||
/// Unevaluated sum of monomials. `*` distributes over `+`. A public
|
||
/// coefficient scales a term. Nothing is sampled until `session::operator()`.
|
||
template <typename Ring>
|
||
struct expr
|
||
{
|
||
struct term
|
||
{
|
||
Ring coeff{};
|
||
/// Positive exponents, sorted by wire id.
|
||
std::vector<std::pair<std::uint32_t, std::uint8_t>> powers;
|
||
};
|
||
|
||
session<Ring> * sess = nullptr;
|
||
std::vector<term> terms;
|
||
};
|
||
|
||
template <typename T>
|
||
struct is_beaver_wire : std::false_type {};
|
||
|
||
template <typename Ring>
|
||
struct is_beaver_wire<wire<Ring>> : std::true_type {};
|
||
|
||
template <typename T>
|
||
struct is_beaver_expr : std::false_type {};
|
||
|
||
template <typename Ring>
|
||
struct is_beaver_expr<expr<Ring>> : std::true_type {};
|
||
|
||
template <typename Ring>
|
||
expr<Ring> wire_expr(wire<Ring> w);
|
||
|
||
template <typename Ring>
|
||
expr<Ring> horner_expr(wire<Ring> x, std::initializer_list<Ring> coeffs);
|
||
|
||
/// One PRG lane per blind role, plus the share-mask stream for that role.
|
||
/// `blind(role, index)` and `share(role, index, value)` do not depend on
|
||
/// call order. Walking `index` forward stays inside a refilled window.
|
||
/// `PRG` defaults to `dpf::prg::aes128`.
|
||
template <typename Ring, typename PRG = dpf::prg::aes128>
|
||
class oracle
|
||
{
|
||
static_assert(std::is_trivially_copyable_v<Ring>,
|
||
"prg oracle lanes require a trivially copyable ring");
|
||
|
||
public:
|
||
using prg_type = PRG;
|
||
using seed_type = typename PRG::block_type;
|
||
using traits = ring_traits<Ring>;
|
||
|
||
/// Monomial roles sit above wire ids. Dot-cross roles sit above those.
|
||
static constexpr std::uint32_t mono_role_base = 0x40000000u;
|
||
static constexpr std::uint32_t dot_role_base = 0x80000000u;
|
||
|
||
static constexpr std::uint32_t wire_role(std::uint32_t id) noexcept { return id; }
|
||
|
||
static constexpr std::uint32_t mono_role(std::uint32_t i) noexcept
|
||
{
|
||
return mono_role_base + i;
|
||
}
|
||
|
||
static constexpr std::uint32_t dot_role(std::uint32_t gate) noexcept
|
||
{
|
||
return dot_role_base + gate;
|
||
}
|
||
|
||
/// Fused within-polynomial λ combinations (Appendix E groupings).
|
||
static constexpr std::uint32_t bundle_role_base = 0xC0000000u;
|
||
|
||
static constexpr std::uint32_t bundle_role(std::uint32_t i) noexcept
|
||
{
|
||
return bundle_role_base + i;
|
||
}
|
||
|
||
explicit oracle(std::size_t window = 256u)
|
||
: lanes_(window)
|
||
{ }
|
||
|
||
explicit oracle(seed_type seed, std::size_t window = 256u)
|
||
: lanes_(std::move(seed), window)
|
||
{ }
|
||
|
||
const seed_type & seed() const noexcept { return lanes_.seed(); }
|
||
|
||
Ring blind(std::uint32_t role, std::uint64_t index) const
|
||
{
|
||
return lanes_.value_at(role, index);
|
||
}
|
||
|
||
Ring mask(std::uint32_t role, std::uint64_t index) const
|
||
{
|
||
return lanes_.mask_at(role, index);
|
||
}
|
||
|
||
/// Additive split of `value`. The mask is the role's mask stream at `index`.
|
||
split<Ring> share(std::uint32_t role, std::uint64_t index, const Ring & value) const
|
||
{
|
||
Ring p0 = mask(role, index);
|
||
return split<Ring>{p0, traits::sub(value, p0)};
|
||
}
|
||
|
||
void fill_blinds(std::uint32_t role, std::uint64_t index, Ring * out, std::size_t n) const
|
||
{
|
||
lanes_.fill_values(role, index, out, n);
|
||
}
|
||
|
||
void fill_masks(std::uint32_t role, std::uint64_t index, Ring * out, std::size_t n) const
|
||
{
|
||
lanes_.fill_masks(role, index, out, n);
|
||
}
|
||
|
||
private:
|
||
dpf::randomness::lane_table<Ring, PRG> lanes_;
|
||
};
|
||
|
||
/// Shares produced for one copy index of a recorded formula.
|
||
template <typename Ring>
|
||
struct prg_material
|
||
{
|
||
std::vector<split<Ring>> lambda;
|
||
std::vector<split<Ring>> monomial;
|
||
/// Fused λ-combinations for polynomial gates, in bundle index order.
|
||
std::vector<split<Ring>> bundles;
|
||
/// Parallel to the session's gates. Empty split when the gate is not a dot.
|
||
std::vector<split<Ring>> dot_cross;
|
||
};
|
||
|
||
/// Dealer session: record formulae, `sample` blinds and monomials, `bind`
|
||
/// input secrets, `evaluate` every round.
|
||
template <typename Ring>
|
||
class session
|
||
{
|
||
static_assert(!std::is_same_v<std::remove_cv_t<Ring>, bool>,
|
||
"beaver bit wires use bit(), not bool arithmetic");
|
||
static_assert(!(std::is_integral_v<Ring> && std::is_signed_v<Ring>),
|
||
"beaver rings are unsigned mod 2^n; use uintN_t or dpf::modint");
|
||
|
||
public:
|
||
using traits = ring_traits<Ring>;
|
||
using exp_list = std::vector<std::pair<std::uint32_t, std::uint8_t>>;
|
||
using wire = ::dpf::beavers::wire<Ring>;
|
||
|
||
/// One factor of a monomial query: `{{x, 2}, {a, 1}}`.
|
||
struct power
|
||
{
|
||
wire base{};
|
||
unsigned exp = 1;
|
||
};
|
||
|
||
session() = default;
|
||
session(const session &) = delete;
|
||
session & operator=(const session &) = delete;
|
||
session(session &&) = delete;
|
||
session & operator=(session &&) = delete;
|
||
|
||
/// Arithmetic input. Its blind is sampled once and then reused.
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire input()
|
||
{
|
||
return emplace_wire(0, true, false);
|
||
}
|
||
|
||
/// Sample this wire's blind even if no recorded formula opens it.
|
||
/// One-shot triples use this for the product wire, so it can be reused.
|
||
void pin(wire w)
|
||
{
|
||
wires_[check(w)].pinned = true;
|
||
}
|
||
|
||
/// 0/1 wire in this ring. `bind` accepts only `zero()` or `one()`
|
||
/// (`1` for integer rings, the all-ones word for `xor_wrapper`).
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire bit()
|
||
{
|
||
return emplace_wire(0, true, true);
|
||
}
|
||
|
||
/// One-round product. Repeated wires share a blind.
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire product(std::initializer_list<wire> factors)
|
||
{
|
||
std::vector<std::uint32_t> ids;
|
||
ids.reserve(factors.size());
|
||
for (auto w : factors)
|
||
ids.push_back(check(w));
|
||
return commit_product(std::move(ids));
|
||
}
|
||
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire product(wire a, wire b)
|
||
{
|
||
return commit_product({check(a), check(b)});
|
||
}
|
||
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire product(wire a, wire b, wire c)
|
||
{
|
||
return commit_product({check(a), check(b), check(c)});
|
||
}
|
||
|
||
/// Record a sum of monomials. Like terms share one blind product.
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire operator()(const expr<Ring> & e)
|
||
{
|
||
if (e.sess != this)
|
||
throw std::invalid_argument("beaver expression is from a different session");
|
||
return commit_poly(e);
|
||
}
|
||
|
||
/// `c[0] + c[1] x + c[2] x^2 + ...` in one round.
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire horner(wire x, std::initializer_list<Ring> coeffs)
|
||
{
|
||
return (*this)(horner_expr(x, coeffs));
|
||
}
|
||
|
||
/// Sign-corrected Horner: `sign * (c[0] + c[1] x + ...)`, still one round
|
||
/// when `sign` and `x` are inputs.
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire horner(wire sign, wire x, std::initializer_list<Ring> coeffs)
|
||
{
|
||
return (*this)(wire_expr(sign) * horner_expr(x, coeffs));
|
||
}
|
||
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire square(wire x)
|
||
{
|
||
return product(x, x);
|
||
}
|
||
|
||
/// One-round `a * x * x` (one blind for `x`).
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire mul_square(wire a, wire x)
|
||
{
|
||
return product({a, x, x});
|
||
}
|
||
|
||
template <typename ContX, typename ContY>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire dot(const ContX & xs, const ContY & ys)
|
||
{
|
||
std::vector<std::uint32_t> x;
|
||
std::vector<std::uint32_t> y;
|
||
for (const auto & w : xs)
|
||
x.push_back(check(w));
|
||
for (const auto & w : ys)
|
||
y.push_back(check(w));
|
||
return commit_dot(std::move(x), std::move(y));
|
||
}
|
||
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire dot(std::initializer_list<wire> xs, std::initializer_list<wire> ys)
|
||
{
|
||
std::vector<std::uint32_t> x;
|
||
std::vector<std::uint32_t> y;
|
||
x.reserve(xs.size());
|
||
y.reserve(ys.size());
|
||
for (auto w : xs)
|
||
x.push_back(check(w));
|
||
for (auto w : ys)
|
||
y.push_back(check(w));
|
||
return commit_dot(std::move(x), std::move(y));
|
||
}
|
||
|
||
/// `z_i = scalar * lanes[i]`, one output wire per lane. The scalar blind
|
||
/// is shared. Each lane gets its own `λ_s λ_i` share.
|
||
template <typename Cont>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
std::vector<wire> scale(wire scalar, const Cont & lanes)
|
||
{
|
||
std::vector<wire> out;
|
||
for (const auto & lane : lanes)
|
||
out.push_back(product(scalar, lane));
|
||
return out;
|
||
}
|
||
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
std::vector<wire> scale(wire scalar, std::initializer_list<wire> lanes)
|
||
{
|
||
std::vector<wire> out;
|
||
out.reserve(lanes.size());
|
||
for (auto lane : lanes)
|
||
out.push_back(product(scalar, lane));
|
||
return out;
|
||
}
|
||
|
||
/// `bit * scalar`. `bit` must come from `bit()`.
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire bit_mul(wire selector, wire scalar)
|
||
{
|
||
auto id = check(selector);
|
||
if (!wires_[id].is_bit)
|
||
throw std::invalid_argument("bit_mul selector must come from bit()");
|
||
return commit_product({id, check(scalar)});
|
||
}
|
||
|
||
/// One-round `selector ? when1 : when0`, i.e. `when0 + selector * (when1 - when0)`.
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
wire mux(wire selector, wire when1, wire when0)
|
||
{
|
||
auto ib = check(selector);
|
||
auto ix = check(when1);
|
||
auto iy = check(when0);
|
||
int round = 1;
|
||
round = std::max(round, wires_[ib].ready_round + 1);
|
||
round = std::max(round, wires_[ix].ready_round + 1);
|
||
round = std::max(round, wires_[iy].ready_round + 1);
|
||
auto out = emplace_wire(round, false, false);
|
||
require_product_monomials({ib, ix});
|
||
require_product_monomials({ib, iy});
|
||
gate g;
|
||
g.kind = gate_kind::mux;
|
||
g.out = out.id_;
|
||
g.lhs = {ib, ix, iy};
|
||
gates_.push_back(std::move(g));
|
||
wires_[out.id_].gate = static_cast<int>(gates_.size() - 1);
|
||
return out;
|
||
}
|
||
|
||
/// Sample every missing wire blind and every missing monomial.
|
||
/// Blinds already sampled are left alone.
|
||
template <typename Sample>
|
||
void sample(Sample && sampler)
|
||
{
|
||
auto & rng = sampler;
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
auto & w = wires_[id];
|
||
if (w.lambda_ready || !needs_blind(id))
|
||
continue;
|
||
w.lambda_full = rng();
|
||
w.lambda = share_of(w.lambda_full, rng);
|
||
w.lambda_ready = true;
|
||
}
|
||
for (auto & m : monos_)
|
||
{
|
||
if (m.ready)
|
||
continue;
|
||
Ring full = traits::one();
|
||
for (auto [id, e] : m.key)
|
||
{
|
||
if (!wires_[id].lambda_ready)
|
||
throw std::logic_error("beaver blind is missing");
|
||
full = traits::mul(full, pow_small(wires_[id].lambda_full, e));
|
||
}
|
||
m.share = share_of(full, rng);
|
||
m.ready = true;
|
||
}
|
||
for (auto & b : bundles_)
|
||
{
|
||
if (b.ready)
|
||
continue;
|
||
b.share = share_of(bundle_value(b), rng);
|
||
b.ready = true;
|
||
}
|
||
for (auto & g : gates_)
|
||
{
|
||
if (g.kind != gate_kind::dot || g.cross_ready)
|
||
continue;
|
||
Ring sum = traits::zero();
|
||
for (std::size_t i = 0; i < g.lhs.size(); ++i)
|
||
{
|
||
sum = traits::add(sum, traits::mul(
|
||
wires_[g.lhs[i]].lambda_full,
|
||
wires_[g.rhs[i]].lambda_full));
|
||
}
|
||
g.cross = share_of(sum, rng);
|
||
g.cross_ready = true;
|
||
}
|
||
}
|
||
|
||
void sample()
|
||
{
|
||
sample(default_sampler<Ring>{});
|
||
}
|
||
|
||
/// Install missing blinds and product shares from copy `index` of `src`.
|
||
/// Already-sampled wires keep their λ. New product shares are built from
|
||
/// those stored blinds, then split with the oracle's share lane.
|
||
template <typename PRG>
|
||
void sample_from(const oracle<Ring, PRG> & src, std::uint64_t index = 0)
|
||
{
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
auto & w = wires_[id];
|
||
if (w.lambda_ready || !needs_blind(id))
|
||
continue;
|
||
w.lambda_full = src.blind(oracle<Ring>::wire_role(id), index);
|
||
w.lambda = src.share(oracle<Ring>::wire_role(id), index, w.lambda_full);
|
||
w.lambda_ready = true;
|
||
}
|
||
for (std::uint32_t i = 0; i < monos_.size(); ++i)
|
||
{
|
||
auto & m = monos_[i];
|
||
if (m.ready)
|
||
continue;
|
||
Ring full = mono_full(m.key);
|
||
m.share = src.share(oracle<Ring>::mono_role(i), index, full);
|
||
m.ready = true;
|
||
}
|
||
for (std::uint32_t i = 0; i < bundles_.size(); ++i)
|
||
{
|
||
auto & b = bundles_[i];
|
||
if (b.ready)
|
||
continue;
|
||
b.share = src.share(oracle<Ring>::bundle_role(i), index, bundle_value(b));
|
||
b.ready = true;
|
||
}
|
||
for (std::uint32_t g = 0; g < gates_.size(); ++g)
|
||
{
|
||
auto & gate = gates_[g];
|
||
if (gate.kind != gate_kind::dot || gate.cross_ready)
|
||
continue;
|
||
gate.cross = src.share(oracle<Ring>::dot_role(g), index, dot_full(gate));
|
||
gate.cross_ready = true;
|
||
}
|
||
}
|
||
|
||
/// Every wire, monomial, and dot cross of this formula at copy `index`.
|
||
/// Does not change the session. Copies are independent lanes samples, so
|
||
/// `material_at(src, 5)` does not depend on having asked for 0..4.
|
||
template <typename PRG>
|
||
prg_material<Ring> material_at(const oracle<Ring, PRG> & src, std::uint64_t index) const
|
||
{
|
||
prg_material<Ring> out;
|
||
std::vector<Ring> full(wires_.size());
|
||
out.lambda.resize(wires_.size());
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
full[id] = src.blind(oracle<Ring>::wire_role(id), index);
|
||
out.lambda[id] = src.share(oracle<Ring>::wire_role(id), index, full[id]);
|
||
}
|
||
out.monomial.resize(monos_.size());
|
||
for (std::uint32_t i = 0; i < monos_.size(); ++i)
|
||
{
|
||
Ring prod = traits::one();
|
||
for (auto [id, e] : monos_[i].key)
|
||
prod = traits::mul(prod, pow_small(full[id], e));
|
||
out.monomial[i] = src.share(oracle<Ring>::mono_role(i), index, prod);
|
||
}
|
||
out.bundles.resize(bundles_.size());
|
||
for (std::uint32_t i = 0; i < bundles_.size(); ++i)
|
||
{
|
||
Ring value = traits::zero();
|
||
for (const auto & part : bundles_[i].parts)
|
||
{
|
||
Ring m = traits::one();
|
||
for (auto [id, e] : part.lam)
|
||
m = traits::mul(m, pow_small(full[id], e));
|
||
value = traits::add(value, traits::mul(part.coeff, m));
|
||
}
|
||
out.bundles[i] = src.share(oracle<Ring>::bundle_role(i), index, value);
|
||
}
|
||
out.dot_cross.resize(gates_.size());
|
||
for (std::uint32_t g = 0; g < gates_.size(); ++g)
|
||
{
|
||
if (gates_[g].kind != gate_kind::dot)
|
||
continue;
|
||
Ring sum = traits::zero();
|
||
for (std::size_t i = 0; i < gates_[g].lhs.size(); ++i)
|
||
{
|
||
sum = traits::add(sum, traits::mul(
|
||
full[gates_[g].lhs[i]], full[gates_[g].rhs[i]]));
|
||
}
|
||
out.dot_cross[g] = src.share(oracle<Ring>::dot_role(g), index, sum);
|
||
}
|
||
return out;
|
||
}
|
||
|
||
/// Split `secret` into fresh additive shares and bind them to an input.
|
||
template <typename Sample>
|
||
void bind(wire w, Ring secret, Sample && sampler)
|
||
{
|
||
auto id = check(w);
|
||
if (wires_[id].ready_round != 0)
|
||
throw std::invalid_argument("only input wires can be bound");
|
||
if (wires_[id].value_ready)
|
||
throw std::logic_error("input wire is already bound");
|
||
if (wires_[id].is_bit && !is_bit_secret(secret))
|
||
throw std::invalid_argument("bit wire must be 0 or 1");
|
||
auto & rng = sampler;
|
||
wires_[id].value = share_of(secret, rng);
|
||
wires_[id].value_ready = true;
|
||
}
|
||
|
||
void bind(wire w, Ring secret)
|
||
{
|
||
bind(w, secret, default_sampler<Ring>{});
|
||
}
|
||
|
||
/// Bind shares the caller already holds. Their sum is the secret.
|
||
void bind_shares(wire w, Ring p0, Ring p1)
|
||
{
|
||
auto id = check(w);
|
||
if (wires_[id].ready_round != 0)
|
||
throw std::invalid_argument("only input wires can be bound");
|
||
if (wires_[id].value_ready)
|
||
throw std::logic_error("input wire is already bound");
|
||
if (wires_[id].is_bit && !is_bit_secret(traits::add(p0, p1)))
|
||
throw std::invalid_argument("bit wire must be 0 or 1");
|
||
wires_[id].value = split<Ring>{p0, p1};
|
||
wires_[id].value_ready = true;
|
||
}
|
||
|
||
/// Open every ready round. Inputs used by round-1 gates are opened
|
||
/// together; a gate output is a later round's input and keeps the blind
|
||
/// chosen in `sample`.
|
||
void evaluate()
|
||
{
|
||
int max_round = 0;
|
||
for (const auto & w : wires_)
|
||
max_round = std::max(max_round, w.ready_round);
|
||
for (int r = 1; r <= max_round; ++r)
|
||
{
|
||
for (const auto & g : gates_)
|
||
{
|
||
if (wires_[g.out].ready_round != r)
|
||
continue;
|
||
if (wires_[g.out].value_ready)
|
||
continue;
|
||
split<Ring> val;
|
||
if (g.kind == gate_kind::product)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta(id);
|
||
val = eval_product(g.lhs);
|
||
}
|
||
else if (g.kind == gate_kind::dot)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta(id);
|
||
for (auto id : g.rhs)
|
||
ensure_delta(id);
|
||
val = eval_dot(g);
|
||
}
|
||
else if (g.kind == gate_kind::poly)
|
||
{
|
||
for (const auto & term : g.terms)
|
||
for (auto id : term.factors)
|
||
ensure_delta(id);
|
||
val = eval_poly(g);
|
||
}
|
||
else
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta(id);
|
||
val = eval_mux(g);
|
||
}
|
||
wires_[g.out].value = val;
|
||
wires_[g.out].value_ready = true;
|
||
if (wires_[g.out].lambda_ready)
|
||
ensure_delta(g.out);
|
||
}
|
||
}
|
||
}
|
||
|
||
split<Ring> lambda(wire w) const
|
||
{
|
||
auto id = check(w);
|
||
if (!wires_[id].lambda_ready)
|
||
throw std::logic_error("call sample() before reading a beaver blind");
|
||
return wires_[id].lambda;
|
||
}
|
||
|
||
/// Share of `Π λ_i^{e_i}`. A lone `λ_w` is the wire blind itself.
|
||
split<Ring> monomial(const std::vector<power> & spec) const
|
||
{
|
||
std::vector<std::pair<std::uint32_t, unsigned>> raw;
|
||
raw.reserve(spec.size());
|
||
for (auto p : spec)
|
||
{
|
||
if (p.exp == 0)
|
||
throw std::invalid_argument("beaver power is zero");
|
||
raw.emplace_back(check(p.base), p.exp);
|
||
}
|
||
return term_share(merge_exponents(raw));
|
||
}
|
||
|
||
split<Ring> monomial(std::initializer_list<power> spec) const
|
||
{
|
||
return monomial(std::vector<power>(spec));
|
||
}
|
||
|
||
split<Ring> value(wire w) const
|
||
{
|
||
auto id = check(w);
|
||
if (!wires_[id].value_ready)
|
||
throw std::logic_error("beaver wire has no value yet");
|
||
return wires_[id].value;
|
||
}
|
||
|
||
Ring open(wire w) const
|
||
{
|
||
auto v = value(w);
|
||
return traits::add(v.p0, v.p1);
|
||
}
|
||
|
||
/// Public ABY2.0 mask δ = x + λ, after `evaluate` has opened the wire.
|
||
Ring delta(wire w) const
|
||
{
|
||
auto id = check(w);
|
||
if (!wires_[id].delta_ready)
|
||
throw std::logic_error("beaver wire has not been opened");
|
||
return wires_[id].delta;
|
||
}
|
||
|
||
split<Ring> dot_cross(wire w) const
|
||
{
|
||
auto id = check(w);
|
||
int g = wires_[id].gate;
|
||
if (g < 0 || gates_[static_cast<std::size_t>(g)].kind != gate_kind::dot)
|
||
throw std::invalid_argument("wire is not a dot output");
|
||
if (!gates_[static_cast<std::size_t>(g)].cross_ready)
|
||
throw std::logic_error("call sample() before reading a dot cross term");
|
||
return gates_[static_cast<std::size_t>(g)].cross;
|
||
}
|
||
|
||
int round_of(wire w) const
|
||
{
|
||
return wires_[check(w)].ready_round;
|
||
}
|
||
|
||
std::size_t wire_count() const noexcept { return wires_.size(); }
|
||
|
||
/// Product shares beyond the per-wire blinds: subset monomials from
|
||
/// `product` gates, plus one fused bundle per public-δ class in a
|
||
/// polynomial (Appendix E). A lone mask is not counted.
|
||
std::size_t monomial_count() const noexcept
|
||
{
|
||
return monos_.size() + bundles_.size();
|
||
}
|
||
|
||
/// Wire blinds that the recorded formulae actually open, plus product shares.
|
||
std::size_t preprocessing_count() const
|
||
{
|
||
std::size_t n = monomial_count();
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
if (needs_blind(id))
|
||
++n;
|
||
}
|
||
return n;
|
||
}
|
||
|
||
private:
|
||
enum class gate_kind : unsigned char { product, dot, mux, poly };
|
||
|
||
struct poly_term
|
||
{
|
||
Ring coeff{};
|
||
std::vector<std::uint32_t> factors;
|
||
};
|
||
|
||
/// One λ-monomial in a fused preprocessing share.
|
||
struct bundle_part
|
||
{
|
||
Ring coeff{};
|
||
exp_list lam;
|
||
};
|
||
|
||
struct bundle
|
||
{
|
||
std::vector<bundle_part> parts;
|
||
split<Ring> share{};
|
||
bool ready = false;
|
||
};
|
||
|
||
/// Online: `public(δ) * scale * share`, where share is a wire mask,
|
||
/// a raw monomial, or a fused sum of monomials.
|
||
struct poly_step
|
||
{
|
||
exp_list delta;
|
||
Ring scale{};
|
||
int bundle = -1;
|
||
int mask_wire = -1;
|
||
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_;
|
||
for (const auto & piece : pieces)
|
||
(void)compile_poly(piece);
|
||
std::size_t added = bundles_.size() - saved.size();
|
||
bundles_ = std::move(saved);
|
||
std::map<std::uint32_t, char> blinds;
|
||
for (const auto & piece : pieces)
|
||
{
|
||
for (const auto & term : piece)
|
||
{
|
||
for (auto id : term.factors)
|
||
{
|
||
if (id >= already.size() || already[id] == 0)
|
||
blinds[id] = 1;
|
||
}
|
||
}
|
||
}
|
||
return added + blinds.size();
|
||
}
|
||
|
||
wire schedule_terms(std::vector<poly_term> terms)
|
||
{
|
||
std::vector<std::vector<poly_term>> best{terms};
|
||
std::size_t best_cost = estimate_pieces(best);
|
||
|
||
std::map<std::uint32_t, char> seen;
|
||
for (const auto & term : terms)
|
||
for (auto id : term.factors)
|
||
seen[id] = 1;
|
||
for (auto [wire_id, _] : seen)
|
||
{
|
||
bool common = true;
|
||
for (const auto & term : terms)
|
||
{
|
||
if (factor_exp(term, wire_id) == 0)
|
||
common = false;
|
||
}
|
||
if (!common)
|
||
continue;
|
||
std::vector<poly_term> quot = terms;
|
||
for (auto & term : quot)
|
||
{
|
||
auto it = std::find(term.factors.begin(), term.factors.end(), wire_id);
|
||
if (it != term.factors.end())
|
||
term.factors.erase(it);
|
||
}
|
||
const std::uint32_t mid = static_cast<std::uint32_t>(wires_.size());
|
||
poly_term mul;
|
||
mul.coeff = traits::one();
|
||
mul.factors = {mid, wire_id};
|
||
std::vector<std::vector<poly_term>> seq{std::move(quot), {std::move(mul)}};
|
||
std::size_t cost = estimate_pieces(seq);
|
||
if (cost < best_cost)
|
||
{
|
||
best = std::move(seq);
|
||
best_cost = cost;
|
||
}
|
||
}
|
||
|
||
std::map<std::vector<std::uint8_t>, std::vector<std::uint32_t>> clusters;
|
||
for (auto [wire_id, _] : seen)
|
||
{
|
||
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[shape].push_back(wire_id);
|
||
}
|
||
for (auto & [shape, group] : clusters)
|
||
{
|
||
if (group.size() < 2)
|
||
continue;
|
||
const std::uint32_t mid = static_cast<std::uint32_t>(wires_.size());
|
||
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(mid);
|
||
rewritten.push_back(std::move(next));
|
||
}
|
||
std::vector<std::vector<poly_term>> seq{{std::move(prod)}, std::move(rewritten)};
|
||
std::size_t cost = estimate_pieces(seq);
|
||
if (cost < best_cost)
|
||
{
|
||
best = std::move(seq);
|
||
best_cost = cost;
|
||
}
|
||
}
|
||
|
||
if (best.size() != 1 && terms.size() == 1 && terms[0].factors.size() <= 4)
|
||
{
|
||
std::fprintf(stderr, "split factors=%zu pieces=%zu cost=%zu flat_factors=",
|
||
terms[0].factors.size(), best.size(), best_cost);
|
||
for (auto f : terms[0].factors)
|
||
std::fprintf(stderr, "%u ", f);
|
||
std::fprintf(stderr, "\n");
|
||
}
|
||
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 (key " + std::to_string(key.size())
|
||
+ " bundles " + std::to_string(bundles_.size())
|
||
+ " gates " + std::to_string(gates_.size()) + ")");
|
||
}
|
||
|
||
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;
|
||
for (const auto & term : terms)
|
||
{
|
||
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);
|
||
});
|
||
}
|
||
|
||
std::vector<poly_step> steps;
|
||
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)
|
||
{
|
||
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__
|