Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
5197 lines
173 KiB
C++
5197 lines
173 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. 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 draw from `sample_bit_mul`
|
||
/// / `sample_scale`. A callable sampler and an `oracle` are the same
|
||
/// construction: `sample` walks the callable, `sample_from` reads the
|
||
/// oracle, and the one-shot helpers accept either. 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.
|
||
///
|
||
/// Optional Shark/SPDZ-style IT-MACs: `session::set_mac_key` before
|
||
/// `sample`/`bind` tags every λ, monomial, and value under a global
|
||
/// `Δ`. Classic triples use `sample_auth_beaver2` /
|
||
/// `authenticate_beaver2`. Openings of δ = x + λ carry tag shares and
|
||
/// check with `verify_delta` / `verify_auth_opening`.
|
||
///
|
||
/// `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.
|
||
/// IT-MAC tag masks are a third stream from the same seed, so an
|
||
/// authenticated `sample_from` does not fall back to the system RNG.
|
||
/// Copy `index` along every lane is one fresh triple. A thousand
|
||
/// copies are a thousand indices, not a thousand lanes. `seed()`
|
||
/// replays every blind, every share mask, and every tag mask.
|
||
/// Constant-round arithmetic is a different object. Free addition, free
|
||
/// scaling by a public constant, and a unary projection are Ball, Malkin,
|
||
/// and Rosulek, CCS 2016 (ePrint 2016/969), in `dpf/arith_garble.hpp`.
|
||
/// A session does not garble those gadgets: it still opens δ once per wire.
|
||
/// @note A session follows Patra, Schneider, Suresh, and Yalame, USENIX Security 2021 (full version ePrint 2020/1225): one public reconstruction per newly opened wire. `sample_beaver2` follows Donald Beaver, "Efficient Multiparty Protocols Using Circuit Randomization," CRYPTO 1991 (LNCS 576, pp. 420–432), which reconstructs both masked factors. Ball, Malkin, and Rosulek, "Garbling Gadgets for Boolean and Arithmetic Circuits," CCS 2016 (ePrint 2016/969), is the constant-round projection gadget in `arith_garble.hpp`, not this session.
|
||
/// @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 <cstring>
|
||
#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/leaf_arithmetic.hpp"
|
||
#include "dpf/packed_lane.hpp"
|
||
#include "dpf/random.hpp"
|
||
#include "dpf/verifiable.hpp"
|
||
#include "dpf/xor_wrapper.hpp"
|
||
#include "grotto/fixedpoint.hpp"
|
||
|
||
namespace dpf
|
||
{
|
||
|
||
/// @brief Integers mod `2^w`, `w` in `1 .. 128`.
|
||
/// @details `w` is the thread-local width set by `beavers::mod2k_width_scope`
|
||
/// for the duration of a sample. Arithmetic masks to that width.
|
||
struct mod2k
|
||
{
|
||
simde_uint128 raw{};
|
||
|
||
friend bool operator==(mod2k a, mod2k b) noexcept
|
||
{
|
||
return a.raw == b.raw;
|
||
}
|
||
|
||
friend bool operator!=(mod2k a, mod2k b) noexcept
|
||
{
|
||
return !(a == b);
|
||
}
|
||
};
|
||
|
||
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 Fixed-point product widens the fractional field; cast back for the ring.
|
||
/// @tparam FractionalBits fractional bits
|
||
/// @tparam IntegralType backend integer
|
||
template <unsigned FractionalBits, typename IntegralType>
|
||
struct ring_traits<grotto::fixedpoint<FractionalBits, IntegralType>>
|
||
{
|
||
using Ring = grotto::fixedpoint<FractionalBits, IntegralType>;
|
||
|
||
static Ring zero() { return Ring{}; }
|
||
|
||
static Ring one() { return Ring{1}; }
|
||
|
||
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 grotto::precision_cast<FractionalBits>(a * b);
|
||
}
|
||
|
||
static Ring neg(const Ring & a) { return -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>(); }
|
||
};
|
||
|
||
/// @brief IEEE `float` leaf group is XOR/AND of the bit pattern, matching
|
||
/// `multiply_t<float>` / `add_t<float, NodeT>`.
|
||
template <>
|
||
struct ring_traits<float>
|
||
{
|
||
using ring = float;
|
||
|
||
static ring zero() { return 0.f; }
|
||
|
||
static ring one() { return dpf::leaf_group_one<float>(); }
|
||
|
||
static ring add(const ring & a, const ring & b)
|
||
{
|
||
return dpf::leaf_group_add(a, b);
|
||
}
|
||
|
||
static ring sub(const ring & a, const ring & b)
|
||
{
|
||
return dpf::leaf_group_add(a, b);
|
||
}
|
||
|
||
static ring mul(const ring & a, const ring & b)
|
||
{
|
||
return dpf::leaf_group_mul(a, b);
|
||
}
|
||
|
||
static ring neg(const ring & a) { return a; }
|
||
|
||
static ring sample() { return dpf::uniform_sample<ring>(); }
|
||
};
|
||
|
||
/// @brief IEEE `double` leaf group is XOR/AND of the bit pattern.
|
||
template <>
|
||
struct ring_traits<double>
|
||
{
|
||
using ring = double;
|
||
|
||
static ring zero() { return 0.0; }
|
||
|
||
static ring one() { return dpf::leaf_group_one<double>(); }
|
||
|
||
static ring add(const ring & a, const ring & b)
|
||
{
|
||
return dpf::leaf_group_add(a, b);
|
||
}
|
||
|
||
static ring sub(const ring & a, const ring & b)
|
||
{
|
||
return dpf::leaf_group_add(a, b);
|
||
}
|
||
|
||
static ring mul(const ring & a, const ring & b)
|
||
{
|
||
return dpf::leaf_group_mul(a, b);
|
||
}
|
||
|
||
static ring neg(const ring & a) { return a; }
|
||
|
||
static ring sample() { return dpf::uniform_sample<ring>(); }
|
||
};
|
||
|
||
/// @brief Width for `dpf::mod2k` arithmetic. Restored when the scope ends.
|
||
inline unsigned & mod2k_width() noexcept
|
||
{
|
||
static thread_local unsigned bits = 128u;
|
||
return bits;
|
||
}
|
||
|
||
struct mod2k_width_scope
|
||
{
|
||
unsigned previous;
|
||
|
||
explicit mod2k_width_scope(unsigned bits) noexcept
|
||
: previous(mod2k_width())
|
||
{
|
||
mod2k_width() = bits == 0u ? 128u : bits;
|
||
}
|
||
|
||
~mod2k_width_scope() { mod2k_width() = previous; }
|
||
|
||
mod2k_width_scope(const mod2k_width_scope &) = delete;
|
||
mod2k_width_scope & operator=(const mod2k_width_scope &) = delete;
|
||
};
|
||
|
||
/// @brief `Z/2^w Z` with `w` from `mod2k_width()`.
|
||
template <>
|
||
struct ring_traits<dpf::mod2k>
|
||
{
|
||
using ring = dpf::mod2k;
|
||
|
||
static unsigned bits() noexcept { return mod2k_width(); }
|
||
|
||
static simde_uint128 bit_mask(unsigned width) noexcept
|
||
{
|
||
if (width == 0u)
|
||
return 0;
|
||
if (width >= 128u)
|
||
return ~simde_uint128{0};
|
||
return (simde_uint128{1} << width) - 1;
|
||
}
|
||
|
||
static ring fit(simde_uint128 v) noexcept
|
||
{
|
||
return ring{v & bit_mask(bits())};
|
||
}
|
||
|
||
static ring zero() { return ring{}; }
|
||
|
||
static ring one() { return fit(1); }
|
||
|
||
static ring add(const ring & a, const ring & b)
|
||
{
|
||
return fit(a.raw + b.raw);
|
||
}
|
||
|
||
static ring sub(const ring & a, const ring & b)
|
||
{
|
||
return fit(a.raw - b.raw);
|
||
}
|
||
|
||
static ring mul(const ring & a, const ring & b)
|
||
{
|
||
const unsigned width = bits();
|
||
const auto m = bit_mask(width);
|
||
const simde_uint128 x = a.raw & m;
|
||
const simde_uint128 y = b.raw & m;
|
||
if (width <= 64u)
|
||
{
|
||
return fit(simde_uint128{static_cast<std::uint64_t>(x)}
|
||
* static_cast<std::uint64_t>(y));
|
||
}
|
||
const std::uint64_t a0 = static_cast<std::uint64_t>(x);
|
||
const std::uint64_t a1 = static_cast<std::uint64_t>(x >> 64);
|
||
const std::uint64_t b0 = static_cast<std::uint64_t>(y);
|
||
const std::uint64_t b1 = static_cast<std::uint64_t>(y >> 64);
|
||
const simde_uint128 p00 = simde_uint128{a0} * b0;
|
||
const simde_uint128 mid = (p00 >> 64)
|
||
+ static_cast<std::uint64_t>(simde_uint128{a0} * b1)
|
||
+ static_cast<std::uint64_t>(simde_uint128{a1} * b0);
|
||
return fit(simde_uint128{static_cast<std::uint64_t>(p00)} | (mid << 64));
|
||
}
|
||
|
||
static ring neg(const ring & a) { return sub(zero(), a); }
|
||
|
||
static ring sample()
|
||
{
|
||
return fit(dpf::uniform_sample<simde_uint128>());
|
||
}
|
||
};
|
||
|
||
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);
|
||
}
|
||
};
|
||
|
||
/// @brief Additive split of `(y, y·Δ)` under a Shark/SPDZ-style IT-MAC.
|
||
/// @details `value` is the usual Beaver/ABY2.0 share of `y`. `tag` is an
|
||
/// additive share of `y·Δ`. Party `p` holds `party(p)`.
|
||
/// @tparam Ring payload ring
|
||
template <typename Ring>
|
||
struct auth_split
|
||
{
|
||
split<Ring> value{};
|
||
split<Ring> tag{};
|
||
|
||
HEDLEY_NO_THROW
|
||
Ring open() const { return value.open(); }
|
||
|
||
HEDLEY_NO_THROW
|
||
dpf::mac_share<Ring> party(unsigned p) const
|
||
{
|
||
if (p == 0)
|
||
return dpf::mac_share<Ring>{value.p0, tag.p0};
|
||
return dpf::mac_share<Ring>{value.p1, tag.p1};
|
||
}
|
||
|
||
HEDLEY_NO_THROW
|
||
bool verify(const dpf::mac_key<Ring> & key) const
|
||
{
|
||
return dpf::mac_verify(party(0), party(1), key);
|
||
}
|
||
|
||
friend bool operator==(const auth_split & a, const auth_split & b)
|
||
{
|
||
return a.value == b.value && a.tag == b.tag;
|
||
}
|
||
|
||
friend bool operator!=(const auth_split & a, const auth_split & b)
|
||
{
|
||
return !(a == b);
|
||
}
|
||
};
|
||
|
||
/// @brief One party's contribution when opening an authenticated δ = x + λ.
|
||
/// @tparam Ring payload ring
|
||
template <typename Ring>
|
||
struct auth_opening
|
||
{
|
||
Ring value{};
|
||
Ring tag{};
|
||
};
|
||
|
||
/// @brief Attach an IT-MAC to an existing additive split under `key`.
|
||
/// @tparam Ring payload ring
|
||
/// @tparam Sample randomness for the tag mask
|
||
template <typename Ring, typename Sample>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_split<Ring> authenticate(const split<Ring> & s,
|
||
const dpf::mac_key<Ring> & key, Sample && rng)
|
||
{
|
||
using traits = ring_traits<Ring>;
|
||
const Ring y = traits::add(s.p0, s.p1);
|
||
const Ring t0 = rng();
|
||
const Ring t1 = traits::sub(traits::mul(y, key.delta), t0);
|
||
return auth_split<Ring>{s, split<Ring>{t0, t1}};
|
||
}
|
||
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_split<Ring> authenticate(const split<Ring> & s,
|
||
const dpf::mac_key<Ring> & key)
|
||
{
|
||
return authenticate(s, key, default_sampler<Ring>{});
|
||
}
|
||
|
||
/// @brief Authenticate a cleartext `y`, returning the dealer auth_split.
|
||
/// @tparam Ring payload ring
|
||
/// @tparam Sample randomness for both value and tag masks
|
||
template <typename Ring, typename Sample>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_split<Ring> auth_share(const Ring & y, const dpf::mac_key<Ring> & key,
|
||
Sample && rng)
|
||
{
|
||
using traits = ring_traits<Ring>;
|
||
auto & sampler = rng;
|
||
const Ring p0 = sampler();
|
||
const split<Ring> s{p0, traits::sub(y, p0)};
|
||
return authenticate(s, key, sampler);
|
||
}
|
||
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_split<Ring> auth_share(const Ring & y, const dpf::mac_key<Ring> & key)
|
||
{
|
||
return auth_share(y, key, default_sampler<Ring>{});
|
||
}
|
||
|
||
/// @brief Whether two parties' authenticated δ openings check under `key`.
|
||
/// @tparam Ring payload ring
|
||
template <typename Ring>
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_ALWAYS_INLINE
|
||
bool verify_auth_opening(auth_opening<Ring> a, auth_opening<Ring> b,
|
||
const dpf::mac_key<Ring> & key) noexcept
|
||
{
|
||
using traits = ring_traits<Ring>;
|
||
const Ring d = traits::add(a.value, b.value);
|
||
const Ring t = traits::add(a.tag, b.tag);
|
||
return t == traits::mul(d, key.delta);
|
||
}
|
||
|
||
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`. Tag masks are a third stream,
|
||
/// `tag_mask(role, index)`, domain-separated from blinds and share masks.
|
||
/// @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)
|
||
{
|
||
note_experiment_seed("beavers::oracle", seed());
|
||
}
|
||
|
||
explicit oracle(seed_type seed, std::size_t window = 256u)
|
||
: lanes_(seed, window),
|
||
tags_(dpf::randomness::detail::tag_master<PRG>(seed), window)
|
||
{
|
||
note_experiment_seed("beavers::oracle", this->seed());
|
||
}
|
||
|
||
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 IT-MAC tag-share mask for `role` at `index`. Same seed as the blinds.
|
||
Ring tag_mask(std::uint32_t role, std::uint64_t index) const
|
||
{
|
||
return tags_.value_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_;
|
||
dpf::randomness::lane_table<Ring, PRG> tags_;
|
||
};
|
||
|
||
/// @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 One party's view of sampled blinds and product shares.
|
||
/// @details Index order matches the session: wire id, monomial index, bundle
|
||
/// index, and gate index for dot crosses. A `*_ready` entry of false means
|
||
/// that slot was not sampled (no blind required). When `has_mac` is true the
|
||
/// parallel `*_tag` vectors hold this party's IT-MAC tag share for the same
|
||
/// slot (Shark/SPDZ-style `y·Δ`).
|
||
/// @tparam Ring payload ring
|
||
template <typename Ring>
|
||
struct party_tape
|
||
{
|
||
std::vector<Ring> lambda;
|
||
std::vector<std::uint8_t> lambda_ready;
|
||
std::vector<Ring> monomial;
|
||
std::vector<std::uint8_t> monomial_ready;
|
||
std::vector<Ring> bundles;
|
||
std::vector<std::uint8_t> bundles_ready;
|
||
std::vector<Ring> dot_cross;
|
||
std::vector<std::uint8_t> dot_ready;
|
||
|
||
/// @brief IT-MAC tag shares; sized like `lambda` when `has_mac`.
|
||
std::vector<Ring> lambda_tag;
|
||
std::vector<Ring> monomial_tag;
|
||
std::vector<Ring> bundles_tag;
|
||
std::vector<Ring> dot_cross_tag;
|
||
bool has_mac = false;
|
||
};
|
||
|
||
/// @brief How `schedule_terms` ranks factorizations.
|
||
/// @details `prep` minimizes preprocessing (bundles + blinds), then rounds —
|
||
/// the ABY2.0 default matching Appendix-E peels. `rounds` minimizes
|
||
/// interactive depth first (Pika / Grotto online latency), then prep.
|
||
enum class schedule_objective : unsigned char
|
||
{
|
||
prep = 0,
|
||
rounds = 1
|
||
};
|
||
|
||
/// @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 Rank factorizations by prep (default) or by interactive rounds.
|
||
HEDLEY_NO_THROW
|
||
void set_schedule_objective(schedule_objective o) noexcept
|
||
{
|
||
schedule_objective_ = o;
|
||
}
|
||
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
schedule_objective get_schedule_objective() const noexcept
|
||
{
|
||
return schedule_objective_;
|
||
}
|
||
|
||
/// @brief Enable Shark/SPDZ-style IT-MAC tags on blinds and values.
|
||
/// @details Call before `sample` / `bind`. Subsequent sampling authenticates
|
||
/// every λ, monomial, bundle, and dot-cross under `key`. Bound and
|
||
/// evaluated value shares are tagged the same way. Party tapes then
|
||
/// carry tag lanes; δ openings can be checked with `verify_delta`.
|
||
/// @param key global MAC key `Δ` (dealer-held for verification)
|
||
void set_mac_key(dpf::mac_key<Ring> key)
|
||
{
|
||
mac_key_ = std::move(key);
|
||
has_mac_ = true;
|
||
mac_key_known_ = true;
|
||
}
|
||
|
||
HEDLEY_NO_THROW
|
||
bool has_mac() const noexcept { return has_mac_; }
|
||
|
||
/// @brief Dealer MAC key. Undefined when `has_mac()` is false.
|
||
/// @throws std::logic_error if MACs were not enabled
|
||
const dpf::mac_key<Ring> & mac_key() const
|
||
{
|
||
if (!mac_key_known_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
return mac_key_;
|
||
}
|
||
|
||
/// @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`
|
||
/// \complexity One draw per wire that `needs_blind`, plus one share per monomial. Local; no messages.
|
||
/// \preprocessing That tape is what `evaluate_party` later opens: a full blind and a two-party split per such wire.
|
||
template <typename Sample>
|
||
void sample(Sample && sampler)
|
||
{
|
||
struct draws
|
||
{
|
||
Sample && rng;
|
||
|
||
Ring blind(std::uint32_t, std::uint64_t) { return rng(); }
|
||
|
||
Ring mask(std::uint32_t, std::uint64_t) { return rng(); }
|
||
|
||
Ring tag_mask(std::uint32_t, std::uint64_t) { return rng(); }
|
||
};
|
||
sample_with(draws{std::forward<Sample>(sampler)}, 0);
|
||
}
|
||
|
||
void sample()
|
||
{
|
||
sample(default_sampler<Ring>{});
|
||
}
|
||
|
||
/// @brief Drop sampled blinds and product shares. The recorded formula stays.
|
||
/// @details The next `sample` / `sample_from` draws that formula again.
|
||
void clear_sample()
|
||
{
|
||
for (auto & w : wires_)
|
||
w.lambda_ready = false;
|
||
for (auto & m : monos_)
|
||
m.ready = false;
|
||
for (auto & b : bundles_)
|
||
b.ready = false;
|
||
for (auto & g : gates_)
|
||
g.cross_ready = false;
|
||
}
|
||
|
||
/// @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. MAC tag masks
|
||
/// come from `src.tag_mask` for the same role.
|
||
/// @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)
|
||
{
|
||
sample_with(src, index);
|
||
}
|
||
|
||
/// @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;
|
||
if (has_mac_)
|
||
tag_split(wires_[id].value, wires_[id].value_tag, rng);
|
||
}
|
||
|
||
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;
|
||
if (has_mac_)
|
||
{
|
||
default_sampler<Ring> rng;
|
||
tag_split(wires_[id].value, wires_[id].value_tag, rng);
|
||
}
|
||
}
|
||
|
||
/// @brief One party's additive shares of every sampled blind and product.
|
||
/// @param party `0` or `1`
|
||
/// @return Tape suitable for `install_party` / network deal
|
||
/// @throws std::invalid_argument if `party` is not 0 or 1
|
||
party_tape<Ring> export_party(unsigned party) const
|
||
{
|
||
if (party > 1)
|
||
throw std::invalid_argument("beaver party must be 0 or 1");
|
||
party_tape<Ring> out;
|
||
out.has_mac = has_mac_;
|
||
out.lambda.resize(wires_.size());
|
||
out.lambda_ready.resize(wires_.size());
|
||
if (has_mac_)
|
||
out.lambda_tag.resize(wires_.size());
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
out.lambda_ready[id] = wires_[id].lambda_ready ? 1 : 0;
|
||
if (wires_[id].lambda_ready)
|
||
{
|
||
out.lambda[id] = party == 0 ? wires_[id].lambda.p0 : wires_[id].lambda.p1;
|
||
if (has_mac_)
|
||
out.lambda_tag[id] = party == 0
|
||
? wires_[id].lambda_tag.p0
|
||
: wires_[id].lambda_tag.p1;
|
||
}
|
||
}
|
||
out.monomial.resize(monos_.size());
|
||
out.monomial_ready.resize(monos_.size());
|
||
if (has_mac_)
|
||
out.monomial_tag.resize(monos_.size());
|
||
for (std::size_t i = 0; i < monos_.size(); ++i)
|
||
{
|
||
out.monomial_ready[i] = monos_[i].ready ? 1 : 0;
|
||
if (monos_[i].ready)
|
||
{
|
||
out.monomial[i] = party == 0 ? monos_[i].share.p0 : monos_[i].share.p1;
|
||
if (has_mac_)
|
||
out.monomial_tag[i] = party == 0
|
||
? monos_[i].tag.p0
|
||
: monos_[i].tag.p1;
|
||
}
|
||
}
|
||
out.bundles.resize(bundles_.size());
|
||
out.bundles_ready.resize(bundles_.size());
|
||
if (has_mac_)
|
||
out.bundles_tag.resize(bundles_.size());
|
||
for (std::size_t i = 0; i < bundles_.size(); ++i)
|
||
{
|
||
out.bundles_ready[i] = bundles_[i].ready ? 1 : 0;
|
||
if (bundles_[i].ready)
|
||
{
|
||
out.bundles[i] = party == 0 ? bundles_[i].share.p0 : bundles_[i].share.p1;
|
||
if (has_mac_)
|
||
out.bundles_tag[i] = party == 0
|
||
? bundles_[i].tag.p0
|
||
: bundles_[i].tag.p1;
|
||
}
|
||
}
|
||
out.dot_cross.resize(gates_.size());
|
||
out.dot_ready.resize(gates_.size());
|
||
if (has_mac_)
|
||
out.dot_cross_tag.resize(gates_.size());
|
||
for (std::size_t g = 0; g < gates_.size(); ++g)
|
||
{
|
||
out.dot_ready[g] = gates_[g].cross_ready ? 1 : 0;
|
||
if (gates_[g].cross_ready)
|
||
{
|
||
out.dot_cross[g] = party == 0 ? gates_[g].cross.p0 : gates_[g].cross.p1;
|
||
if (has_mac_)
|
||
out.dot_cross_tag[g] = party == 0
|
||
? gates_[g].cross_tag.p0
|
||
: gates_[g].cross_tag.p1;
|
||
}
|
||
}
|
||
return out;
|
||
}
|
||
|
||
/// @brief Install this party's tape after the circuit is recorded. Does not sample.
|
||
/// @param party `0` or `1`
|
||
/// @param tape shares from `export_party` / the dealer
|
||
/// @throws std::invalid_argument if sizes disagree or `party` is bad
|
||
/// @throws std::logic_error if blinds were already sampled
|
||
void install_party(unsigned party, const party_tape<Ring> & tape)
|
||
{
|
||
if (party > 1)
|
||
throw std::invalid_argument("beaver party must be 0 or 1");
|
||
if (tape.lambda.size() != wires_.size()
|
||
|| tape.lambda_ready.size() != wires_.size()
|
||
|| tape.monomial.size() != monos_.size()
|
||
|| tape.monomial_ready.size() != monos_.size()
|
||
|| tape.bundles.size() != bundles_.size()
|
||
|| tape.bundles_ready.size() != bundles_.size()
|
||
|| tape.dot_cross.size() != gates_.size()
|
||
|| tape.dot_ready.size() != gates_.size())
|
||
throw std::invalid_argument("beaver party tape size mismatch");
|
||
if (tape.has_mac)
|
||
{
|
||
if (tape.lambda_tag.size() != wires_.size()
|
||
|| tape.monomial_tag.size() != monos_.size()
|
||
|| tape.bundles_tag.size() != bundles_.size()
|
||
|| tape.dot_cross_tag.size() != gates_.size())
|
||
throw std::invalid_argument("beaver party MAC tape size mismatch");
|
||
has_mac_ = true;
|
||
}
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
if (wires_[id].lambda_ready)
|
||
throw std::logic_error("call install_party before sample");
|
||
if (!tape.lambda_ready[id])
|
||
continue;
|
||
wires_[id].lambda = party == 0
|
||
? split<Ring>{tape.lambda[id], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.lambda[id]};
|
||
wires_[id].lambda_ready = true;
|
||
if (tape.has_mac)
|
||
{
|
||
wires_[id].lambda_tag = party == 0
|
||
? split<Ring>{tape.lambda_tag[id], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.lambda_tag[id]};
|
||
}
|
||
}
|
||
for (std::size_t i = 0; i < monos_.size(); ++i)
|
||
{
|
||
if (monos_[i].ready)
|
||
throw std::logic_error("call install_party before sample");
|
||
if (!tape.monomial_ready[i])
|
||
continue;
|
||
monos_[i].share = party == 0
|
||
? split<Ring>{tape.monomial[i], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.monomial[i]};
|
||
monos_[i].ready = true;
|
||
if (tape.has_mac)
|
||
{
|
||
monos_[i].tag = party == 0
|
||
? split<Ring>{tape.monomial_tag[i], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.monomial_tag[i]};
|
||
}
|
||
}
|
||
for (std::size_t i = 0; i < bundles_.size(); ++i)
|
||
{
|
||
if (bundles_[i].ready)
|
||
throw std::logic_error("call install_party before sample");
|
||
if (!tape.bundles_ready[i])
|
||
continue;
|
||
bundles_[i].share = party == 0
|
||
? split<Ring>{tape.bundles[i], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.bundles[i]};
|
||
bundles_[i].ready = true;
|
||
if (tape.has_mac)
|
||
{
|
||
bundles_[i].tag = party == 0
|
||
? split<Ring>{tape.bundles_tag[i], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.bundles_tag[i]};
|
||
}
|
||
}
|
||
for (std::size_t g = 0; g < gates_.size(); ++g)
|
||
{
|
||
if (gates_[g].cross_ready)
|
||
throw std::logic_error("call install_party before sample");
|
||
if (!tape.dot_ready[g])
|
||
continue;
|
||
gates_[g].cross = party == 0
|
||
? split<Ring>{tape.dot_cross[g], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.dot_cross[g]};
|
||
gates_[g].cross_ready = true;
|
||
if (tape.has_mac)
|
||
{
|
||
gates_[g].cross_tag = party == 0
|
||
? split<Ring>{tape.dot_cross_tag[g], traits::zero()}
|
||
: split<Ring>{traits::zero(), tape.dot_cross_tag[g]};
|
||
}
|
||
}
|
||
party_mode_ = true;
|
||
party_id_ = party;
|
||
}
|
||
|
||
/// @brief Bind this party's additive share of an input secret.
|
||
/// @param w input wire
|
||
/// @param share this party's share
|
||
/// @throws std::invalid_argument if `w` is not an input
|
||
/// @throws std::logic_error if already bound or `install_party` was not called
|
||
void bind_party(wire w, Ring share)
|
||
{
|
||
if (!party_mode_)
|
||
throw std::logic_error("call install_party before bind_party");
|
||
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");
|
||
wires_[id].value = party_id_ == 0
|
||
? split<Ring>{share, traits::zero()}
|
||
: split<Ring>{traits::zero(), share};
|
||
wires_[id].value_ready = true;
|
||
}
|
||
|
||
/// @brief Bind an authenticated input share `(value, tag)` for this party.
|
||
/// @param w input wire
|
||
/// @param share this party's `mac_share` of the secret
|
||
/// @throws std::logic_error if MACs were not enabled on the tape
|
||
void bind_party(wire w, dpf::mac_share<Ring> share)
|
||
{
|
||
if (!party_mode_)
|
||
throw std::logic_error("call install_party before bind_party");
|
||
if (!has_mac_)
|
||
throw std::logic_error("bind_party MAC share requires an authenticated tape");
|
||
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");
|
||
wires_[id].value = party_id_ == 0
|
||
? split<Ring>{share.value, traits::zero()}
|
||
: split<Ring>{traits::zero(), share.value};
|
||
wires_[id].value_tag = party_id_ == 0
|
||
? split<Ring>{share.tag, traits::zero()}
|
||
: split<Ring>{traits::zero(), share.tag};
|
||
wires_[id].value_ready = true;
|
||
}
|
||
|
||
/// @brief Open every ready round using peer exchanges for public δ.
|
||
/// @details `exchange_delta(mine)` returns the peer's contribution
|
||
/// `value_share + lambda_share`. Public monomial terms land on party 0.
|
||
/// When the session is MAC-authenticated, use the `auth_opening` overload
|
||
/// so tag shares of δ are exchanged too.
|
||
/// @tparam Exchange callable `Ring(Ring)`
|
||
/// @param exchange_delta peer opening callback
|
||
/// @throws std::logic_error if not in party mode or MACs are enabled
|
||
/// \complexity O(R G) gate steps. R is max `ready_round`, G is the number of gates.
|
||
/// \rounds R. A product of fresh inputs is round 1. A gate whose inputs are earlier outputs waits for `ready_round`.
|
||
/// \communication One `Ring` per wire whose delta is opened (`ensure_delta_party`). A wire already opened is not sent again.
|
||
/// \preprocessing `sample` drew one full blind and a two-party split per wire that `needs_blind`, and one share per monomial.
|
||
template <typename Exchange>
|
||
void evaluate_party(Exchange && exchange_delta)
|
||
{
|
||
if (!party_mode_)
|
||
throw std::logic_error("call install_party before evaluate_party");
|
||
if (has_mac_)
|
||
throw std::logic_error("authenticated session needs evaluate_party with auth_opening");
|
||
evaluate_party_impl(std::forward<Exchange>(exchange_delta));
|
||
}
|
||
|
||
/// @brief Like `evaluate_party`, but open many δ values per ready-round in one vector exchange.
|
||
/// @details At the start of each ready-round, every already-valued input that
|
||
/// the round will open is exchanged once (shared order on both
|
||
/// parties). Gate outputs that need δ are batched at the end of
|
||
/// the round, flushing earlier if a later same-round gate reads
|
||
/// that δ. Logical `round_of` is unchanged; channel barriers drop
|
||
/// from one-per-wire to a small constant per ready-round.
|
||
/// @tparam BatchExchange callable `std::vector<Ring>(std::vector<Ring>)`
|
||
/// @param exchange_deltas peer opening callback for a packed vector
|
||
template <typename BatchExchange>
|
||
void evaluate_party_batch(BatchExchange && exchange_deltas)
|
||
{
|
||
if (!party_mode_)
|
||
throw std::logic_error("call install_party before evaluate_party_batch");
|
||
if (has_mac_)
|
||
throw std::logic_error("authenticated session needs evaluate_party_auth_batch");
|
||
evaluate_party_batch_impl(std::forward<BatchExchange>(exchange_deltas));
|
||
}
|
||
|
||
/// @brief Authenticated opening path: exchange `(value+λ, tag_value+tag_λ)`.
|
||
/// @tparam Exchange callable `auth_opening<Ring>(auth_opening<Ring>)`
|
||
/// @param exchange_delta peer opening callback
|
||
/// \complexity Same gate walk as `evaluate_party`.
|
||
/// \rounds R, the max `ready_round`.
|
||
/// \communication One `auth_opening` (value and MAC tag) per newly opened wire, instead of one `Ring`.
|
||
/// \preprocessing `sample` also drew a tag split when the session has a MAC key.
|
||
template <typename Exchange>
|
||
void evaluate_party_auth(Exchange && exchange_delta)
|
||
{
|
||
if (!party_mode_)
|
||
throw std::logic_error("call install_party before evaluate_party_auth");
|
||
if (!has_mac_)
|
||
throw std::logic_error("evaluate_party_auth requires an authenticated tape");
|
||
evaluate_party_auth_impl(std::forward<Exchange>(exchange_delta));
|
||
}
|
||
|
||
/// @brief Authenticated batched opening path (vector of `auth_opening`).
|
||
/// @tparam BatchExchange callable `std::vector<auth_opening<Ring>>(std::vector<auth_opening<Ring>>)`
|
||
template <typename BatchExchange>
|
||
void evaluate_party_auth_batch(BatchExchange && exchange_deltas)
|
||
{
|
||
if (!party_mode_)
|
||
throw std::logic_error("call install_party before evaluate_party_auth_batch");
|
||
if (!has_mac_)
|
||
throw std::logic_error("evaluate_party_auth_batch requires an authenticated tape");
|
||
evaluate_party_auth_batch_impl(std::forward<BatchExchange>(exchange_deltas));
|
||
}
|
||
|
||
/// @brief This party's value share after `evaluate_party`.
|
||
Ring value_party(wire w) const
|
||
{
|
||
auto id = check(w);
|
||
if (!wires_[id].value_ready)
|
||
throw std::logic_error("beaver wire has no value yet");
|
||
return party_id_ == 0 ? wires_[id].value.p0 : wires_[id].value.p1;
|
||
}
|
||
|
||
/// @brief This party's MAC share of the wire value after binding / evaluate.
|
||
/// @throws std::logic_error if MACs are off or the value is missing
|
||
dpf::mac_share<Ring> value_mac_party(wire w) const
|
||
{
|
||
if (!has_mac_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
auto id = check(w);
|
||
if (!wires_[id].value_ready)
|
||
throw std::logic_error("beaver wire has no value yet");
|
||
return party_id_ == 0
|
||
? dpf::mac_share<Ring>{wires_[id].value.p0, wires_[id].value_tag.p0}
|
||
: dpf::mac_share<Ring>{wires_[id].value.p1, wires_[id].value_tag.p1};
|
||
}
|
||
|
||
/// @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`
|
||
/// \complexity O(R G) local additions. R is max `ready_round`, G is the number of gates. The dealer holds both shares, so this sends nothing.
|
||
/// \rounds R, counted locally. No messages.
|
||
/// \preprocessing The blinds from `sample`.
|
||
void evaluate()
|
||
{
|
||
int max_round = 0;
|
||
for (const auto & w : wires_)
|
||
max_round = std::max(max_round, w.ready_round);
|
||
default_sampler<Ring> rng;
|
||
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 (has_mac_)
|
||
tag_split(wires_[g.out].value, wires_[g.out].value_tag, rng);
|
||
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));
|
||
}
|
||
|
||
/// @brief Authenticated share of `Π λ_i^{e_i}` under the session MAC key.
|
||
/// @throws std::logic_error if MACs are off or the monomial was not sampled
|
||
auth_split<Ring> monomial_auth(const std::vector<power> & spec) const
|
||
{
|
||
if (!has_mac_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
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);
|
||
}
|
||
auto located = locate_term(merge_exponents(raw));
|
||
return auth_split<Ring>{*located.share, *located.tag};
|
||
}
|
||
|
||
auth_split<Ring> monomial_auth(std::initializer_list<power> spec) const
|
||
{
|
||
return monomial_auth(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;
|
||
}
|
||
|
||
/// @brief Authenticated blind `[λ]` under the session MAC key.
|
||
/// @throws std::logic_error if MACs are off or the blind is missing
|
||
auth_split<Ring> lambda_auth(wire w) const
|
||
{
|
||
if (!has_mac_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
auto id = check(w);
|
||
if (!wires_[id].lambda_ready)
|
||
throw std::logic_error("call sample() before reading a beaver blind");
|
||
return auth_split<Ring>{wires_[id].lambda, wires_[id].lambda_tag};
|
||
}
|
||
|
||
/// @brief Authenticated value `[x]` under the session MAC key.
|
||
/// @throws std::logic_error if MACs are off or the value is missing
|
||
auth_split<Ring> value_auth(wire w) const
|
||
{
|
||
if (!has_mac_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
auto id = check(w);
|
||
if (!wires_[id].value_ready)
|
||
throw std::logic_error("beaver wire has no value yet");
|
||
return auth_split<Ring>{wires_[id].value, wires_[id].value_tag};
|
||
}
|
||
|
||
/// @brief Authenticated opening of δ = x + λ (value and tag sums).
|
||
/// @details The opened public δ equals `delta(w)`. Tags must satisfy
|
||
/// `tag = δ · Δ` when both parties' shares are combined.
|
||
auth_split<Ring> delta_auth(wire w) const
|
||
{
|
||
if (!has_mac_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
auto id = check(w);
|
||
if (!wires_[id].delta_ready)
|
||
throw std::logic_error("beaver wire has not been opened");
|
||
if (!wires_[id].lambda_ready || !wires_[id].value_ready)
|
||
throw std::logic_error("beaver wire is not ready to open");
|
||
return auth_split<Ring>{
|
||
split<Ring>{
|
||
traits::add(wires_[id].value.p0, wires_[id].lambda.p0),
|
||
traits::add(wires_[id].value.p1, wires_[id].lambda.p1)},
|
||
split<Ring>{
|
||
traits::add(wires_[id].value_tag.p0, wires_[id].lambda_tag.p0),
|
||
traits::add(wires_[id].value_tag.p1, wires_[id].lambda_tag.p1)}};
|
||
}
|
||
|
||
/// @brief Whether `[x]` verifies under the session MAC key.
|
||
bool verify_value(wire w) const
|
||
{
|
||
if (!mac_key_known_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
return value_auth(w).verify(mac_key_);
|
||
}
|
||
|
||
/// @brief Whether the opened δ = x + λ verifies under the session MAC key.
|
||
bool verify_delta(wire w) const
|
||
{
|
||
if (!mac_key_known_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
auto id = check(w);
|
||
if (wires_[id].delta_open_ready)
|
||
{
|
||
const auto & mine = wires_[id].delta_mine;
|
||
const auto & peer = wires_[id].delta_peer;
|
||
const Ring d = traits::add(mine.value, peer.value);
|
||
const Ring t = traits::add(mine.tag, peer.tag);
|
||
return d == wires_[id].delta
|
||
&& t == traits::mul(d, mac_key_.delta);
|
||
}
|
||
auto a = delta_auth(w);
|
||
if (a.open() != wires_[id].delta)
|
||
return false;
|
||
return a.verify(mac_key_);
|
||
}
|
||
|
||
/// @brief Check a party opening under an explicit key (dealer audit).
|
||
bool verify_delta(wire w, const dpf::mac_key<Ring> & key) const
|
||
{
|
||
auto id = check(w);
|
||
if (wires_[id].delta_open_ready)
|
||
return verify_auth_opening(wires_[id].delta_mine, wires_[id].delta_peer, key)
|
||
&& traits::add(wires_[id].delta_mine.value, wires_[id].delta_peer.value)
|
||
== wires_[id].delta;
|
||
auth_split<Ring> a{
|
||
split<Ring>{
|
||
traits::add(wires_[id].value.p0, wires_[id].lambda.p0),
|
||
traits::add(wires_[id].value.p1, wires_[id].lambda.p1)},
|
||
split<Ring>{
|
||
traits::add(wires_[id].value_tag.p0, wires_[id].lambda_tag.p0),
|
||
traits::add(wires_[id].value_tag.p1, wires_[id].lambda_tag.p1)}};
|
||
return a.open() == wires_[id].delta && a.verify(key);
|
||
}
|
||
|
||
/// @brief Party-local view of the last authenticated δ opening for `w`.
|
||
auth_opening<Ring> delta_opening_mine(wire w) const
|
||
{
|
||
auto id = check(w);
|
||
if (!wires_[id].delta_open_ready)
|
||
throw std::logic_error("beaver wire has no authenticated opening");
|
||
return wires_[id].delta_mine;
|
||
}
|
||
|
||
auth_opening<Ring> delta_opening_peer(wire w) const
|
||
{
|
||
auto id = check(w);
|
||
if (!wires_[id].delta_open_ready)
|
||
throw std::logic_error("beaver wire has no authenticated opening");
|
||
return wires_[id].delta_peer;
|
||
}
|
||
|
||
/// @brief Verify every opened wire's value and δ MAC (dealer / test path).
|
||
bool verify_all() const
|
||
{
|
||
if (!has_mac_)
|
||
return true;
|
||
if (!mac_key_known_)
|
||
throw std::logic_error("beaver session has no MAC key");
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
if (wires_[id].lambda_ready
|
||
&& !auth_split<Ring>{wires_[id].lambda, wires_[id].lambda_tag}
|
||
.verify(mac_key_))
|
||
return false;
|
||
if (wires_[id].value_ready
|
||
&& !auth_split<Ring>{wires_[id].value, wires_[id].value_tag}
|
||
.verify(mac_key_))
|
||
return false;
|
||
if (wires_[id].delta_ready)
|
||
{
|
||
if (wires_[id].delta_open_ready)
|
||
{
|
||
if (!verify_auth_opening(wires_[id].delta_mine,
|
||
wires_[id].delta_peer, mac_key_))
|
||
return false;
|
||
}
|
||
else
|
||
{
|
||
auth_split<Ring> d{
|
||
split<Ring>{
|
||
traits::add(wires_[id].value.p0, wires_[id].lambda.p0),
|
||
traits::add(wires_[id].value.p1, wires_[id].lambda.p1)},
|
||
split<Ring>{
|
||
traits::add(wires_[id].value_tag.p0, wires_[id].lambda_tag.p0),
|
||
traits::add(wires_[id].value_tag.p1, wires_[id].lambda_tag.p1)}};
|
||
if (d.open() != wires_[id].delta || !d.verify(mac_key_))
|
||
return false;
|
||
}
|
||
}
|
||
}
|
||
for (const auto & m : monos_)
|
||
{
|
||
if (m.ready
|
||
&& !auth_split<Ring>{m.share, m.tag}.verify(mac_key_))
|
||
return false;
|
||
}
|
||
for (const auto & b : bundles_)
|
||
{
|
||
if (b.ready
|
||
&& !auth_split<Ring>{b.share, b.tag}.verify(mac_key_))
|
||
return false;
|
||
}
|
||
for (const auto & g : gates_)
|
||
{
|
||
if (g.cross_ready
|
||
&& !auth_split<Ring>{g.cross, g.cross_tag}.verify(mac_key_))
|
||
return false;
|
||
}
|
||
return true;
|
||
}
|
||
|
||
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;
|
||
}
|
||
|
||
/// @brief Rebuild a wire handle from its id (same session).
|
||
wire wire_at(std::uint32_t id) const
|
||
{
|
||
if (static_cast<std::size_t>(id) >= wires_.size())
|
||
throw std::out_of_range("beaver wire_at");
|
||
wire w;
|
||
w.sess_ = const_cast<session *>(this);
|
||
w.id_ = id;
|
||
return w;
|
||
}
|
||
|
||
/// @brief Maximum `ready_round` over every wire (0 when the circuit is empty).
|
||
int max_ready_round() const
|
||
{
|
||
int m = 0;
|
||
for (const auto & w : wires_)
|
||
m = std::max(m, w.ready_round);
|
||
return m;
|
||
}
|
||
|
||
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;
|
||
}
|
||
|
||
/// @brief One δ-vector exchange performed by `evaluate_party_batch`.
|
||
struct delta_barrier
|
||
{
|
||
int ready_round = 0;
|
||
std::vector<std::uint32_t> wire_ids;
|
||
};
|
||
|
||
/// @brief Every flush `evaluate_party_batch` will perform, in order.
|
||
/// @details Dry-runs the batch control flow on the gate graph. Input wires
|
||
/// are treated as bound; blinds are treated as sampled for every
|
||
/// wire that `needs_blind`. Same-round splits match the live
|
||
/// evaluator. Used by the protocol composer to place δ opens on
|
||
/// RoundSink waves.
|
||
std::vector<delta_barrier> exchange_barriers() const
|
||
{
|
||
return compile_batch_barriers();
|
||
}
|
||
|
||
template <typename R>
|
||
friend class party_batch_stepper;
|
||
template <typename R>
|
||
friend class party_auth_batch_stepper;
|
||
|
||
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{};
|
||
split<Ring> tag{};
|
||
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> lambda_tag{};
|
||
split<Ring> value{};
|
||
split<Ring> value_tag{};
|
||
Ring delta{};
|
||
auth_opening<Ring> delta_mine{};
|
||
auth_opening<Ring> delta_peer{};
|
||
bool delta_open_ready = false;
|
||
};
|
||
|
||
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{};
|
||
split<Ring> cross_tag{};
|
||
bool cross_ready = false;
|
||
};
|
||
|
||
struct mono
|
||
{
|
||
exp_list key;
|
||
split<Ring> share{};
|
||
split<Ring> tag{};
|
||
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_;
|
||
bool party_mode_ = false;
|
||
unsigned party_id_ = 0;
|
||
bool has_mac_ = false;
|
||
bool mac_key_known_ = false;
|
||
dpf::mac_key<Ring> mac_key_{};
|
||
schedule_objective schedule_objective_ = schedule_objective::prep;
|
||
|
||
Ring side_of(const split<Ring> & s) const
|
||
{
|
||
return party_id_ == 0 ? s.p0 : s.p1;
|
||
}
|
||
|
||
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)};
|
||
}
|
||
|
||
template <typename Sample>
|
||
void tag_split(const split<Ring> & value, split<Ring> & tag, Sample & rng) const
|
||
{
|
||
tag_assigned(value, tag, rng());
|
||
}
|
||
|
||
void tag_assigned(const split<Ring> & value, split<Ring> & tag, const Ring & mask) const
|
||
{
|
||
const Ring y = traits::add(value.p0, value.p1);
|
||
tag.p0 = mask;
|
||
tag.p1 = traits::sub(traits::mul(y, mac_key_.delta), mask);
|
||
}
|
||
|
||
/// @brief One construction for a callable sampler and for an oracle.
|
||
/// @details `src.blind` / `src.mask` / `src.tag_mask` are called in that
|
||
/// order for each new wire, monomial, bundle, and dot cross.
|
||
/// A sequential sampler ignores the role and the index.
|
||
template <typename Source>
|
||
void sample_with(Source && src, std::uint64_t index)
|
||
{
|
||
for (std::uint32_t id = 0; id < wires_.size(); ++id)
|
||
{
|
||
auto & w = wires_[id];
|
||
if (w.lambda_ready || !needs_blind(id))
|
||
continue;
|
||
const auto role = oracle<Ring>::wire_role(id);
|
||
w.lambda_full = src.blind(role, index);
|
||
const Ring p0 = src.mask(role, index);
|
||
w.lambda = split<Ring>{p0, traits::sub(w.lambda_full, p0)};
|
||
w.lambda_ready = true;
|
||
if (has_mac_)
|
||
tag_assigned(w.lambda, w.lambda_tag, src.tag_mask(role, index));
|
||
}
|
||
for (std::uint32_t i = 0; i < monos_.size(); ++i)
|
||
{
|
||
auto & m = monos_[i];
|
||
if (m.ready)
|
||
continue;
|
||
const auto role = oracle<Ring>::mono_role(i);
|
||
const Ring full = mono_full(m.key);
|
||
const Ring p0 = src.mask(role, index);
|
||
m.share = split<Ring>{p0, traits::sub(full, p0)};
|
||
m.ready = true;
|
||
if (has_mac_)
|
||
tag_assigned(m.share, m.tag, src.tag_mask(role, index));
|
||
}
|
||
for (std::uint32_t i = 0; i < bundles_.size(); ++i)
|
||
{
|
||
auto & b = bundles_[i];
|
||
if (b.ready)
|
||
continue;
|
||
const auto role = oracle<Ring>::bundle_role(i);
|
||
const Ring full = bundle_value(b);
|
||
const Ring p0 = src.mask(role, index);
|
||
b.share = split<Ring>{p0, traits::sub(full, p0)};
|
||
b.ready = true;
|
||
if (has_mac_)
|
||
tag_assigned(b.share, b.tag, src.tag_mask(role, index));
|
||
}
|
||
for (std::uint32_t g = 0; g < gates_.size(); ++g)
|
||
{
|
||
auto & gate = gates_[g];
|
||
if (gate.kind != gate_kind::dot || gate.cross_ready)
|
||
continue;
|
||
const auto role = oracle<Ring>::dot_role(g);
|
||
const Ring full = dot_full(gate);
|
||
const Ring p0 = src.mask(role, index);
|
||
gate.cross = split<Ring>{p0, traits::sub(full, p0)};
|
||
gate.cross_ready = true;
|
||
if (has_mac_)
|
||
tag_assigned(gate.cross, gate.cross_tag, src.tag_mask(role, index));
|
||
}
|
||
}
|
||
|
||
template <typename Exchange>
|
||
void evaluate_party_impl(Exchange && exchange_delta)
|
||
{
|
||
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;
|
||
Ring val;
|
||
if (g.kind == gate_kind::product)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta_party(id, exchange_delta);
|
||
val = eval_product_party(g.lhs);
|
||
}
|
||
else if (g.kind == gate_kind::dot)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta_party(id, exchange_delta);
|
||
for (auto id : g.rhs)
|
||
ensure_delta_party(id, exchange_delta);
|
||
val = eval_dot_party(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_party(id, exchange_delta);
|
||
}
|
||
if (step.mask_wire >= 0)
|
||
ensure_delta_party(
|
||
static_cast<std::uint32_t>(step.mask_wire), exchange_delta);
|
||
}
|
||
val = eval_poly_party(g);
|
||
}
|
||
else
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta_party(id, exchange_delta);
|
||
val = eval_mux_party(g);
|
||
}
|
||
wires_[g.out].value = party_id_ == 0
|
||
? split<Ring>{val, traits::zero()}
|
||
: split<Ring>{traits::zero(), val};
|
||
wires_[g.out].value_ready = true;
|
||
if (wires_[g.out].lambda_ready)
|
||
ensure_delta_party(g.out, exchange_delta);
|
||
}
|
||
}
|
||
}
|
||
|
||
/// @brief One step of the compiled batch schedule: a δ flush or a gate eval.
|
||
struct batch_op
|
||
{
|
||
enum class kind : unsigned char { flush, eval_gate } type = kind::flush;
|
||
int ready_round = 0;
|
||
std::vector<std::uint32_t> wire_ids;
|
||
std::size_t gate_index = 0;
|
||
};
|
||
|
||
/// @brief Compile the batch control flow without opening wires.
|
||
std::vector<batch_op> compile_batch_ops() const
|
||
{
|
||
std::vector<batch_op> ops;
|
||
const std::size_t n = wires_.size();
|
||
std::vector<char> value_ready(n, 0);
|
||
std::vector<char> delta_ready(n, 0);
|
||
std::vector<char> lambda_ready(n, 0);
|
||
for (std::uint32_t id = 0; id < n; ++id)
|
||
{
|
||
if (wires_[id].ready_round == 0)
|
||
value_ready[id] = 1;
|
||
if (needs_blind(id) || wires_[id].lambda_ready)
|
||
lambda_ready[id] = 1;
|
||
}
|
||
|
||
int max_round = 0;
|
||
for (const auto & w : wires_)
|
||
max_round = std::max(max_round, w.ready_round);
|
||
|
||
std::vector<char> queued(n, 0);
|
||
std::vector<std::uint32_t> pending;
|
||
pending.reserve(n);
|
||
|
||
auto note = [&](std::uint32_t id) {
|
||
if (delta_ready[id] || queued[id])
|
||
return;
|
||
if (!lambda_ready[id] || !value_ready[id])
|
||
return;
|
||
queued[id] = 1;
|
||
pending.push_back(id);
|
||
};
|
||
|
||
auto record_flush = [&](int round) {
|
||
if (pending.empty())
|
||
return;
|
||
batch_op op;
|
||
op.type = batch_op::kind::flush;
|
||
op.ready_round = round;
|
||
op.wire_ids = pending;
|
||
ops.push_back(std::move(op));
|
||
for (auto id : pending)
|
||
{
|
||
delta_ready[id] = 1;
|
||
queued[id] = 0;
|
||
}
|
||
pending.clear();
|
||
};
|
||
|
||
auto note_gate_inputs = [&](const gate & g) {
|
||
if (g.kind == gate_kind::product)
|
||
{
|
||
for (auto id : g.lhs)
|
||
note(id);
|
||
}
|
||
else if (g.kind == gate_kind::dot)
|
||
{
|
||
for (auto id : g.lhs)
|
||
note(id);
|
||
for (auto id : g.rhs)
|
||
note(id);
|
||
}
|
||
else if (g.kind == gate_kind::poly)
|
||
{
|
||
for (const auto & step : g.steps)
|
||
{
|
||
if (step.value_wire >= 0)
|
||
continue;
|
||
for (auto [id, exp] : step.delta)
|
||
{
|
||
(void)exp;
|
||
note(id);
|
||
}
|
||
if (step.mask_wire >= 0)
|
||
note(static_cast<std::uint32_t>(step.mask_wire));
|
||
}
|
||
}
|
||
else
|
||
{
|
||
for (auto id : g.lhs)
|
||
note(id);
|
||
}
|
||
};
|
||
|
||
auto ensure_ready = [&](std::uint32_t id, int round) {
|
||
if (delta_ready[id])
|
||
return;
|
||
if (!lambda_ready[id] || !value_ready[id])
|
||
throw std::logic_error("beaver wire is not ready to open");
|
||
record_flush(round);
|
||
note(id);
|
||
record_flush(round);
|
||
};
|
||
|
||
auto ensure_gate_inputs = [&](const gate & g, int round) {
|
||
if (g.kind == gate_kind::product)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_ready(id, round);
|
||
}
|
||
else if (g.kind == gate_kind::dot)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_ready(id, round);
|
||
for (auto id : g.rhs)
|
||
ensure_ready(id, round);
|
||
}
|
||
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 (!value_ready[id])
|
||
throw std::logic_error("beaver wire is not ready to open");
|
||
continue;
|
||
}
|
||
for (auto [id, exp] : step.delta)
|
||
{
|
||
(void)exp;
|
||
ensure_ready(id, round);
|
||
}
|
||
if (step.mask_wire >= 0)
|
||
ensure_ready(static_cast<std::uint32_t>(step.mask_wire), round);
|
||
}
|
||
}
|
||
else
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_ready(id, round);
|
||
}
|
||
};
|
||
|
||
for (int r = 1; r <= max_round; ++r)
|
||
{
|
||
pending.clear();
|
||
for (auto & q : queued)
|
||
q = 0;
|
||
for (const auto & g : gates_)
|
||
{
|
||
if (wires_[g.out].ready_round != r)
|
||
continue;
|
||
if (value_ready[g.out])
|
||
continue;
|
||
note_gate_inputs(g);
|
||
}
|
||
record_flush(r);
|
||
|
||
for (std::size_t gi = 0; gi < gates_.size(); ++gi)
|
||
{
|
||
const auto & g = gates_[gi];
|
||
if (wires_[g.out].ready_round != r)
|
||
continue;
|
||
if (value_ready[g.out])
|
||
continue;
|
||
ensure_gate_inputs(g, r);
|
||
batch_op ev;
|
||
ev.type = batch_op::kind::eval_gate;
|
||
ev.ready_round = r;
|
||
ev.gate_index = gi;
|
||
ops.push_back(std::move(ev));
|
||
value_ready[g.out] = 1;
|
||
if (lambda_ready[g.out])
|
||
note(g.out);
|
||
}
|
||
record_flush(r);
|
||
}
|
||
return ops;
|
||
}
|
||
|
||
std::vector<delta_barrier> compile_batch_barriers() const
|
||
{
|
||
std::vector<delta_barrier> out;
|
||
for (const auto & op : compile_batch_ops())
|
||
{
|
||
if (op.type != batch_op::kind::flush)
|
||
continue;
|
||
delta_barrier b;
|
||
b.ready_round = op.ready_round;
|
||
b.wire_ids = op.wire_ids;
|
||
out.push_back(std::move(b));
|
||
}
|
||
return out;
|
||
}
|
||
|
||
void eval_gate_party(const gate & g)
|
||
{
|
||
Ring val;
|
||
if (g.kind == gate_kind::product)
|
||
val = eval_product_party(g.lhs);
|
||
else if (g.kind == gate_kind::dot)
|
||
val = eval_dot_party(g);
|
||
else if (g.kind == gate_kind::poly)
|
||
val = eval_poly_party(g);
|
||
else
|
||
val = eval_mux_party(g);
|
||
wires_[g.out].value = party_id_ == 0
|
||
? split<Ring>{val, traits::zero()}
|
||
: split<Ring>{traits::zero(), val};
|
||
wires_[g.out].value_ready = true;
|
||
}
|
||
|
||
template <typename BatchExchange>
|
||
void flush_delta_batch(std::vector<std::uint32_t> & ids,
|
||
std::vector<char> & queued, BatchExchange && exchange_deltas)
|
||
{
|
||
if (ids.empty())
|
||
return;
|
||
std::vector<Ring> mine;
|
||
mine.reserve(ids.size());
|
||
for (auto id : ids)
|
||
{
|
||
auto & w = wires_[id];
|
||
if (w.delta_ready)
|
||
throw std::logic_error("beaver delta batch contains an opened wire");
|
||
if (!w.lambda_ready || !w.value_ready)
|
||
throw std::logic_error("beaver wire is not ready to open");
|
||
mine.push_back(traits::add(side_of(w.value), side_of(w.lambda)));
|
||
}
|
||
std::vector<Ring> peer = exchange_deltas(mine);
|
||
if (peer.size() != mine.size())
|
||
throw std::runtime_error("beaver delta batch size mismatch");
|
||
for (std::size_t i = 0; i < ids.size(); ++i)
|
||
{
|
||
auto & w = wires_[ids[i]];
|
||
w.delta = traits::add(mine[i], peer[i]);
|
||
w.delta_ready = true;
|
||
queued[ids[i]] = 0;
|
||
}
|
||
ids.clear();
|
||
}
|
||
|
||
template <typename BatchExchange>
|
||
void evaluate_party_batch_impl(BatchExchange && exchange_deltas)
|
||
{
|
||
const auto ops = compile_batch_ops();
|
||
std::vector<char> queued(wires_.size(), 0);
|
||
for (const auto & op : ops)
|
||
{
|
||
if (op.type == batch_op::kind::flush)
|
||
{
|
||
std::vector<std::uint32_t> ids = op.wire_ids;
|
||
for (auto id : ids)
|
||
queued[id] = 1;
|
||
flush_delta_batch(ids, queued, exchange_deltas);
|
||
}
|
||
else
|
||
{
|
||
const auto & g = gates_[op.gate_index];
|
||
if (wires_[g.out].value_ready)
|
||
continue;
|
||
eval_gate_party(g);
|
||
}
|
||
}
|
||
}
|
||
|
||
template <typename Exchange>
|
||
void ensure_delta_party_auth(std::uint32_t id, Exchange && exchange_delta)
|
||
{
|
||
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");
|
||
auth_opening<Ring> mine{
|
||
traits::add(side_of(w.value), side_of(w.lambda)),
|
||
traits::add(side_of(w.value_tag), side_of(w.lambda_tag))};
|
||
auth_opening<Ring> peer = exchange_delta(mine);
|
||
w.delta = traits::add(mine.value, peer.value);
|
||
w.delta_mine = mine;
|
||
w.delta_peer = peer;
|
||
w.delta_open_ready = true;
|
||
w.delta_ready = true;
|
||
}
|
||
|
||
template <typename Exchange>
|
||
void evaluate_party_auth_impl(Exchange && exchange_delta)
|
||
{
|
||
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;
|
||
Ring val;
|
||
if (g.kind == gate_kind::product)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta_party_auth(id, exchange_delta);
|
||
val = eval_product_party(g.lhs);
|
||
}
|
||
else if (g.kind == gate_kind::dot)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta_party_auth(id, exchange_delta);
|
||
for (auto id : g.rhs)
|
||
ensure_delta_party_auth(id, exchange_delta);
|
||
val = eval_dot_party(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_party_auth(id, exchange_delta);
|
||
}
|
||
if (step.mask_wire >= 0)
|
||
ensure_delta_party_auth(
|
||
static_cast<std::uint32_t>(step.mask_wire), exchange_delta);
|
||
}
|
||
val = eval_poly_party(g);
|
||
}
|
||
else
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_delta_party_auth(id, exchange_delta);
|
||
val = eval_mux_party(g);
|
||
}
|
||
wires_[g.out].value = party_id_ == 0
|
||
? split<Ring>{val, traits::zero()}
|
||
: split<Ring>{traits::zero(), val};
|
||
wires_[g.out].value_ready = true;
|
||
if (wires_[g.out].lambda_ready)
|
||
ensure_delta_party_auth(g.out, exchange_delta);
|
||
}
|
||
}
|
||
}
|
||
|
||
template <typename BatchExchange>
|
||
void flush_delta_batch_auth(std::vector<std::uint32_t> & ids,
|
||
std::vector<char> & queued, BatchExchange && exchange_deltas)
|
||
{
|
||
if (ids.empty())
|
||
return;
|
||
std::vector<auth_opening<Ring>> mine;
|
||
mine.reserve(ids.size());
|
||
for (auto id : ids)
|
||
{
|
||
auto & w = wires_[id];
|
||
if (w.delta_ready)
|
||
throw std::logic_error("beaver delta batch contains an opened wire");
|
||
if (!w.lambda_ready || !w.value_ready)
|
||
throw std::logic_error("beaver wire is not ready to open");
|
||
mine.push_back(auth_opening<Ring>{
|
||
traits::add(side_of(w.value), side_of(w.lambda)),
|
||
traits::add(side_of(w.value_tag), side_of(w.lambda_tag))});
|
||
}
|
||
std::vector<auth_opening<Ring>> peer = exchange_deltas(mine);
|
||
if (peer.size() != mine.size())
|
||
throw std::runtime_error("beaver auth delta batch size mismatch");
|
||
for (std::size_t i = 0; i < ids.size(); ++i)
|
||
{
|
||
auto & w = wires_[ids[i]];
|
||
w.delta = traits::add(mine[i].value, peer[i].value);
|
||
w.delta_mine = mine[i];
|
||
w.delta_peer = peer[i];
|
||
w.delta_open_ready = true;
|
||
w.delta_ready = true;
|
||
queued[ids[i]] = 0;
|
||
}
|
||
ids.clear();
|
||
}
|
||
|
||
template <typename BatchExchange>
|
||
void evaluate_party_auth_batch_impl(BatchExchange && exchange_deltas)
|
||
{
|
||
int max_round = 0;
|
||
for (const auto & w : wires_)
|
||
max_round = std::max(max_round, w.ready_round);
|
||
std::vector<char> queued(wires_.size(), 0);
|
||
std::vector<std::uint32_t> pending;
|
||
pending.reserve(wires_.size());
|
||
|
||
auto note = [&](std::uint32_t id) {
|
||
auto & w = wires_[id];
|
||
if (w.delta_ready || queued[id])
|
||
return;
|
||
if (!w.lambda_ready || !w.value_ready)
|
||
return;
|
||
queued[id] = 1;
|
||
pending.push_back(id);
|
||
};
|
||
|
||
auto note_gate_inputs = [&](const gate & g) {
|
||
if (g.kind == gate_kind::product)
|
||
{
|
||
for (auto id : g.lhs)
|
||
note(id);
|
||
}
|
||
else if (g.kind == gate_kind::dot)
|
||
{
|
||
for (auto id : g.lhs)
|
||
note(id);
|
||
for (auto id : g.rhs)
|
||
note(id);
|
||
}
|
||
else if (g.kind == gate_kind::poly)
|
||
{
|
||
for (const auto & step : g.steps)
|
||
{
|
||
if (step.value_wire >= 0)
|
||
continue;
|
||
for (auto [id, exp] : step.delta)
|
||
{
|
||
(void)exp;
|
||
note(id);
|
||
}
|
||
if (step.mask_wire >= 0)
|
||
note(static_cast<std::uint32_t>(step.mask_wire));
|
||
}
|
||
}
|
||
else
|
||
{
|
||
for (auto id : g.lhs)
|
||
note(id);
|
||
}
|
||
};
|
||
|
||
auto ensure_ready = [&](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");
|
||
flush_delta_batch_auth(pending, queued, exchange_deltas);
|
||
note(id);
|
||
flush_delta_batch_auth(pending, queued, exchange_deltas);
|
||
};
|
||
|
||
auto ensure_gate_inputs = [&](const gate & g) {
|
||
if (g.kind == gate_kind::product)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_ready(id);
|
||
}
|
||
else if (g.kind == gate_kind::dot)
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_ready(id);
|
||
for (auto id : g.rhs)
|
||
ensure_ready(id);
|
||
}
|
||
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_ready(id);
|
||
}
|
||
if (step.mask_wire >= 0)
|
||
ensure_ready(static_cast<std::uint32_t>(step.mask_wire));
|
||
}
|
||
}
|
||
else
|
||
{
|
||
for (auto id : g.lhs)
|
||
ensure_ready(id);
|
||
}
|
||
};
|
||
|
||
for (int r = 1; r <= max_round; ++r)
|
||
{
|
||
pending.clear();
|
||
for (auto & q : queued)
|
||
q = 0;
|
||
for (const auto & g : gates_)
|
||
{
|
||
if (wires_[g.out].ready_round != r)
|
||
continue;
|
||
if (wires_[g.out].value_ready)
|
||
continue;
|
||
note_gate_inputs(g);
|
||
}
|
||
flush_delta_batch_auth(pending, queued, exchange_deltas);
|
||
|
||
for (const auto & g : gates_)
|
||
{
|
||
if (wires_[g.out].ready_round != r)
|
||
continue;
|
||
if (wires_[g.out].value_ready)
|
||
continue;
|
||
ensure_gate_inputs(g);
|
||
Ring val;
|
||
if (g.kind == gate_kind::product)
|
||
val = eval_product_party(g.lhs);
|
||
else if (g.kind == gate_kind::dot)
|
||
val = eval_dot_party(g);
|
||
else if (g.kind == gate_kind::poly)
|
||
val = eval_poly_party(g);
|
||
else
|
||
val = eval_mux_party(g);
|
||
wires_[g.out].value = party_id_ == 0
|
||
? split<Ring>{val, traits::zero()}
|
||
: split<Ring>{traits::zero(), val};
|
||
wires_[g.out].value_ready = true;
|
||
if (wires_[g.out].lambda_ready)
|
||
note(g.out);
|
||
}
|
||
flush_delta_batch_auth(pending, queued, exchange_deltas);
|
||
}
|
||
}
|
||
|
||
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));
|
||
|
||
// Latency path: keep one interactive round (Pika / online Grotto).
|
||
// Factor peels that buy prep at the cost of an extra round are skipped.
|
||
if (schedule_objective_ == schedule_objective::rounds)
|
||
return emit_terms(std::move(terms));
|
||
|
||
std::vector<std::vector<poly_term>> best{terms};
|
||
std::size_t best_prep = estimate_pieces(best);
|
||
std::size_t best_pieces = best.size();
|
||
auto consider = [&](std::vector<std::vector<poly_term>> seq) {
|
||
const std::size_t prep = estimate_pieces(seq);
|
||
const std::size_t pieces = seq.size();
|
||
// Prep first (Appendix E): fewer blinds/bundles, then fewer rounds.
|
||
const bool take = (prep < best_prep)
|
||
|| (prep == best_prep && pieces < best_pieces);
|
||
if (take)
|
||
{
|
||
best = std::move(seq);
|
||
best_prep = prep;
|
||
best_pieces = pieces;
|
||
}
|
||
};
|
||
|
||
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;
|
||
}
|
||
|
||
template <typename Exchange>
|
||
void ensure_delta_party(std::uint32_t id, Exchange && exchange_delta)
|
||
{
|
||
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");
|
||
Ring mine = traits::add(side_of(w.value), side_of(w.lambda));
|
||
Ring peer = exchange_delta(mine);
|
||
w.delta = traits::add(mine, peer);
|
||
w.delta_ready = true;
|
||
}
|
||
|
||
struct located_term
|
||
{
|
||
const split<Ring> * share = nullptr;
|
||
const split<Ring> * tag = nullptr;
|
||
};
|
||
|
||
located_term locate_term(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 located_term{&w.lambda, &w.lambda_tag};
|
||
}
|
||
auto it = mono_index_.find(key);
|
||
if (it != mono_index_.end() && monos_[it->second].ready)
|
||
return located_term{&monos_[it->second].share, &monos_[it->second].tag};
|
||
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 located_term{&b.share, &b.tag};
|
||
}
|
||
throw std::logic_error("beaver monomial was not prepared");
|
||
}
|
||
|
||
split<Ring> term_share(const exp_list & key) const
|
||
{
|
||
return *locate_term(key).share;
|
||
}
|
||
|
||
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};
|
||
}
|
||
|
||
Ring eval_product_party(const std::vector<std::uint32_t> & factors) const
|
||
{
|
||
auto groups = group_exponents(factors);
|
||
Ring s = 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)
|
||
{
|
||
if (party_id_ == 0)
|
||
s = traits::add(s, scale_int(pub, coeff));
|
||
return;
|
||
}
|
||
auto sh = term_share(key);
|
||
s = traits::add(s, scale_int(traits::mul(pub, side_of(sh)), coeff));
|
||
});
|
||
return s;
|
||
}
|
||
|
||
Ring eval_dot_party(const gate & g) const
|
||
{
|
||
if (!g.cross_ready)
|
||
throw std::logic_error("beaver dot cross term is not sampled");
|
||
Ring s = 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]];
|
||
if (party_id_ == 0)
|
||
s = traits::add(s, traits::mul(x.delta, y.delta));
|
||
s = traits::sub(s, traits::mul(x.delta, side_of(y.lambda)));
|
||
s = traits::sub(s, traits::mul(y.delta, side_of(x.lambda)));
|
||
}
|
||
s = traits::add(s, side_of(g.cross));
|
||
return s;
|
||
}
|
||
|
||
Ring eval_poly_party(const gate & g) const
|
||
{
|
||
Ring s = 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;
|
||
s = traits::add(s, traits::mul(step.scale, side_of(val)));
|
||
continue;
|
||
}
|
||
Ring pub = pow_delta(step.delta);
|
||
if (step.public_only)
|
||
{
|
||
if (party_id_ == 0)
|
||
s = traits::add(s, traits::mul(pub, step.scale));
|
||
continue;
|
||
}
|
||
if (step.mask_wire >= 0)
|
||
{
|
||
const auto & lam = wires_[static_cast<std::size_t>(step.mask_wire)].lambda;
|
||
s = traits::add(s, traits::mul(pub, traits::mul(step.scale, side_of(lam))));
|
||
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");
|
||
s = traits::add(s, traits::mul(pub, traits::mul(step.scale, side_of(share))));
|
||
}
|
||
return s;
|
||
}
|
||
|
||
Ring eval_mux_party(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);
|
||
Ring ld = traits::sub(side_of(x.lambda), side_of(y.lambda));
|
||
auto bx = term_share(pair_key(g.lhs[0], g.lhs[1]));
|
||
auto by = term_share(pair_key(g.lhs[0], g.lhs[2]));
|
||
Ring cross = traits::sub(side_of(bx), side_of(by));
|
||
Ring s = traits::zero();
|
||
if (party_id_ == 0)
|
||
s = traits::mul(b.delta, dd);
|
||
s = traits::sub(s, traits::mul(b.delta, ld));
|
||
s = traits::sub(s, traits::mul(dd, side_of(b.lambda)));
|
||
s = traits::add(s, cross);
|
||
s = traits::add(s, side_of(y.value));
|
||
return s;
|
||
}
|
||
};
|
||
|
||
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 static_cast<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;
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator*(expr<Ring> a, wire<Ring> b)
|
||
{
|
||
return detail::mul_exprs(a, wire_expr(b));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator*(wire<Ring> a, expr<Ring> b)
|
||
{
|
||
return detail::mul_exprs(wire_expr(a), b);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator*(expr<Ring> a, expr<Ring> b)
|
||
{
|
||
return detail::mul_exprs(a, b);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator+(wire<Ring> a, expr<Ring> b)
|
||
{
|
||
return detail::add_exprs(wire_expr(a), b);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator+(expr<Ring> a, expr<Ring> b)
|
||
{
|
||
return detail::add_exprs(std::move(a), b);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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()));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator-(wire<Ring> w)
|
||
{
|
||
return -wire_expr(w);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator-(expr<Ring> a, expr<Ring> b)
|
||
{
|
||
return std::move(a) + (-std::move(b));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator-(expr<Ring> a, wire<Ring> b)
|
||
{
|
||
return std::move(a) + (-b);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator-(wire<Ring> a, expr<Ring> b)
|
||
{
|
||
return wire_expr(a) + (-std::move(b));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
expr<Ring> operator-(wire<Ring> a, wire<Ring> b)
|
||
{
|
||
return wire_expr(a) + (-b);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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;
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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));
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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;
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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);
|
||
}
|
||
|
||
/// \complexity O(1) to record the gate on the session. No messages. Opening is `evaluate` / `evaluate_party`.
|
||
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{};
|
||
};
|
||
|
||
/// @brief Drive `evaluate_party_batch` one δ barrier at a time.
|
||
/// @details Construct after `install_party` and binding inputs. Each
|
||
/// `take_local` / `apply_peer` pair matches one entry of
|
||
/// `session::exchange_barriers()`. Gate evaluations between barriers
|
||
/// run inside `apply_peer`. Wire bytes match `evaluate_party_batch`.
|
||
template <typename Ring>
|
||
class party_batch_stepper
|
||
{
|
||
public:
|
||
using traits = ring_traits<Ring>;
|
||
|
||
explicit party_batch_stepper(session<Ring> & s)
|
||
: s_(&s),
|
||
ops_(s.compile_batch_ops()),
|
||
pc_(0),
|
||
waiting_(false)
|
||
{
|
||
if (!s.party_mode_)
|
||
throw std::logic_error("call install_party before party_batch_stepper");
|
||
advance_to_flush_or_end();
|
||
}
|
||
|
||
std::size_t barrier_count() const
|
||
{
|
||
std::size_t n = 0;
|
||
for (const auto & op : ops_)
|
||
{
|
||
if (op.type == session<Ring>::batch_op::kind::flush)
|
||
++n;
|
||
}
|
||
return n;
|
||
}
|
||
|
||
bool done() const noexcept
|
||
{
|
||
return pc_ >= ops_.size() && !waiting_;
|
||
}
|
||
|
||
/// @brief Local δ shares for the next barrier (same order as `barrier().wire_ids`).
|
||
std::vector<Ring> take_local()
|
||
{
|
||
if (done())
|
||
throw std::logic_error("party_batch_stepper is done");
|
||
if (waiting_)
|
||
throw std::logic_error("party_batch_stepper already waiting for peer");
|
||
if (pc_ >= ops_.size()
|
||
|| ops_[pc_].type != session<Ring>::batch_op::kind::flush)
|
||
throw std::logic_error("party_batch_stepper expected a flush");
|
||
const auto & ids = ops_[pc_].wire_ids;
|
||
current_.ready_round = ops_[pc_].ready_round;
|
||
current_.wire_ids = ids;
|
||
std::vector<Ring> mine;
|
||
mine.reserve(ids.size());
|
||
for (auto id : ids)
|
||
{
|
||
auto & w = s_->wires_[id];
|
||
if (w.delta_ready)
|
||
throw std::logic_error("beaver delta batch contains an opened wire");
|
||
if (!w.lambda_ready || !w.value_ready)
|
||
throw std::logic_error("beaver wire is not ready to open");
|
||
mine.push_back(traits::add(s_->side_of(w.value), s_->side_of(w.lambda)));
|
||
}
|
||
waiting_ = true;
|
||
return mine;
|
||
}
|
||
|
||
const typename session<Ring>::delta_barrier & barrier() const
|
||
{
|
||
if (!waiting_)
|
||
throw std::logic_error("party_batch_stepper::barrier requires take_local");
|
||
return current_;
|
||
}
|
||
|
||
/// @brief Apply peer shares and evaluate until the next flush or completion.
|
||
void apply_peer(const std::vector<Ring> & peer)
|
||
{
|
||
if (!waiting_)
|
||
throw std::logic_error("party_batch_stepper::apply_peer without take_local");
|
||
const auto & ids = current_.wire_ids;
|
||
if (peer.size() != ids.size())
|
||
throw std::runtime_error("beaver delta batch size mismatch");
|
||
for (std::size_t i = 0; i < ids.size(); ++i)
|
||
{
|
||
auto & w = s_->wires_[ids[i]];
|
||
Ring mine = traits::add(s_->side_of(w.value), s_->side_of(w.lambda));
|
||
w.delta = traits::add(mine, peer[i]);
|
||
w.delta_ready = true;
|
||
}
|
||
waiting_ = false;
|
||
++pc_;
|
||
advance_to_flush_or_end();
|
||
}
|
||
|
||
private:
|
||
void advance_to_flush_or_end()
|
||
{
|
||
while (pc_ < ops_.size()
|
||
&& ops_[pc_].type == session<Ring>::batch_op::kind::eval_gate)
|
||
{
|
||
const auto & g = s_->gates_[ops_[pc_].gate_index];
|
||
if (!s_->wires_[g.out].value_ready)
|
||
s_->eval_gate_party(g);
|
||
++pc_;
|
||
}
|
||
}
|
||
|
||
session<Ring> * s_;
|
||
std::vector<typename session<Ring>::batch_op> ops_;
|
||
std::size_t pc_ = 0;
|
||
bool waiting_ = false;
|
||
typename session<Ring>::delta_barrier current_{};
|
||
};
|
||
|
||
/// @brief Authenticated twin of `party_batch_stepper` (`auth_opening` per wire).
|
||
template <typename Ring>
|
||
class party_auth_batch_stepper
|
||
{
|
||
public:
|
||
using traits = ring_traits<Ring>;
|
||
|
||
explicit party_auth_batch_stepper(session<Ring> & s)
|
||
: s_(&s),
|
||
ops_(s.compile_batch_ops()),
|
||
pc_(0),
|
||
waiting_(false)
|
||
{
|
||
if (!s.party_mode_)
|
||
throw std::logic_error(
|
||
"call install_party before party_auth_batch_stepper");
|
||
if (!s.has_mac_)
|
||
throw std::logic_error(
|
||
"party_auth_batch_stepper requires an authenticated tape");
|
||
advance_to_flush_or_end();
|
||
}
|
||
|
||
bool done() const noexcept
|
||
{
|
||
return pc_ >= ops_.size() && !waiting_;
|
||
}
|
||
|
||
std::vector<auth_opening<Ring>> take_local()
|
||
{
|
||
if (done())
|
||
throw std::logic_error("party_auth_batch_stepper is done");
|
||
if (waiting_)
|
||
throw std::logic_error("party_auth_batch_stepper already waiting");
|
||
if (pc_ >= ops_.size()
|
||
|| ops_[pc_].type != session<Ring>::batch_op::kind::flush)
|
||
throw std::logic_error("party_auth_batch_stepper expected a flush");
|
||
const auto & ids = ops_[pc_].wire_ids;
|
||
current_.ready_round = ops_[pc_].ready_round;
|
||
current_.wire_ids = ids;
|
||
mine_.clear();
|
||
mine_.reserve(ids.size());
|
||
for (auto id : ids)
|
||
{
|
||
auto & w = s_->wires_[id];
|
||
if (w.delta_ready)
|
||
throw std::logic_error("beaver delta batch contains an opened wire");
|
||
if (!w.lambda_ready || !w.value_ready)
|
||
throw std::logic_error("beaver wire is not ready to open");
|
||
mine_.push_back(auth_opening<Ring>{
|
||
traits::add(s_->side_of(w.value), s_->side_of(w.lambda)),
|
||
traits::add(s_->side_of(w.value_tag), s_->side_of(w.lambda_tag))});
|
||
}
|
||
waiting_ = true;
|
||
return mine_;
|
||
}
|
||
|
||
void apply_peer(const std::vector<auth_opening<Ring>> & peer)
|
||
{
|
||
if (!waiting_)
|
||
throw std::logic_error(
|
||
"party_auth_batch_stepper::apply_peer without take_local");
|
||
const auto & ids = current_.wire_ids;
|
||
if (peer.size() != ids.size() || peer.size() != mine_.size())
|
||
throw std::runtime_error("beaver auth delta batch size mismatch");
|
||
for (std::size_t i = 0; i < ids.size(); ++i)
|
||
{
|
||
auto & w = s_->wires_[ids[i]];
|
||
w.delta = traits::add(mine_[i].value, peer[i].value);
|
||
w.delta_mine = mine_[i];
|
||
w.delta_peer = peer[i];
|
||
w.delta_open_ready = true;
|
||
w.delta_ready = true;
|
||
if (s_->mac_key_known_
|
||
&& !verify_auth_opening(w.delta_mine, w.delta_peer, s_->mac_key_))
|
||
throw std::runtime_error("beaver auth opening failed");
|
||
}
|
||
waiting_ = false;
|
||
++pc_;
|
||
advance_to_flush_or_end();
|
||
}
|
||
|
||
private:
|
||
void advance_to_flush_or_end()
|
||
{
|
||
while (pc_ < ops_.size()
|
||
&& ops_[pc_].type == session<Ring>::batch_op::kind::eval_gate)
|
||
{
|
||
const auto & g = s_->gates_[ops_[pc_].gate_index];
|
||
if (!s_->wires_[g.out].value_ready)
|
||
s_->eval_gate_party(g);
|
||
++pc_;
|
||
}
|
||
}
|
||
|
||
session<Ring> * s_;
|
||
std::vector<typename session<Ring>::batch_op> ops_;
|
||
std::size_t pc_ = 0;
|
||
bool waiting_ = false;
|
||
typename session<Ring>::delta_barrier current_{};
|
||
std::vector<auth_opening<Ring>> mine_;
|
||
};
|
||
|
||
|
||
|
||
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{};
|
||
};
|
||
|
||
template <typename Ring, typename Apply>
|
||
void with_beaver2_formula(Apply && apply)
|
||
{
|
||
session<Ring> s;
|
||
auto a = s.input();
|
||
auto b = s.input();
|
||
auto out = s(a * b);
|
||
s.pin(out);
|
||
apply(s, a, b, out);
|
||
}
|
||
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
beaver2<Ring> beaver2_read(session<Ring> & s,
|
||
typename session<Ring>::wire a,
|
||
typename session<Ring>::wire b,
|
||
typename session<Ring>::wire out)
|
||
{
|
||
return beaver2<Ring>{
|
||
s.lambda(a),
|
||
s.lambda(b),
|
||
s.monomial({{a, 1u}, {b, 1u}}),
|
||
s.lambda(out)};
|
||
}
|
||
|
||
/// @brief One fresh Beaver pair from copy `index` of an oracle.
|
||
/// @details Records the same `a * b` formula as `sample_beaver2` and samples it
|
||
/// with `session::sample_from`.
|
||
/// @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)
|
||
{
|
||
beaver2<Ring> got;
|
||
with_beaver2_formula<Ring>([&](auto & s, auto a, auto b, auto out) {
|
||
s.sample_from(src, index);
|
||
got = beaver2_read(s, a, b, out);
|
||
});
|
||
return got;
|
||
}
|
||
|
||
/// @brief `n` copies starting at `begin`, each one `sample_from` of the pair formula.
|
||
/// @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");
|
||
with_beaver2_formula<Ring>([&](auto & s, auto a, auto b, auto product) {
|
||
for (std::size_t i = 0; i < n; ++i)
|
||
{
|
||
if (i != 0)
|
||
s.clear_sample();
|
||
s.sample_from(src, begin + static_cast<std::uint64_t>(i));
|
||
out[i] = beaver2_read(s, a, b, product);
|
||
}
|
||
});
|
||
}
|
||
|
||
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>{});
|
||
}
|
||
|
||
/// @brief Beaver triple with IT-MAC tags on every share.
|
||
/// @tparam Ring payload ring
|
||
template <typename Ring>
|
||
struct auth_beaver2
|
||
{
|
||
auth_split<Ring> a{};
|
||
auth_split<Ring> b{};
|
||
auth_split<Ring> ab{};
|
||
auth_split<Ring> out{};
|
||
|
||
HEDLEY_NO_THROW
|
||
bool verify(const dpf::mac_key<Ring> & key) const
|
||
{
|
||
return a.verify(key) && b.verify(key) && ab.verify(key) && out.verify(key)
|
||
&& ab.open() == ring_traits<Ring>::mul(a.open(), b.open());
|
||
}
|
||
};
|
||
|
||
template <typename Ring>
|
||
struct auth_beaver3
|
||
{
|
||
auth_split<Ring> a{};
|
||
auth_split<Ring> b{};
|
||
auth_split<Ring> c{};
|
||
auth_split<Ring> ab{};
|
||
auth_split<Ring> ac{};
|
||
auth_split<Ring> bc{};
|
||
auth_split<Ring> abc{};
|
||
auth_split<Ring> out{};
|
||
|
||
HEDLEY_NO_THROW
|
||
bool verify(const dpf::mac_key<Ring> & key) const
|
||
{
|
||
using traits = ring_traits<Ring>;
|
||
return a.verify(key) && b.verify(key) && c.verify(key)
|
||
&& ab.verify(key) && ac.verify(key) && bc.verify(key)
|
||
&& abc.verify(key) && out.verify(key)
|
||
&& ab.open() == traits::mul(a.open(), b.open())
|
||
&& ac.open() == traits::mul(a.open(), c.open())
|
||
&& bc.open() == traits::mul(b.open(), c.open())
|
||
&& abc.open() == traits::mul(ab.open(), c.open());
|
||
}
|
||
};
|
||
|
||
/// @brief Authenticate an existing beaver2 under `key`.
|
||
template <typename Ring, typename Sample>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver2<Ring> authenticate_beaver2(const beaver2<Ring> & t,
|
||
const dpf::mac_key<Ring> & key, Sample && rng)
|
||
{
|
||
auto & sampler = rng;
|
||
return auth_beaver2<Ring>{
|
||
authenticate(t.a, key, sampler),
|
||
authenticate(t.b, key, sampler),
|
||
authenticate(t.ab, key, sampler),
|
||
authenticate(t.out, key, sampler)};
|
||
}
|
||
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver2<Ring> authenticate_beaver2(const beaver2<Ring> & t,
|
||
const dpf::mac_key<Ring> & key)
|
||
{
|
||
return authenticate_beaver2(t, key, default_sampler<Ring>{});
|
||
}
|
||
|
||
template <typename Ring, typename Sample>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver2<Ring> sample_auth_beaver2(const dpf::mac_key<Ring> & key,
|
||
Sample && rng)
|
||
{
|
||
auto & sampler = rng;
|
||
return authenticate_beaver2(sample_beaver2<Ring>(sampler), key, sampler);
|
||
}
|
||
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver2<Ring> sample_auth_beaver2(const dpf::mac_key<Ring> & key)
|
||
{
|
||
return sample_auth_beaver2<Ring>(key, default_sampler<Ring>{});
|
||
}
|
||
|
||
template <typename Ring, typename Sample>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver3<Ring> authenticate_beaver3(const beaver3<Ring> & t,
|
||
const dpf::mac_key<Ring> & key, Sample && rng)
|
||
{
|
||
auto & sampler = rng;
|
||
return auth_beaver3<Ring>{
|
||
authenticate(t.a, key, sampler),
|
||
authenticate(t.b, key, sampler),
|
||
authenticate(t.c, key, sampler),
|
||
authenticate(t.ab, key, sampler),
|
||
authenticate(t.ac, key, sampler),
|
||
authenticate(t.bc, key, sampler),
|
||
authenticate(t.abc, key, sampler),
|
||
authenticate(t.out, key, sampler)};
|
||
}
|
||
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver3<Ring> authenticate_beaver3(const beaver3<Ring> & t,
|
||
const dpf::mac_key<Ring> & key)
|
||
{
|
||
return authenticate_beaver3(t, key, default_sampler<Ring>{});
|
||
}
|
||
|
||
template <typename Ring, typename Sample>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver3<Ring> sample_auth_beaver3(const dpf::mac_key<Ring> & key,
|
||
Sample && rng)
|
||
{
|
||
auto & sampler = rng;
|
||
return authenticate_beaver3(sample_beaver3<Ring>(sampler), key, sampler);
|
||
}
|
||
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver3<Ring> sample_auth_beaver3(const dpf::mac_key<Ring> & key)
|
||
{
|
||
return sample_auth_beaver3<Ring>(key, default_sampler<Ring>{});
|
||
}
|
||
|
||
/// @brief Classic Beaver multiply of authenticated inputs with an auth triple.
|
||
/// @details Opens `d = x-a`, `e = y-b`, then
|
||
/// `[xy] = [ab] + d[b] + e[a] + de` on values and tags. The public
|
||
/// `de` value and `de·Δ` tag land on party 1 (same side as the clear
|
||
/// beaver product).
|
||
/// @return Party shares of the authenticated product (no ABY2 out mask).
|
||
template <typename Ring>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
std::pair<dpf::mac_share<Ring>, dpf::mac_share<Ring>> auth_beaver_mul(
|
||
dpf::mac_share<Ring> x0, dpf::mac_share<Ring> x1,
|
||
dpf::mac_share<Ring> y0, dpf::mac_share<Ring> y1,
|
||
const auth_beaver2<Ring> & bev, const dpf::mac_key<Ring> & key)
|
||
{
|
||
using traits = ring_traits<Ring>;
|
||
const Ring x = traits::add(x0.value, x1.value);
|
||
const Ring y = traits::add(y0.value, y1.value);
|
||
const Ring d = traits::sub(x, bev.a.open());
|
||
const Ring e = traits::sub(y, bev.b.open());
|
||
const auto a0 = bev.a.party(0);
|
||
const auto a1 = bev.a.party(1);
|
||
const auto b0 = bev.b.party(0);
|
||
const auto b1 = bev.b.party(1);
|
||
const auto ab0 = bev.ab.party(0);
|
||
const auto ab1 = bev.ab.party(1);
|
||
const Ring de = traits::mul(d, e);
|
||
dpf::mac_share<Ring> z0{
|
||
traits::add(ab0.value,
|
||
traits::add(traits::mul(d, b0.value), traits::mul(e, a0.value))),
|
||
traits::add(ab0.tag,
|
||
traits::add(traits::mul(d, b0.tag), traits::mul(e, a0.tag)))};
|
||
dpf::mac_share<Ring> z1{
|
||
traits::add(ab1.value,
|
||
traits::add(traits::mul(d, b1.value),
|
||
traits::add(traits::mul(e, a1.value), de))),
|
||
traits::add(ab1.tag,
|
||
traits::add(traits::mul(d, b1.tag),
|
||
traits::add(traits::mul(e, a1.tag),
|
||
traits::mul(de, key.delta))))};
|
||
return {z0, z1};
|
||
}
|
||
|
||
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.resize(n);
|
||
t.y.resize(n);
|
||
for (std::size_t i = 0; i < n; ++i)
|
||
{
|
||
t.x[i] = s.lambda(x[i]);
|
||
t.y[i] = 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 Concrete, typename LeafT, typename Sample>
|
||
void fill_wildcard_scale_blinds(Concrete & out0, Concrete & out1, LeafT & vec0,
|
||
LeafT & vec1, std::size_t nlanes, Sample && rng)
|
||
{
|
||
auto sc = sample_scale<Concrete>(nlanes, std::forward<Sample>(rng));
|
||
out0 = sc.scalar.p0;
|
||
out1 = sc.scalar.p1;
|
||
vec0 = LeafT{};
|
||
vec1 = LeafT{};
|
||
if constexpr (utils::is_packed_subbyte_v<Concrete>)
|
||
{
|
||
for (std::size_t li = 0; li < sc.lanes.size(); ++li)
|
||
{
|
||
packed::deposit_lane(vec0, li, sc.lanes[li].p0);
|
||
packed::deposit_lane(vec1, li, sc.lanes[li].p1);
|
||
}
|
||
}
|
||
else
|
||
{
|
||
for (std::size_t li = 0; li < sc.lanes.size(); ++li)
|
||
{
|
||
const auto off = li * sizeof(Concrete);
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wstringop-overflow")
|
||
std::memcpy(reinterpret_cast<unsigned char *>(std::addressof(vec0)) + off,
|
||
std::addressof(sc.lanes[li].p0), sizeof(Concrete));
|
||
std::memcpy(reinterpret_cast<unsigned char *>(std::addressof(vec1)) + off,
|
||
std::addressof(sc.lanes[li].p1), sizeof(Concrete));
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
}
|
||
}
|
||
}
|
||
|
||
template <typename Concrete, typename LeafT>
|
||
void fill_wildcard_scale_blinds(Concrete & out0, Concrete & out1, LeafT & vec0,
|
||
LeafT & vec1, std::size_t nlanes)
|
||
{
|
||
fill_wildcard_scale_blinds(out0, out1, vec0, vec1, nlanes,
|
||
default_sampler<Concrete>{});
|
||
}
|
||
|
||
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>{});
|
||
}
|
||
|
||
/// @brief XOR share of a bit and an additive share of that same bit in `uint64_t`.
|
||
/// @details Both shares are bound on a session. `pad.bit()` chooses the bit and
|
||
/// fills the XOR-ring sampler; `pad.block()` fills the integer sampler.
|
||
struct bit_arith_shares
|
||
{
|
||
std::uint8_t xor0 = 0;
|
||
std::uint8_t xor1 = 0;
|
||
std::uint64_t add0 = 0;
|
||
std::uint64_t add1 = 0;
|
||
};
|
||
|
||
template <typename Pad>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
bit_arith_shares sample_bit_arith(Pad & pad)
|
||
{
|
||
using bit_ring = dpf::xor_wrapper<std::uint8_t>;
|
||
using btraits = ring_traits<bit_ring>;
|
||
const std::uint8_t bit = static_cast<std::uint8_t>(pad.bit() & 1u);
|
||
auto bits = [&pad]() -> bit_ring {
|
||
std::uint8_t packed = 0;
|
||
for (int i = 0; i < 8; ++i)
|
||
{
|
||
packed = static_cast<std::uint8_t>(
|
||
packed | (static_cast<std::uint8_t>(pad.bit() & 1u) << i));
|
||
}
|
||
return bit_ring{packed};
|
||
};
|
||
auto words = [&pad]() -> std::uint64_t {
|
||
const auto block = pad.block();
|
||
std::uint64_t word = 0;
|
||
std::memcpy(&word, &block, sizeof(word));
|
||
return word;
|
||
};
|
||
|
||
session<bit_ring> xb;
|
||
auto bw = xb.bit();
|
||
xb.pin(bw);
|
||
xb.sample(bits);
|
||
xb.bind(bw, bit ? btraits::one() : btraits::zero(), bits);
|
||
const auto xs = xb.value(bw);
|
||
|
||
session<std::uint64_t> ab;
|
||
auto aw = ab.input();
|
||
ab.pin(aw);
|
||
ab.sample(words);
|
||
ab.bind(aw, std::uint64_t{bit}, words);
|
||
const auto as = ab.value(aw);
|
||
|
||
auto low = [](bit_ring x) {
|
||
return static_cast<std::uint8_t>(static_cast<std::uint8_t>(x) & 1u);
|
||
};
|
||
return bit_arith_shares{low(xs.p0), low(xs.p1), as.p0, as.p1};
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
beaver2<Ring> sample_beaver2(const oracle<Ring, PRG> & src, std::uint64_t index)
|
||
{
|
||
return beaver2_at<Ring, PRG>(src, index);
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
beaver3<Ring> sample_beaver3(const oracle<Ring, PRG> & src, std::uint64_t index)
|
||
{
|
||
session<Ring> s;
|
||
auto a = s.input();
|
||
auto b = s.input();
|
||
auto c = s.input();
|
||
auto out = s(a * b * c);
|
||
s.pin(out);
|
||
s.sample_from(src, index);
|
||
return beaver3<Ring>{
|
||
s.lambda(a), s.lambda(b), s.lambda(c),
|
||
s.monomial({{a, 1u}, {b, 1u}}),
|
||
s.monomial({{a, 1u}, {c, 1u}}),
|
||
s.monomial({{b, 1u}, {c, 1u}}),
|
||
s.monomial({{a, 1u}, {b, 1u}, {c, 1u}}),
|
||
s.lambda(out)};
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver2<Ring> sample_auth_beaver2(const dpf::mac_key<Ring> & key,
|
||
const oracle<Ring, PRG> & src, std::uint64_t index)
|
||
{
|
||
auth_beaver2<Ring> got;
|
||
with_beaver2_formula<Ring>([&](auto & s, auto a, auto b, auto out) {
|
||
s.set_mac_key(key);
|
||
s.sample_from(src, index);
|
||
got = auth_beaver2<Ring>{
|
||
s.lambda_auth(a),
|
||
s.lambda_auth(b),
|
||
s.monomial_auth({{a, 1u}, {b, 1u}}),
|
||
s.lambda_auth(out)};
|
||
});
|
||
return got;
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auth_beaver3<Ring> sample_auth_beaver3(const dpf::mac_key<Ring> & key,
|
||
const oracle<Ring, PRG> & src, std::uint64_t index)
|
||
{
|
||
session<Ring> s;
|
||
s.set_mac_key(key);
|
||
auto a = s.input();
|
||
auto b = s.input();
|
||
auto c = s.input();
|
||
auto out = s(a * b * c);
|
||
s.pin(out);
|
||
s.sample_from(src, index);
|
||
return auth_beaver3<Ring>{
|
||
s.lambda_auth(a), s.lambda_auth(b), s.lambda_auth(c),
|
||
s.monomial_auth({{a, 1u}, {b, 1u}}),
|
||
s.monomial_auth({{a, 1u}, {c, 1u}}),
|
||
s.monomial_auth({{b, 1u}, {c, 1u}}),
|
||
s.monomial_auth({{a, 1u}, {b, 1u}, {c, 1u}}),
|
||
s.lambda_auth(out)};
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
square_beaver<Ring> sample_square(const oracle<Ring, PRG> & src, std::uint64_t index)
|
||
{
|
||
session<Ring> s;
|
||
auto x = s.input();
|
||
auto out = s(x * x);
|
||
s.pin(out);
|
||
s.sample_from(src, index);
|
||
return square_beaver<Ring>{s.lambda(x), s.monomial({{x, 2u}}), s.lambda(out)};
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
mul_square_beaver<Ring> sample_mul_square(const oracle<Ring, PRG> & src,
|
||
std::uint64_t index)
|
||
{
|
||
session<Ring> s;
|
||
auto a = s.input();
|
||
auto x = s.input();
|
||
auto out = s(a * x * x);
|
||
s.pin(out);
|
||
s.sample_from(src, index);
|
||
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, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
dot_beaver<Ring> sample_dot(std::size_t n, const oracle<Ring, PRG> & src,
|
||
std::uint64_t index)
|
||
{
|
||
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_from(src, index);
|
||
dot_beaver<Ring> t;
|
||
t.cross = s.dot_cross(out);
|
||
t.out = s.lambda(out);
|
||
t.x.resize(n);
|
||
t.y.resize(n);
|
||
for (std::size_t i = 0; i < n; ++i)
|
||
{
|
||
t.x[i] = s.lambda(x[i]);
|
||
t.y[i] = s.lambda(y[i]);
|
||
}
|
||
return t;
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
scale_beaver<Ring> sample_scale(std::size_t n, const oracle<Ring, PRG> & src,
|
||
std::uint64_t index)
|
||
{
|
||
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_from(src, index);
|
||
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 Concrete, typename LeafT, typename PRG>
|
||
void fill_wildcard_scale_blinds(Concrete & out0, Concrete & out1, LeafT & vec0,
|
||
LeafT & vec1, std::size_t nlanes, const oracle<Concrete, PRG> & src,
|
||
std::uint64_t index)
|
||
{
|
||
auto sc = sample_scale<Concrete, PRG>(nlanes, src, index);
|
||
out0 = sc.scalar.p0;
|
||
out1 = sc.scalar.p1;
|
||
vec0 = LeafT{};
|
||
vec1 = LeafT{};
|
||
if constexpr (utils::is_packed_subbyte_v<Concrete>)
|
||
{
|
||
for (std::size_t li = 0; li < sc.lanes.size(); ++li)
|
||
{
|
||
packed::deposit_lane(vec0, li, sc.lanes[li].p0);
|
||
packed::deposit_lane(vec1, li, sc.lanes[li].p1);
|
||
}
|
||
}
|
||
else
|
||
{
|
||
for (std::size_t li = 0; li < sc.lanes.size(); ++li)
|
||
{
|
||
const auto off = li * sizeof(Concrete);
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wstringop-overflow")
|
||
std::memcpy(reinterpret_cast<unsigned char *>(std::addressof(vec0)) + off,
|
||
std::addressof(sc.lanes[li].p0), sizeof(Concrete));
|
||
std::memcpy(reinterpret_cast<unsigned char *>(std::addressof(vec1)) + off,
|
||
std::addressof(sc.lanes[li].p1), sizeof(Concrete));
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
}
|
||
}
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
bit_mul_beaver<Ring> sample_bit_mul(const oracle<Ring, PRG> & src, std::uint64_t index)
|
||
{
|
||
session<Ring> s;
|
||
auto b = s.bit();
|
||
auto x = s.input();
|
||
auto out = s.bit_mul(b, x);
|
||
s.pin(out);
|
||
s.sample_from(src, index);
|
||
return bit_mul_beaver<Ring>{
|
||
s.lambda(b), s.lambda(x),
|
||
s.monomial({{b, 1u}, {x, 1u}}),
|
||
s.lambda(out)};
|
||
}
|
||
|
||
template <typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
mux_beaver<Ring> sample_mux(const oracle<Ring, PRG> & src, std::uint64_t index)
|
||
{
|
||
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_from(src, index);
|
||
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 <std::size_t Arity, typename Ring, typename PRG>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
fresh_beaver<Arity, Ring> sample_fresh(const oracle<Ring, PRG> & src,
|
||
std::uint64_t index)
|
||
{
|
||
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_from(src, index);
|
||
|
||
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;
|
||
}
|
||
|
||
} // namespace beavers
|
||
} // namespace dpf
|
||
|
||
#endif // LIBDPF_INCLUDE_DPF_BEAVER_HPP__
|