/// @file grotto/fixedpoint_mul.hpp /// @brief Fixed-point product with a caller-chosen integer and fraction width. /// @details The plaintext operation an ABY2.0 / Beaver multiplier reproduces. /// Included from grotto/fixedpoint.hpp. #ifndef LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_MUL_HPP__ #define LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_MUL_HPP__ #include "hedley/hedley.h" #ifndef LIBDPF_INCLUDE_DPF_FIXEDPOINT_HPP__ #include "grotto/fixedpoint.hpp" #endif #include #include #include #include namespace grotto { /// @brief Shape of `fixed_mul`: the ring width and the public shift. /// /// The product of the raw integers is taken modulo `2^multiply_bits`, which is /// the narrowest ring whose low bits contain the requested window. Bits of /// that product at and above `align_shift` are the output; bits below it are /// the discarded fraction, and bits above the window are the discarded /// integer part. Both discards are the low or high residue modulo a power of /// two, so a negative value is floored onto the output ulp. /// /// A Beaver triple replaces only the multiply in `Z/2^multiply_bits Z`. /// Reducing each operand into that ring is local when it is a truncation or a /// zero-extend. Sign-extending a narrower signed operand, and replicating the /// product sign when `modulus_bits > multiply_bits`, are plaintext steps the /// MPC protocol has to reproduce (they are not local on additive shares). /// @tparam IntegerBits number of integer bits, including the sign /// @tparam FractionalBits number of fractional bits /// @tparam LhsFractionalBits lhs fractional bits /// @tparam LhsIntegral lhs integral /// @tparam RhsFractionalBits rhs fractional bits /// @tparam RhsIntegral rhs integral template struct fixed_mul_plan { static constexpr unsigned integer_bits = IntegerBits; static constexpr unsigned fractional_bits = FractionalBits; static constexpr unsigned out_bits = IntegerBits + FractionalBits; static constexpr unsigned lhs_width = dpf::utils::bitlength_of_v; static constexpr unsigned rhs_width = dpf::utils::bitlength_of_v; static constexpr bool lhs_signed = std::is_signed_v; static constexpr bool rhs_signed = std::is_signed_v; static constexpr bool operands_signed = lhs_signed || rhs_signed; /// @brief Right shift applied to the raw product. Negative means a left shift. static constexpr int align_shift = static_cast(LhsFractionalBits) + static_cast(RhsFractionalBits) - static_cast(FractionalBits); /// @brief Bits of the product that the shift reads. Zero when a left shift /// moves every product bit out of the output. static constexpr int modulus_bits_signed = align_shift >= 0 ? align_shift + static_cast(out_bits) : static_cast(out_bits) + align_shift; static constexpr unsigned modulus_bits = modulus_bits_signed > 0 ? static_cast(modulus_bits_signed) : 0u; /// @brief Full two's-complement product fits in this many bits. static constexpr unsigned product_bits = lhs_width + rhs_width; static constexpr unsigned multiply_bits = modulus_bits < product_bits ? modulus_bits : product_bits; static constexpr unsigned limbs = multiply_bits == 0u ? 0u : (multiply_bits + 63u) / 64u; /// @brief Signed storage exists through 128 bits. A wider window is the same /// residue held in an unsigned fixed-point. static constexpr bool result_is_signed = operands_signed && out_bits <= 128u; static_assert(out_bits >= 1u && out_bits <= 256u, "fixed-point product window must be between 1 and 256 bits"); static_assert(modulus_bits <= 768u, "fixed-point product window exceeds 768 bits"); static_assert(limbs <= 8u, "fixed-point multiply uses at most 8 limbs"); private: template struct storage { static constexpr unsigned width = Bits <= 8u ? 8u : Bits <= 16u ? 16u : Bits <= 32u ? 32u : Bits <= 64u ? 64u : Bits <= 128u ? 128u : 256u; using type = std::conditional_t>>>, dpf::utils::integral_type_from_bitlength_t>; }; public: using integral_type = typename storage::type; using result_type = fixedpoint; }; namespace detail { inline constexpr std::size_t fixed_mul_buf_limbs = 12; HEDLEY_NO_THROW constexpr void mask_to_bits(std::uint64_t * limbs, std::size_t nlimbs, unsigned bits) noexcept { if (bits >= nlimbs * 64u) { return; } const unsigned limb = bits / 64u; const unsigned rem = bits % 64u; if (rem == 0u) { for (std::size_t i = limb; i < nlimbs; ++i) { limbs[i] = 0; } return; } limbs[limb] &= (std::uint64_t{1} << rem) - 1u; for (std::size_t i = limb + 1; i < nlimbs; ++i) { limbs[i] = 0; } } HEDLEY_NO_THROW constexpr bool test_bit(const std::uint64_t * limbs, unsigned bit) noexcept { return ((limbs[bit / 64u] >> (bit % 64u)) & 1u) != 0u; } HEDLEY_NO_THROW constexpr void fill_ones(std::uint64_t * limbs, unsigned from, unsigned to) noexcept { for (unsigned bit = from; bit < to; ) { const unsigned limb = bit / 64u; const unsigned rem = bit % 64u; const unsigned count = std::min(64u - rem, to - bit); const std::uint64_t ones = count == 64u ? ~std::uint64_t{0} : (std::uint64_t{1} << count) - 1u; limbs[limb] |= ones << rem; bit += count; } } HEDLEY_NO_THROW constexpr void sign_extend_range(std::uint64_t * limbs, unsigned from_bits, unsigned to_bits) noexcept { if (to_bits <= from_bits || from_bits == 0u) { return; } if (test_bit(limbs, from_bits - 1u)) { fill_ones(limbs, from_bits, to_bits); } } template HEDLEY_NO_THROW constexpr void store_raw_limbs(const T & value, std::uint64_t out[4]) noexcept { out[0] = out[1] = out[2] = out[3] = 0; if constexpr (std::is_same_v) { out[0] = value.lower().lower(); out[1] = value.lower().upper(); out[2] = value.upper().lower(); out[3] = value.upper().upper(); } else if constexpr (std::is_same_v) { out[0] = value.lower(); out[1] = value.upper(); } else if constexpr (std::is_same_v || std::is_same_v) { const simde_uint128 bits = static_cast(value); out[0] = static_cast(bits); out[1] = static_cast(bits >> 64); } else { using unsigned_same = std::make_unsigned_t; out[0] = static_cast(static_cast(value)); } } /// @brief Low `dest_bits` of `value`, sign-extended when `value` is a narrower signed integer. /// @tparam T value type /// @param value the value to convert or store /// @param src_bits the `src_bits` /// @param is_signed the `is_signed` /// @param dest_bits the `dest_bits` /// @param dest the destination /// @param nlimbs the `nlimbs` template HEDLEY_NO_THROW constexpr void reduce_operand(const T & value, unsigned src_bits, bool is_signed, unsigned dest_bits, std::uint64_t * dest, unsigned nlimbs) noexcept { for (unsigned i = 0; i < nlimbs; ++i) { dest[i] = 0; } std::uint64_t raw[4]; store_raw_limbs(value, raw); mask_to_bits(raw, 4, src_bits); const unsigned copying = nlimbs < 4u ? nlimbs : 4u; for (unsigned i = 0; i < copying; ++i) { dest[i] = raw[i]; } if (dest_bits > src_bits && is_signed) { sign_extend_range(dest, src_bits, dest_bits); } mask_to_bits(dest, nlimbs, dest_bits); } /// @brief Product modulo `2^(64*nlimbs)`, using exactly `nlimbs` limbs of each operand. /// @param out the output buffer /// @param lhs the left-hand operand /// @param rhs the right-hand operand /// @param nlimbs the `nlimbs` HEDLEY_NO_THROW constexpr void mul_low_limbs(std::uint64_t * out, const std::uint64_t * lhs, const std::uint64_t * rhs, unsigned nlimbs) noexcept { for (unsigned i = 0; i < nlimbs; ++i) { out[i] = 0; } for (unsigned i = 0; i < nlimbs; ++i) { simde_uint128 carry = 0; for (unsigned j = 0; i + j < nlimbs; ++j) { const simde_uint128 prod = simde_uint128(lhs[i]) * simde_uint128(rhs[j]) + simde_uint128(out[i + j]) + carry; out[i + j] = static_cast(prod); carry = prod >> 64; } } } HEDLEY_NO_THROW constexpr void shift_left_limbs(std::uint64_t * limbs, std::size_t nlimbs, unsigned shift) noexcept { if (shift == 0u) { return; } const unsigned limb_shift = shift / 64u; const unsigned bit_shift = shift % 64u; std::uint64_t tmp[fixed_mul_buf_limbs] = {}; for (std::size_t i = limb_shift; i < nlimbs; ++i) { const std::size_t src = i - limb_shift; std::uint64_t hi = limbs[src] << bit_shift; std::uint64_t lo = 0; if (bit_shift != 0u && src > 0u) { lo = limbs[src - 1u] >> (64u - bit_shift); } tmp[i] = hi | lo; } for (std::size_t i = 0; i < nlimbs; ++i) { limbs[i] = tmp[i]; } } HEDLEY_NO_THROW constexpr void shift_right_limbs(std::uint64_t * limbs, std::size_t nlimbs, unsigned shift) noexcept { if (shift == 0u) { return; } const unsigned limb_shift = shift / 64u; const unsigned bit_shift = shift % 64u; std::uint64_t tmp[fixed_mul_buf_limbs] = {}; for (std::size_t i = 0; i + limb_shift < nlimbs; ++i) { const std::size_t src = i + limb_shift; std::uint64_t lo = limbs[src] >> bit_shift; std::uint64_t hi = 0; if (bit_shift != 0u && src + 1u < nlimbs) { hi = limbs[src + 1u] << (64u - bit_shift); } tmp[i] = lo | hi; } for (std::size_t i = 0; i < nlimbs; ++i) { limbs[i] = tmp[i]; } } template HEDLEY_NO_THROW constexpr T limbs_to_integral(const std::uint64_t * limbs) noexcept { if constexpr (std::is_same_v) { return uint256_t{ uint128_t{limbs[3], limbs[2]}, uint128_t{limbs[1], limbs[0]}}; } else if constexpr (std::is_same_v) { return uint128_t{limbs[1], limbs[0]}; } else if constexpr (std::is_same_v || std::is_same_v) { const simde_uint128 bits = simde_uint128(limbs[0]) | (simde_uint128(limbs[1]) << 64); return static_cast(bits); } else { using unsigned_same = std::make_unsigned_t; return static_cast(static_cast(limbs[0])); } } } // namespace detail /// @brief Multiply two fixed-point values into a chosen integer and fraction width. /// @details The result is held in the smallest fixed-point word that can store /// `IntegerBits + FractionalBits`. A signed word is used when either operand /// is signed and the window is at most 128 bits; otherwise the window is the /// unsigned residue. /// @tparam IntegerBits Integer bits kept in the result, including the sign bit /// when the result is signed. Bits above this wrap. /// @tparam FractionalBits Fraction bits kept in the result. Lower fraction bits /// of the exact product are discarded (floored). /// @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 /// @param lhs the left-hand operand /// @param rhs the right-hand operand /// @return the product at the requested width template HEDLEY_NO_THROW constexpr auto fixed_mul( fixedpoint lhs, fixedpoint rhs) noexcept -> typename fixed_mul_plan::result_type { using plan = fixed_mul_plan; using integral = typename plan::integral_type; if constexpr (plan::modulus_bits == 0u || plan::limbs == 0u) { return make_fixed_from_integral_type(static_cast(0)); } else { // 1. Local reduction into the multiply ring. std::uint64_t left[8] = {}; std::uint64_t right[8] = {}; detail::reduce_operand(lhs.integral_representation(), plan::lhs_width, plan::lhs_signed, plan::multiply_bits, left, plan::limbs); detail::reduce_operand(rhs.integral_representation(), plan::rhs_width, plan::rhs_signed, plan::multiply_bits, right, plan::limbs); // 2. The single non-linear step: product in Z/2^multiply_bits Z. std::uint64_t prod[detail::fixed_mul_buf_limbs] = {}; detail::mul_low_limbs(prod, left, right, plan::limbs); detail::mask_to_bits(prod, detail::fixed_mul_buf_limbs, plan::multiply_bits); // 3. Public extension up to the window, then the public radix shift. if constexpr (plan::operands_signed) { detail::sign_extend_range(prod, plan::multiply_bits, plan::modulus_bits); } if constexpr (plan::align_shift > 0) { detail::shift_right_limbs(prod, detail::fixed_mul_buf_limbs, static_cast(plan::align_shift)); } else if constexpr (plan::align_shift < 0) { detail::shift_left_limbs(prod, detail::fixed_mul_buf_limbs, static_cast(-plan::align_shift)); } detail::mask_to_bits(prod, detail::fixed_mul_buf_limbs, plan::out_bits); if constexpr (plan::result_is_signed) { detail::sign_extend_range(prod, plan::out_bits, dpf::utils::bitlength_of_v); } return make_fixed_from_integral_type( detail::limbs_to_integral(prod)); } } } // namespace grotto #endif // LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_MUL_HPP__