2026-09-24 14:08:32 -06:00
/// @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__
2026-09-24 20:44:07 -06:00
# include "hedley/hedley.h"
2026-09-24 14:08:32 -06:00
# ifndef LIBDPF_INCLUDE_DPF_FIXEDPOINT_HPP__
# include "grotto/fixedpoint.hpp"
# endif
# include <algorithm>
# include <cstddef>
# include <cstdint>
# include <type_traits>
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.
///
2026-09-28 05:59:19 -06:00
/// A Beaver triple replaces only the multiply in `Z/2^multiply_bits Z`
/// (`grotto::eval_fixed_mul_beaver`). Dropping bits above that ring is local.
/// A narrower share is lifted by cancelling the carry out of the operand
/// width; a signed lift also replicates the sign. The same correction
/// replicates the product sign when `modulus_bits > multiply_bits` and
/// supplies the right shift by `align_shift`. Each of those lifts extends
/// by at most 64 bits, and the shift discards at most 64 bits.
2026-09-24 23:18:10 -06:00
/// @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
2026-09-24 14:08:32 -06:00
template < unsigned IntegerBits ,
unsigned FractionalBits ,
unsigned LhsFractionalBits ,
typename LhsIntegral ,
unsigned RhsFractionalBits ,
typename RhsIntegral >
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 < LhsIntegral > ;
static constexpr unsigned rhs_width = dpf : : utils : : bitlength_of_v < RhsIntegral > ;
static constexpr bool lhs_signed = std : : is_signed_v < LhsIntegral > ;
static constexpr bool rhs_signed = std : : is_signed_v < RhsIntegral > ;
static constexpr bool operands_signed = lhs_signed | | rhs_signed ;
2026-09-24 23:18:10 -06:00
/// @brief Right shift applied to the raw product. Negative means a left shift.
2026-09-24 14:08:32 -06:00
static constexpr int align_shift = static_cast < int > ( LhsFractionalBits )
+ static_cast < int > ( RhsFractionalBits )
- static_cast < int > ( FractionalBits ) ;
2026-09-24 23:18:10 -06:00
/// @brief Bits of the product that the shift reads. Zero when a left shift
2026-09-24 14:08:32 -06:00
/// moves every product bit out of the output.
static constexpr int modulus_bits_signed = align_shift > = 0
? align_shift + static_cast < int > ( out_bits )
: static_cast < int > ( out_bits ) + align_shift ;
static constexpr unsigned modulus_bits = modulus_bits_signed > 0
? static_cast < unsigned > ( modulus_bits_signed ) : 0u ;
2026-09-24 23:18:10 -06:00
/// @brief Full two's-complement product fits in this many bits.
2026-09-24 14:08:32 -06:00
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 ;
2026-09-24 23:18:10 -06:00
/// @brief Signed storage exists through 128 bits. A wider window is the same
2026-09-24 14:08:32 -06:00
/// 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 < unsigned Bits , bool Signed >
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 < Signed ,
std : : conditional_t < width < = 8u , std : : int8_t ,
std : : conditional_t < width < = 16u , std : : int16_t ,
std : : conditional_t < width < = 32u , std : : int32_t ,
std : : conditional_t < width < = 64u , std : : int64_t , simde_int128 > > > > ,
dpf : : utils : : integral_type_from_bitlength_t < width > > ;
} ;
public :
using integral_type = typename storage < out_bits , result_is_signed > : : type ;
using result_type = fixedpoint < FractionalBits , integral_type > ;
} ;
namespace detail
{
inline constexpr std : : size_t fixed_mul_buf_limbs = 12 ;
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 ;
}
}
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
constexpr bool test_bit ( const std : : uint64_t * limbs , unsigned bit ) noexcept
{
return ( ( limbs [ bit / 64u ] > > ( bit % 64u ) ) & 1u ) ! = 0u ;
}
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 ;
}
}
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 < typename T >
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 < T , uint256_t > )
{
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 < T , uint128_t > )
{
out [ 0 ] = value . lower ( ) ;
out [ 1 ] = value . upper ( ) ;
}
else if constexpr ( std : : is_same_v < T , simde_uint128 > | | std : : is_same_v < T , simde_int128 > )
{
const simde_uint128 bits = static_cast < simde_uint128 > ( value ) ;
out [ 0 ] = static_cast < std : : uint64_t > ( bits ) ;
out [ 1 ] = static_cast < std : : uint64_t > ( bits > > 64 ) ;
}
else
{
using unsigned_same = std : : make_unsigned_t < T > ;
out [ 0 ] = static_cast < std : : uint64_t > ( static_cast < unsigned_same > ( value ) ) ;
}
}
2026-09-24 23:18:10 -06:00
/// @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
2026-09-28 05:59:19 -06:00
/// @param src_bits width of `value` before extension
/// @param is_signed whether `value` is a signed integer of `src_bits` bits
/// @param dest_bits width of the destination word
2026-09-24 23:18:10 -06:00
/// @param dest the destination
2026-09-28 05:59:19 -06:00
/// @param nlimbs number of 64-bit limbs kept in the product
2026-09-24 14:08:32 -06:00
template < typename T >
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 ) ;
}
2026-09-24 23:18:10 -06:00
/// @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
2026-09-28 05:59:19 -06:00
/// @param nlimbs number of 64-bit limbs kept in the product
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 < std : : uint64_t > ( prod ) ;
carry = prod > > 64 ;
}
}
}
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 ] ;
}
}
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
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 < typename T >
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
constexpr T limbs_to_integral ( const std : : uint64_t * limbs ) noexcept
{
if constexpr ( std : : is_same_v < T , uint256_t > )
{
return uint256_t {
uint128_t { limbs [ 3 ] , limbs [ 2 ] } ,
uint128_t { limbs [ 1 ] , limbs [ 0 ] } } ;
}
else if constexpr ( std : : is_same_v < T , uint128_t > )
{
return uint128_t { limbs [ 1 ] , limbs [ 0 ] } ;
}
else if constexpr ( std : : is_same_v < T , simde_uint128 > | | std : : is_same_v < T , simde_int128 > )
{
const simde_uint128 bits = simde_uint128 ( limbs [ 0 ] )
| ( simde_uint128 ( limbs [ 1 ] ) < < 64 ) ;
return static_cast < T > ( bits ) ;
}
else
{
using unsigned_same = std : : make_unsigned_t < T > ;
return static_cast < T > ( static_cast < unsigned_same > ( limbs [ 0 ] ) ) ;
}
}
} // namespace detail
/// @brief Multiply two fixed-point values into a chosen integer and fraction width.
2026-09-24 23:18:10 -06:00
/// @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.
2026-09-24 14:08:32 -06:00
/// @tparam FractionalBits Fraction bits kept in the result. Lower fraction bits
2026-09-24 23:18:10 -06:00
/// 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
2026-09-28 05:59:19 -06:00
/// \complexity `mul_low_limbs` multiplies `L` 64-bit limbs, `L = fixed_mul_plan::limbs`, with loops `i < L` and `i + j < L`: `Θ(L²)` limb products.
/// Reducing each operand and the align shift are `Θ(L)`. Scratch is a fixed limb buffer in this header (the 8-word arrays plus `fixed_mul_buf_limbs`).
/// @see grotto::fixedpoint
/// @note Plaintext. `grotto::eval_fixed_mul_beaver` is the same window on additive shares: one Beaver product in `Z/2^{multiply_bits}Z`, then the lifts and the aligning shift.
/// @see grotto::eval_fixed_mul_beaver
2026-09-24 14:08:32 -06:00
template < unsigned IntegerBits ,
unsigned FractionalBits ,
unsigned LhsFractionalBits ,
typename LhsIntegral ,
unsigned RhsFractionalBits ,
typename RhsIntegral >
2026-09-24 20:44:07 -06:00
HEDLEY_NO_THROW
2026-09-24 14:08:32 -06:00
constexpr auto fixed_mul (
fixedpoint < LhsFractionalBits , LhsIntegral > lhs ,
fixedpoint < RhsFractionalBits , RhsIntegral > rhs ) noexcept
- > typename fixed_mul_plan < IntegerBits , FractionalBits ,
LhsFractionalBits , LhsIntegral , RhsFractionalBits , RhsIntegral > : : result_type
{
using plan = fixed_mul_plan < IntegerBits , FractionalBits ,
LhsFractionalBits , LhsIntegral , RhsFractionalBits , RhsIntegral > ;
using integral = typename plan : : integral_type ;
if constexpr ( plan : : modulus_bits = = 0u | | plan : : limbs = = 0u )
{
return make_fixed_from_integral_type < FractionalBits , integral > ( static_cast < integral > ( 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 < unsigned > ( plan : : align_shift ) ) ;
}
else if constexpr ( plan : : align_shift < 0 )
{
detail : : shift_left_limbs ( prod , detail : : fixed_mul_buf_limbs ,
static_cast < unsigned > ( - 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 < integral > ) ;
}
return make_fixed_from_integral_type < FractionalBits , integral > (
detail : : limbs_to_integral < integral > ( prod ) ) ;
}
}
} // namespace grotto
# endif // LIBDPF_INCLUDE_GROTTO_FIXEDPOINT_MUL_HPP__