/// @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 #include #include #include #include #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(bit_mask(bits)); return dpf::uniform_sample() & m; } return dpf::uniform_sample() & 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(factor); const std::uint64_t f1 = static_cast(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(a)} * static_cast(b)) & bit_mask(bits); } const std::uint64_t a0 = static_cast(a); const std::uint64_t a1 = static_cast(a >> 64); const std::uint64_t b0 = static_cast(b); const std::uint64_t b1 = static_cast(b >> 64); const u128 p00 = u128{a0} * b0; const u128 mid = (p00 >> 64) + static_cast(u128{a0} * b1) + static_cast(u128{a1} * b0); return ((u128)static_cast(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 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 HEDLEY_WARN_UNUSED_RESULT u128 open_sum(u128 mine, unsigned bits, Exchange & exchange) { const std::uint64_t low = exchange(static_cast(mine)); u128 peer = low; if (bits > 64u) { const std::uint64_t hi = exchange(static_cast(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 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(); k.r_share[1] = k.r - k.r_share[0]; k.wrap = make_bit_key(static_cast(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 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(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(threshold) - 1u); } template 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(bit_mask(keys.src)); const std::uint64_t xs = static_cast(share) & src_mask; const std::uint64_t rs = static_cast(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 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(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 struct fixed_mul_beaver_shape { using plan = fixed_mul_plan; using result_integral = typename plan::integral_type; static constexpr unsigned storage_bits = static_cast(dpf::utils::bitlength_of_v); 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(plan::align_shift) <= 64u && plan::out_bits <= 64u && plan::modulus_bits <= 128u && plan::modulus_bits > static_cast(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 struct fixed_mul_beaver_prep { using shape = fixed_mul_beaver_shape; 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 & 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()); } /// @brief Same triple from copy `index` of a `mod2k` oracle. template HEDLEY_WARN_UNUSED_RESULT inline fixed_mul_beaver_triple sample_fixed_mul_beaver_triple(unsigned multiply_bits, const dpf::beavers::oracle & 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 HEDLEY_WARN_UNUSED_RESULT fixed_mul_beaver_prep make_fixed_mul_beaver_prep() { using shape = fixed_mul_beaver_shape; using prep = fixed_mul_beaver_prep; 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(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(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 HEDLEY_WARN_UNUSED_RESULT auto eval_fixed_mul_beaver( const fixed_mul_beaver_prep & prep, const fixed_mul_beaver_triple & triple, std::size_t party, LhsIntegral lhs_share, RhsIntegral rhs_share, Exchange && exchange) -> typename fixed_mul_plan::result_type { using shape = fixed_mul_beaver_shape; 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( 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(-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(window), static_cast(window >> 64), 0u, 0u}; return make_fixed_from_integral_type( detail::limbs_to_integral(limbs)); } } } // namespace grotto #endif // LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_BEAVER_HPP__