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>
617 lines
22 KiB
C++
617 lines
22 KiB
C++
/// @file grotto/fixedpoint_beaver.hpp
|
|
/// @brief Additive-share evaluation of `fixed_mul`.
|
|
/// @details One Beaver triple in `Z/2^{multiply_bits}Z` is the product. Every
|
|
/// other step is linear, or a masked comparison that lifts a narrower
|
|
/// share or drops low bits:
|
|
///
|
|
/// - bits above `multiply_bits` are discarded locally;
|
|
/// - a narrower operand (and a product narrower than `modulus_bits`)
|
|
/// is lifted by `eta + r - w·2^{src}`, with `w` the carry of the
|
|
/// secret mask, and a signed lift then replicates the sign;
|
|
/// - `align_shift > 0` is a truncate-and-reduce of that modulus;
|
|
/// - `align_shift < 0` is a local left shift.
|
|
///
|
|
/// The opened window matches `fixed_mul`, including the final sign
|
|
/// extension into the storage word. Each lift extends by at most 64
|
|
/// bits and starts from at most 64 bits, so the carry bit stays a
|
|
/// `uint64_t` comparison and scaling it is homomorphic in the
|
|
/// destination ring. A right shift discards at most 64 bits. The
|
|
/// modulus and the multiply ring are at most 128 bits.
|
|
///
|
|
/// `prep` is a dealer record (both key halves, both mask shares),
|
|
/// same as `carry_key_pair`. It is independent of the data and is
|
|
/// reused. `sample_fixed_mul_beaver_triple` is one product and is
|
|
/// fresh each time. `eval_fixed_mul_beaver` is one party: `exchange`
|
|
/// sends one `uint64_t` and returns the peer's word, in the same
|
|
/// order on both sides.
|
|
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
|
|
/// @license Released under a GNU General Public v2.0 (GPLv2) license.
|
|
|
|
#ifndef LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_BEAVER_HPP__
|
|
#define LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_BEAVER_HPP__
|
|
|
|
#include <cstddef>
|
|
#include <cstdint>
|
|
#include <optional>
|
|
#include <stdexcept>
|
|
#include <utility>
|
|
|
|
#include "hedley/hedley.h"
|
|
|
|
#include "dpf.hpp"
|
|
#include "grotto/fixedpoint_mul.hpp"
|
|
|
|
namespace grotto
|
|
{
|
|
|
|
namespace fixed_mul_beaver_detail
|
|
{
|
|
|
|
using u128 = simde_uint128;
|
|
|
|
HEDLEY_NO_THROW
|
|
constexpr u128 bit_mask(unsigned bits) noexcept
|
|
{
|
|
if (bits == 0u)
|
|
return 0;
|
|
if (bits >= 128u)
|
|
return ~u128{0};
|
|
return (u128{1} << bits) - 1;
|
|
}
|
|
|
|
HEDLEY_NO_THROW
|
|
constexpr unsigned u64_limbs(unsigned bits) noexcept
|
|
{
|
|
return bits > 64u ? 2u : 1u;
|
|
}
|
|
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline u128 draw_mod(unsigned bits)
|
|
{
|
|
if (bits == 0u)
|
|
return 0;
|
|
if (bits <= 64u)
|
|
{
|
|
const std::uint64_t m = static_cast<std::uint64_t>(bit_mask(bits));
|
|
return dpf::uniform_sample<std::uint64_t>() & m;
|
|
}
|
|
return dpf::uniform_sample<u128>() & bit_mask(bits);
|
|
}
|
|
|
|
/// @brief Low 128 bits of `(uint64)m * factor`, then reduced mod `2^dest`.
|
|
HEDLEY_NO_THROW
|
|
inline u128 mul_u64(std::uint64_t m, u128 factor, unsigned dest) noexcept
|
|
{
|
|
const std::uint64_t f0 = static_cast<std::uint64_t>(factor);
|
|
const std::uint64_t f1 = static_cast<std::uint64_t>(factor >> 64);
|
|
const u128 p0 = u128{m} * f0;
|
|
const u128 p1 = u128{m} * f1;
|
|
return (p0 + (p1 << 64)) & bit_mask(dest);
|
|
}
|
|
|
|
HEDLEY_NO_THROW
|
|
inline u128 mul_mod(u128 a, u128 b, unsigned bits) noexcept
|
|
{
|
|
a &= bit_mask(bits);
|
|
b &= bit_mask(bits);
|
|
if (bits <= 64u)
|
|
{
|
|
return (u128{static_cast<std::uint64_t>(a)}
|
|
* static_cast<std::uint64_t>(b)) & bit_mask(bits);
|
|
}
|
|
const std::uint64_t a0 = static_cast<std::uint64_t>(a);
|
|
const std::uint64_t a1 = static_cast<std::uint64_t>(a >> 64);
|
|
const std::uint64_t b0 = static_cast<std::uint64_t>(b);
|
|
const std::uint64_t b1 = static_cast<std::uint64_t>(b >> 64);
|
|
const u128 p00 = u128{a0} * b0;
|
|
const u128 mid = (p00 >> 64)
|
|
+ static_cast<std::uint64_t>(u128{a0} * b1)
|
|
+ static_cast<std::uint64_t>(u128{a1} * b0);
|
|
return ((u128)static_cast<std::uint64_t>(p00) | (mid << 64)) & bit_mask(bits);
|
|
}
|
|
|
|
HEDLEY_NO_THROW
|
|
inline u128 add_mod(u128 a, u128 b, unsigned bits) noexcept
|
|
{
|
|
return (a + b) & bit_mask(bits);
|
|
}
|
|
|
|
HEDLEY_NO_THROW
|
|
inline u128 sub_mod(u128 a, u128 b, unsigned bits) noexcept
|
|
{
|
|
return (a - b) & bit_mask(bits);
|
|
}
|
|
|
|
template <typename T>
|
|
HEDLEY_NO_THROW
|
|
u128 share_bits(const T & value) noexcept
|
|
{
|
|
std::uint64_t raw[4] = {};
|
|
detail::store_raw_limbs(value, raw);
|
|
return u128{raw[0]} | (u128{raw[1]} << 64);
|
|
}
|
|
|
|
using cmp_pair = decltype(dpf::make_dpf(std::uint64_t{0},
|
|
dpf::lt(std::uint64_t{1})));
|
|
|
|
struct bit_key
|
|
{
|
|
cmp_pair keys;
|
|
};
|
|
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline bit_key make_bit_key(std::uint64_t alpha)
|
|
{
|
|
return bit_key{dpf::make_dpf(alpha, dpf::lt(std::uint64_t{1}))};
|
|
}
|
|
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline std::uint64_t eval_bit(const bit_key & key, std::size_t party,
|
|
std::uint64_t query)
|
|
{
|
|
if (party == 0u)
|
|
return dpf::eval_point(dpf::cmp, key.keys.first, query).raw();
|
|
return dpf::eval_point(dpf::cmp, key.keys.second, query).raw();
|
|
}
|
|
|
|
/// @brief Public sum of the two parties' words, mod `2^bits`.
|
|
template <typename Exchange>
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
u128 open_sum(u128 mine, unsigned bits, Exchange & exchange)
|
|
{
|
|
const std::uint64_t low = exchange(static_cast<std::uint64_t>(mine));
|
|
u128 peer = low;
|
|
if (bits > 64u)
|
|
{
|
|
const std::uint64_t hi = exchange(static_cast<std::uint64_t>(mine >> 64));
|
|
peer |= u128{hi} << 64;
|
|
}
|
|
return (mine + peer) & bit_mask(bits);
|
|
}
|
|
|
|
struct lift_keys
|
|
{
|
|
bool live = false;
|
|
bool sign = false;
|
|
unsigned src = 0;
|
|
unsigned dest = 0;
|
|
/// @brief Secret mask mod `2^src`. Shares sum to `r` in `Z/2^128`.
|
|
u128 r = 0;
|
|
u128 r_share[2]{};
|
|
/// @brief `lt` at `r`. Eval at `T-1` is a share of `1{r >= T}`.
|
|
std::optional<bit_key> wrap{};
|
|
};
|
|
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline lift_keys make_lift(unsigned src, unsigned dest, bool sign)
|
|
{
|
|
lift_keys k;
|
|
k.live = true;
|
|
k.sign = sign;
|
|
k.src = src;
|
|
k.dest = dest;
|
|
k.r = draw_mod(src);
|
|
k.r_share[0] = dpf::uniform_sample<u128>();
|
|
k.r_share[1] = k.r - k.r_share[0];
|
|
k.wrap = make_bit_key(static_cast<std::uint64_t>(k.r));
|
|
return k;
|
|
}
|
|
|
|
struct shift_keys
|
|
{
|
|
bool live = false;
|
|
unsigned n = 0;
|
|
unsigned s = 0;
|
|
u128 rin = 0;
|
|
u128 rin_share[2]{};
|
|
u128 rout_share[2]{};
|
|
std::optional<bit_key> low{};
|
|
};
|
|
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline shift_keys make_shift(unsigned n, unsigned s)
|
|
{
|
|
shift_keys k;
|
|
k.live = true;
|
|
k.n = n;
|
|
k.s = s;
|
|
k.rin = draw_mod(n);
|
|
k.rin_share[0] = draw_mod(n);
|
|
k.rin_share[1] = (k.rin - k.rin_share[0]) & bit_mask(n);
|
|
const u128 neg = (u128{0} - k.rin) & bit_mask(n);
|
|
const std::uint64_t alpha = static_cast<std::uint64_t>(neg & bit_mask(s));
|
|
k.low = make_bit_key(alpha);
|
|
const unsigned out = n - s;
|
|
const u128 y_hi = neg >> s;
|
|
k.rout_share[0] = draw_mod(out);
|
|
k.rout_share[1] = (y_hi - k.rout_share[0]) & bit_mask(out);
|
|
return k;
|
|
}
|
|
|
|
/// @brief Share of `1{r >= T}` for `r` in `[0, 2^src)`. `T == 0` is the public 1.
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline std::uint64_t ge_mask(const lift_keys & keys, std::size_t party,
|
|
u128 threshold)
|
|
{
|
|
if (threshold == 0)
|
|
return party == 1u ? std::uint64_t{1} : 0u;
|
|
if (threshold > bit_mask(keys.src))
|
|
return 0u;
|
|
return eval_bit(*keys.wrap, party,
|
|
static_cast<std::uint64_t>(threshold) - 1u);
|
|
}
|
|
|
|
template <typename Exchange>
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
u128 apply_lift(u128 share, const lift_keys & keys, std::size_t party,
|
|
Exchange & exchange)
|
|
{
|
|
const std::uint64_t src_mask = static_cast<std::uint64_t>(bit_mask(keys.src));
|
|
const std::uint64_t xs = static_cast<std::uint64_t>(share) & src_mask;
|
|
const std::uint64_t rs = static_cast<std::uint64_t>(keys.r_share[party]) & src_mask;
|
|
const std::uint64_t delta = (xs - rs) & src_mask;
|
|
const std::uint64_t peer = exchange(delta);
|
|
const std::uint64_t eta = (delta + peer) & src_mask;
|
|
// `1{r >= 2^src - eta}` is the carry `r + eta >= 2^src`.
|
|
const u128 mod = u128{1} << keys.src;
|
|
const std::uint64_t w = ge_mask(keys, party, mod - eta);
|
|
u128 y = keys.r_share[party];
|
|
y -= mul_u64(w, mod, keys.dest);
|
|
if (party == 0u)
|
|
y += eta;
|
|
if (keys.sign)
|
|
{
|
|
const u128 half = u128{1} << (keys.src - 1u);
|
|
std::uint64_t msb = 0;
|
|
if (u128{eta} < half)
|
|
msb = ge_mask(keys, party, half - eta) - w;
|
|
else
|
|
msb = (party == 1u ? std::uint64_t{1} : 0u) - w
|
|
+ ge_mask(keys, party, mod + half - eta);
|
|
const u128 high = bit_mask(keys.dest) & ~bit_mask(keys.src);
|
|
y += mul_u64(msb, high, keys.dest);
|
|
}
|
|
return y;
|
|
}
|
|
|
|
template <typename Exchange>
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
u128 apply_shift(u128 share, const shift_keys & keys, std::size_t party,
|
|
Exchange & exchange)
|
|
{
|
|
const u128 opened = open_sum(share + keys.rin_share[party], keys.n, exchange);
|
|
const u128 xs = opened & bit_mask(keys.s);
|
|
const u128 query128 = (u128{1} << keys.s) - xs - 1;
|
|
const std::uint64_t t = eval_bit(*keys.low, party,
|
|
static_cast<std::uint64_t>(query128));
|
|
// `t` sums to the carry in `Z/2^64`. The output modulus is at most 64
|
|
// bits for every window this header accepts, so that extra multiple of
|
|
// `2^64` lands outside the modulus.
|
|
u128 acc = keys.rout_share[party] + t;
|
|
if (party == 1u)
|
|
acc += opened >> keys.s;
|
|
return acc;
|
|
}
|
|
|
|
} // namespace fixed_mul_beaver_detail
|
|
|
|
/// @brief Which steps of one `fixed_mul` window are interactive.
|
|
/// @tparam IntegerBits integer bits kept in the product, including the sign
|
|
/// @tparam FractionalBits fraction bits kept in the product
|
|
/// @tparam LhsFractionalBits fractional bits of the left operand
|
|
/// @tparam LhsIntegral integral type of the left operand
|
|
/// @tparam RhsFractionalBits fractional bits of the right operand
|
|
/// @tparam RhsIntegral integral type of the right operand
|
|
template <unsigned IntegerBits,
|
|
unsigned FractionalBits,
|
|
unsigned LhsFractionalBits,
|
|
typename LhsIntegral,
|
|
unsigned RhsFractionalBits,
|
|
typename RhsIntegral>
|
|
struct fixed_mul_beaver_shape
|
|
{
|
|
using plan = fixed_mul_plan<IntegerBits, FractionalBits,
|
|
LhsFractionalBits, LhsIntegral, RhsFractionalBits, RhsIntegral>;
|
|
using result_integral = typename plan::integral_type;
|
|
|
|
static constexpr unsigned storage_bits =
|
|
static_cast<unsigned>(dpf::utils::bitlength_of_v<result_integral>);
|
|
static constexpr bool active = plan::modulus_bits > 0u
|
|
&& plan::multiply_bits > 0u;
|
|
static constexpr bool lhs_lift = active
|
|
&& plan::multiply_bits > plan::lhs_width;
|
|
static constexpr bool rhs_lift = active
|
|
&& plan::multiply_bits > plan::rhs_width;
|
|
static constexpr bool product_lift = active
|
|
&& plan::modulus_bits > plan::multiply_bits;
|
|
static constexpr bool shift_right = active && plan::align_shift > 0;
|
|
static constexpr bool shift_left = active && plan::align_shift < 0;
|
|
static constexpr bool result_lift = active
|
|
&& storage_bits > plan::out_bits;
|
|
|
|
static constexpr bool lift_ok(bool lift, unsigned src, unsigned dest) noexcept
|
|
{
|
|
if (!lift)
|
|
return true;
|
|
return src >= 1u && src <= 64u && dest > src && dest <= 128u
|
|
&& (dest - src) <= 64u;
|
|
}
|
|
|
|
static constexpr bool fits = !active
|
|
|| (plan::multiply_bits <= 128u
|
|
&& plan::modulus_bits <= 128u
|
|
&& storage_bits <= 128u
|
|
&& lift_ok(lhs_lift, plan::lhs_width, plan::multiply_bits)
|
|
&& lift_ok(rhs_lift, plan::rhs_width, plan::multiply_bits)
|
|
&& lift_ok(product_lift, plan::multiply_bits, plan::modulus_bits)
|
|
&& lift_ok(result_lift, plan::out_bits, storage_bits)
|
|
&& (!shift_right || (static_cast<unsigned>(plan::align_shift) <= 64u
|
|
&& plan::out_bits <= 64u
|
|
&& plan::modulus_bits <= 128u
|
|
&& plan::modulus_bits > static_cast<unsigned>(plan::align_shift))));
|
|
|
|
/// @brief `uint64_t` words exchanged by one product. Zero when the window is empty.
|
|
static constexpr unsigned messages = !active ? 0u
|
|
: (lhs_lift ? 1u : 0u)
|
|
+ (rhs_lift ? 1u : 0u)
|
|
+ 2u * fixed_mul_beaver_detail::u64_limbs(plan::multiply_bits)
|
|
+ (product_lift ? 1u : 0u)
|
|
+ (shift_right
|
|
? fixed_mul_beaver_detail::u64_limbs(plan::modulus_bits)
|
|
: 0u)
|
|
+ (result_lift ? 1u : 0u);
|
|
};
|
|
|
|
/// @brief One Beaver triple in `Z/2^{multiply_bits}Z`.
|
|
struct fixed_mul_beaver_triple
|
|
{
|
|
unsigned multiply_bits = 0;
|
|
fixed_mul_beaver_detail::u128 a[2]{};
|
|
fixed_mul_beaver_detail::u128 b[2]{};
|
|
fixed_mul_beaver_detail::u128 ab[2]{};
|
|
};
|
|
|
|
/// @brief Dealer masks and comparison keys for one window. Reused across products.
|
|
template <unsigned IntegerBits,
|
|
unsigned FractionalBits,
|
|
unsigned LhsFractionalBits,
|
|
typename LhsIntegral,
|
|
unsigned RhsFractionalBits,
|
|
typename RhsIntegral>
|
|
struct fixed_mul_beaver_prep
|
|
{
|
|
using shape = fixed_mul_beaver_shape<IntegerBits, FractionalBits,
|
|
LhsFractionalBits, LhsIntegral, RhsFractionalBits, RhsIntegral>;
|
|
|
|
fixed_mul_beaver_detail::lift_keys lhs{};
|
|
fixed_mul_beaver_detail::lift_keys rhs{};
|
|
fixed_mul_beaver_detail::lift_keys product{};
|
|
fixed_mul_beaver_detail::shift_keys shift{};
|
|
fixed_mul_beaver_detail::lift_keys result{};
|
|
};
|
|
|
|
/// @brief Pack a `mod2k` beaver pair into the fixed-point triple.
|
|
inline fixed_mul_beaver_triple pack_mod2k_beaver(
|
|
unsigned multiply_bits, const dpf::beavers::beaver2<dpf::mod2k> & triple)
|
|
{
|
|
fixed_mul_beaver_triple t;
|
|
t.multiply_bits = multiply_bits;
|
|
t.a[0] = triple.a.p0.raw;
|
|
t.a[1] = triple.a.p1.raw;
|
|
t.b[0] = triple.b.p0.raw;
|
|
t.b[1] = triple.b.p1.raw;
|
|
t.ab[0] = triple.ab.p0.raw;
|
|
t.ab[1] = triple.ab.p1.raw;
|
|
return t;
|
|
}
|
|
|
|
/// @brief Sample one product triple from `sample_beaver2` in `Z/2^w Z`.
|
|
/// @param multiply_bits ring width, `0` or `1 .. 128`
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline fixed_mul_beaver_triple sample_fixed_mul_beaver_triple(unsigned multiply_bits)
|
|
{
|
|
if (multiply_bits == 0u)
|
|
return {};
|
|
if (multiply_bits > 128u)
|
|
throw std::invalid_argument("fixed_mul beaver triple exceeds 128 bits");
|
|
dpf::beavers::mod2k_width_scope width(multiply_bits);
|
|
return pack_mod2k_beaver(multiply_bits,
|
|
dpf::beavers::sample_beaver2<dpf::mod2k>());
|
|
}
|
|
|
|
/// @brief Same triple from copy `index` of a `mod2k` oracle.
|
|
template <typename PRG = dpf::prg::aes128>
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
inline fixed_mul_beaver_triple sample_fixed_mul_beaver_triple(unsigned multiply_bits,
|
|
const dpf::beavers::oracle<dpf::mod2k, PRG> & src, std::uint64_t index = 0)
|
|
{
|
|
if (multiply_bits == 0u)
|
|
return {};
|
|
if (multiply_bits > 128u)
|
|
throw std::invalid_argument("fixed_mul beaver triple exceeds 128 bits");
|
|
dpf::beavers::mod2k_width_scope width(multiply_bits);
|
|
return pack_mod2k_beaver(multiply_bits,
|
|
dpf::beavers::sample_beaver2(src, index));
|
|
}
|
|
|
|
/// @brief Build the reusable comparisons for this window.
|
|
/// \complexity One `dpf::make_dpf` per live lift (operand, product, result)
|
|
/// and one for a right shift. A signed lift reuses that key at a second
|
|
/// query. Each comparison is a `uint64_t` domain. No messages.
|
|
/// \rounds No party interaction.
|
|
/// \communication None.
|
|
/// \preprocessing Those comparison keys and the mask shares.
|
|
template <unsigned IntegerBits,
|
|
unsigned FractionalBits,
|
|
unsigned LhsFractionalBits,
|
|
typename LhsIntegral,
|
|
unsigned RhsFractionalBits,
|
|
typename RhsIntegral>
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
fixed_mul_beaver_prep<IntegerBits, FractionalBits, LhsFractionalBits, LhsIntegral,
|
|
RhsFractionalBits, RhsIntegral>
|
|
make_fixed_mul_beaver_prep()
|
|
{
|
|
using shape = fixed_mul_beaver_shape<IntegerBits, FractionalBits,
|
|
LhsFractionalBits, LhsIntegral, RhsFractionalBits, RhsIntegral>;
|
|
using prep = fixed_mul_beaver_prep<IntegerBits, FractionalBits,
|
|
LhsFractionalBits, LhsIntegral, RhsFractionalBits, RhsIntegral>;
|
|
static_assert(shape::fits,
|
|
"fixed_mul beaver: multiply and modulus must be at most 128 bits, "
|
|
"each lift must start from at most 64 bits and extend by at most 64, "
|
|
"and a right shift must discard at most 64 bits");
|
|
prep out;
|
|
if constexpr (!shape::active)
|
|
return out;
|
|
using plan = typename shape::plan;
|
|
if constexpr (shape::lhs_lift)
|
|
{
|
|
out.lhs = fixed_mul_beaver_detail::make_lift(plan::lhs_width,
|
|
plan::multiply_bits, plan::lhs_signed);
|
|
}
|
|
if constexpr (shape::rhs_lift)
|
|
{
|
|
out.rhs = fixed_mul_beaver_detail::make_lift(plan::rhs_width,
|
|
plan::multiply_bits, plan::rhs_signed);
|
|
}
|
|
if constexpr (shape::product_lift)
|
|
{
|
|
out.product = fixed_mul_beaver_detail::make_lift(plan::multiply_bits,
|
|
plan::modulus_bits, plan::operands_signed);
|
|
}
|
|
if constexpr (shape::shift_right)
|
|
{
|
|
out.shift = fixed_mul_beaver_detail::make_shift(plan::modulus_bits,
|
|
static_cast<unsigned>(plan::align_shift));
|
|
}
|
|
if constexpr (shape::result_lift)
|
|
{
|
|
out.result = fixed_mul_beaver_detail::make_lift(plan::out_bits,
|
|
shape::storage_bits, plan::result_is_signed);
|
|
}
|
|
return out;
|
|
}
|
|
|
|
/// @brief One party's share of `fixed_mul<IntegerBits, FractionalBits>(lhs, rhs)`.
|
|
/// @details `lhs_share` and `rhs_share` are additive shares of each operand's
|
|
/// raw integral word, in that word's own ring. The return value is
|
|
/// this party's share of the product's raw integral word. `exchange(mine)`
|
|
/// returns the peer's matching `uint64_t`. Both parties must call it
|
|
/// the same number of times (`fixed_mul_beaver_shape::messages`).
|
|
/// @tparam Exchange callable `std::uint64_t(std::uint64_t)`
|
|
/// @param prep reusable dealer material from `make_fixed_mul_beaver_prep`
|
|
/// @param triple fresh Beaver triple in the multiply ring
|
|
/// @param party `0` or `1`
|
|
/// @param lhs_share this party's share of the left raw word
|
|
/// @param rhs_share this party's share of the right raw word
|
|
/// @param exchange peer exchange for one `uint64_t`
|
|
/// @return this party's share of the product, as a fixed-point word
|
|
/// \complexity The Beaver product is `O(1)` 128-bit arithmetic. Each live
|
|
/// lift or shift is one to three `eval_point` calls on a `uint64_t` DCF.
|
|
/// \rounds One round if the caller pipelines every `exchange`; the callback
|
|
/// itself is one word at a time. `messages` words in total.
|
|
/// \communication `fixed_mul_beaver_shape::messages` words of `uint64_t`.
|
|
/// \preprocessing `prep` and one `triple`.
|
|
template <unsigned IntegerBits,
|
|
unsigned FractionalBits,
|
|
unsigned LhsFractionalBits,
|
|
typename LhsIntegral,
|
|
unsigned RhsFractionalBits,
|
|
typename RhsIntegral,
|
|
typename Exchange>
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|
auto eval_fixed_mul_beaver(
|
|
const fixed_mul_beaver_prep<IntegerBits, FractionalBits, LhsFractionalBits,
|
|
LhsIntegral, RhsFractionalBits, RhsIntegral> & prep,
|
|
const fixed_mul_beaver_triple & triple,
|
|
std::size_t party,
|
|
LhsIntegral lhs_share,
|
|
RhsIntegral rhs_share,
|
|
Exchange && exchange)
|
|
-> typename fixed_mul_plan<IntegerBits, FractionalBits, LhsFractionalBits,
|
|
LhsIntegral, RhsFractionalBits, RhsIntegral>::result_type
|
|
{
|
|
using shape = fixed_mul_beaver_shape<IntegerBits, FractionalBits,
|
|
LhsFractionalBits, LhsIntegral, RhsFractionalBits, RhsIntegral>;
|
|
using plan = typename shape::plan;
|
|
using integral = typename plan::integral_type;
|
|
static_assert(shape::fits,
|
|
"fixed_mul beaver: multiply and modulus must be at most 128 bits, "
|
|
"each lift must start from at most 64 bits and extend by at most 64, "
|
|
"and a right shift must discard at most 64 bits");
|
|
if (party > 1u)
|
|
throw std::invalid_argument("fixed_mul beaver party is 0 or 1");
|
|
|
|
using namespace fixed_mul_beaver_detail;
|
|
const auto zero = make_fixed_from_integral_type<FractionalBits, integral>(
|
|
integral{});
|
|
if constexpr (!shape::active)
|
|
{
|
|
(void)prep;
|
|
(void)triple;
|
|
(void)lhs_share;
|
|
(void)rhs_share;
|
|
(void)exchange;
|
|
return zero;
|
|
}
|
|
else
|
|
{
|
|
if (triple.multiply_bits != plan::multiply_bits)
|
|
throw std::invalid_argument("fixed_mul beaver triple width does not match the window");
|
|
|
|
auto reduce = [&](u128 share, unsigned width, const lift_keys & lift,
|
|
bool do_lift) -> u128 {
|
|
const unsigned kept = width < 128u ? width : 128u;
|
|
u128 limb = share & bit_mask(kept);
|
|
if (!do_lift)
|
|
return limb & bit_mask(plan::multiply_bits);
|
|
return apply_lift(limb, lift, party, exchange);
|
|
};
|
|
|
|
const u128 left = reduce(share_bits(lhs_share), plan::lhs_width,
|
|
prep.lhs, shape::lhs_lift);
|
|
const u128 right = reduce(share_bits(rhs_share), plan::rhs_width,
|
|
prep.rhs, shape::rhs_lift);
|
|
|
|
const unsigned m = plan::multiply_bits;
|
|
const u128 d = open_sum(sub_mod(left, triple.a[party], m), m, exchange);
|
|
const u128 e = open_sum(sub_mod(right, triple.b[party], m), m, exchange);
|
|
u128 prod = add_mod(triple.ab[party],
|
|
add_mod(mul_mod(d, triple.b[party], m),
|
|
mul_mod(e, triple.a[party], m), m), m);
|
|
if (party == 1u)
|
|
prod = add_mod(prod, mul_mod(d, e, m), m);
|
|
|
|
u128 wide = prod;
|
|
if constexpr (shape::product_lift)
|
|
wide = apply_lift(wide, prep.product, party, exchange);
|
|
|
|
u128 window = wide;
|
|
if constexpr (shape::shift_right)
|
|
window = apply_shift(wide, prep.shift, party, exchange);
|
|
else if constexpr (shape::shift_left)
|
|
{
|
|
constexpr unsigned k = static_cast<unsigned>(-plan::align_shift);
|
|
window = ((wide & bit_mask(plan::modulus_bits)) << k)
|
|
& bit_mask(plan::out_bits);
|
|
}
|
|
else
|
|
window = wide & bit_mask(plan::out_bits);
|
|
|
|
if constexpr (shape::result_lift)
|
|
window = apply_lift(window, prep.result, party, exchange);
|
|
|
|
window &= bit_mask(shape::storage_bits);
|
|
std::uint64_t limbs[4] = {
|
|
static_cast<std::uint64_t>(window),
|
|
static_cast<std::uint64_t>(window >> 64),
|
|
0u,
|
|
0u};
|
|
return make_fixed_from_integral_type<FractionalBits, integral>(
|
|
detail::limbs_to_integral<integral>(limbs));
|
|
}
|
|
}
|
|
|
|
} // namespace grotto
|
|
|
|
#endif // LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_BEAVER_HPP__
|