libdpf/include/grotto/nmod.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
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>
2026-09-28 05:59:19 -06:00

212 lines
8.4 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @file grotto/nmod.hpp
/// @brief Reduction modulo an arbitrary public modulus.
/// @details `nmod` multiplies by a rounded reciprocal `1/M` and splits that
/// product with a floor. The quotient is `floor(x/M)`. The residue
/// is `{x/M}` truncated onto `residue_bits` fractional bits, so it
/// lies in `[0, 2^residue_bits)`. A negative product borrows, which
/// keeps the residue non-negative.
///
/// The reciprocal magnitude is 128 bits, so `1/M` can be carried
/// wider than the input. A power-of-two modulus is the same split
/// with an exact shift (`nmod_pow2`).
/// @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_NMOD_HPP__
#define LIBDPF_INCLUDE_GROTTO_NMOD_HPP__
#include <cstdint>
#include <stdexcept>
#include "hedley/hedley.h"
namespace grotto
{
namespace nmod_detail
{
using u128 = unsigned __int128;
}
/// @brief Quotient and truncated fractional residue of one reduction.
struct nmod_result
{
/// @brief `floor({x/M} * 2^residue_bits)`, in `[0, 2^residue_bits)`.
std::int64_t residue = 0;
/// @brief `floor(x/M)`.
std::int64_t quotient = 0;
};
/// @brief `x_raw / 2^x_bits` modulo `M`, with `1/M ≈ recip_raw / 2^recip_bits`.
/// @details `recip_raw` is a positive magnitude of at most 128 bits.
/// @param x_raw the signed integer significand
/// @param x_bits the fractional width of `x_raw`
/// @param recip_raw the positive magnitude of the rounded reciprocal
/// @param recip_bits the fractional width of `recip_raw`
/// @param residue_bits the fractional bits kept in the residue
/// @return `x_raw / 2^x_bits` modulo `M`, with `1/M ≈ recip_raw / 2^recip_bits`
/// @throws std::invalid_argument if the reciprocal is zero or a width is illegal.
/// @throws std::overflow_error if the quotient does not fit in `int64_t`.
/// \complexity One 64×128 multiply into four 64-bit limbs, then a constant number of limb tests to split the quotient and the residue. `Θ(1)` in that 256-bit width. Extra space is the four-limb array.
/// @see grotto::ring_switch_eval
/// @see grotto::nmod_pow2
/// @note Truncated Barrett. The residue is `{x/M}` on `residue_bits` fractional bits, not an exact `zn64` representative. Exact conversion is `ring_switch`.
HEDLEY_WARN_UNUSED_RESULT
inline nmod_result nmod(std::int64_t x_raw, unsigned x_bits,
unsigned __int128 recip_raw, unsigned recip_bits, unsigned residue_bits)
{
if (recip_raw == 0)
throw std::invalid_argument("nmod: reciprocal must be positive");
if (residue_bits > 63)
throw std::invalid_argument("nmod: residue must fit in int64");
if (x_bits > 100000u || recip_bits > 100000u)
throw std::invalid_argument("nmod: fractional width is too large");
const bool neg = x_raw < 0;
const auto x_mag = static_cast<std::uint64_t>(
neg ? -static_cast<__int128>(x_raw) : x_raw);
const auto recip_lo = static_cast<std::uint64_t>(recip_raw);
const auto recip_hi = static_cast<std::uint64_t>(recip_raw >> 64);
unsigned __int128 low = static_cast<unsigned __int128>(x_mag) * recip_lo;
unsigned __int128 high = static_cast<unsigned __int128>(x_mag) * recip_hi;
std::uint64_t limb[4] = {};
limb[0] = static_cast<std::uint64_t>(low);
const unsigned __int128 mid = (low >> 64) + static_cast<std::uint64_t>(high);
limb[1] = static_cast<std::uint64_t>(mid);
const unsigned __int128 top = (high >> 64) + (mid >> 64);
limb[2] = static_cast<std::uint64_t>(top);
limb[3] = static_cast<std::uint64_t>(top >> 64);
const auto bit_set_at_or_above = [&](unsigned bit) {
if (bit >= 256)
return false;
const unsigned index = bit / 64u;
const unsigned offset = bit % 64u;
if ((limb[index] >> offset) != 0)
return true;
for (unsigned i = index + 1; i < 4; ++i)
{
if (limb[i] != 0)
return true;
}
return false;
};
const auto extract = [&](unsigned low_bit, unsigned count) -> std::uint64_t {
if (count == 0 || low_bit >= 256)
return 0;
const unsigned index = low_bit / 64u;
const unsigned offset = low_bit % 64u;
unsigned __int128 chunk = limb[index];
if (index + 1 < 4)
chunk |= static_cast<unsigned __int128>(limb[index + 1]) << 64;
chunk >>= offset;
if (count == 64)
return static_cast<std::uint64_t>(chunk);
return static_cast<std::uint64_t>(chunk) & ((std::uint64_t{1} << count) - 1);
};
const auto low_bits_set = [&](unsigned width) {
if (width == 0)
return false;
if (width >= 256)
return limb[0] != 0 || limb[1] != 0 || limb[2] != 0 || limb[3] != 0;
const unsigned index = width / 64u;
const unsigned offset = width % 64u;
for (unsigned i = 0; i < index; ++i)
{
if (limb[i] != 0)
return true;
}
if (offset == 0)
return false;
const std::uint64_t mask = (std::uint64_t{1} << offset) - 1;
return (limb[index] & mask) != 0;
};
const unsigned scale = x_bits + recip_bits;
const bool remainder = low_bits_set(scale);
std::uint64_t quotient_mag = 0;
if (scale < 256)
{
if (bit_set_at_or_above(scale + 64))
throw std::overflow_error("nmod: quotient does not fit int64");
quotient_mag = extract(scale, 64);
}
nmod_result out;
if (!neg)
{
if (quotient_mag > static_cast<std::uint64_t>(INT64_MAX))
throw std::overflow_error("nmod: quotient does not fit int64");
out.quotient = static_cast<std::int64_t>(quotient_mag);
}
else if (!remainder)
{
if (quotient_mag > (static_cast<std::uint64_t>(INT64_MAX) + 1))
throw std::overflow_error("nmod: quotient does not fit int64");
out.quotient = quotient_mag == (std::uint64_t{1} << 63)
? INT64_MIN
: -static_cast<std::int64_t>(quotient_mag);
}
else
{
if (quotient_mag > static_cast<std::uint64_t>(INT64_MAX))
throw std::overflow_error("nmod: quotient does not fit int64");
out.quotient = -static_cast<std::int64_t>(quotient_mag) - 1;
}
if (residue_bits == 0 || (!neg && !remainder) || (neg && !remainder))
return out;
// A 128-bit reciprocal times an int64 magnitude stays under 2^192.
// With a residue of at most 63 bits, a scale at or above 256 puts every
// product bit strictly below the residue window.
if (scale >= 256)
{
if (neg)
out.residue = static_cast<std::int64_t>((std::uint64_t{1} << residue_bits) - 1);
return out;
}
std::uint64_t field = 0;
if (scale >= residue_bits)
field = extract(scale - residue_bits, residue_bits);
else
field = extract(0, scale) << (residue_bits - scale);
if (neg)
{
const std::uint64_t unit = std::uint64_t{1} << residue_bits;
const unsigned discarded = scale >= residue_bits ? scale - residue_bits : 0;
const bool borrow = discarded > 0 && low_bits_set(discarded);
field = borrow ? unit - field - 1 : unit - field;
}
out.residue = static_cast<std::int64_t>(field);
return out;
}
/// @brief Modulus `2^{-exp}` by an exact shift. Positive `exp` multiplies by
/// `2^exp` (`exp`'s `2^{-13}` split). Zero splits at the integer (`2^x`).
/// @details Negative `exp` is a modulus above one.
/// @param x_raw the signed integer significand
/// @param x_bits the fractional width of `x_raw`
/// @param exp the power-of-two exponent
/// @param residue_bits the fractional bits kept in the residue
/// @return Modulus `2^{-exp}` by an exact shift
/// @throws std::overflow_error if `exp` does not fit the reciprocal width.
/// \complexity One call to `nmod` with a reciprocal that is a shift (`2^{exp}` or `2^{-exp}`). `Θ(1)`.
/// @see grotto::nmod
HEDLEY_WARN_UNUSED_RESULT
inline nmod_result nmod_pow2(std::int64_t x_raw, unsigned x_bits, int exp,
unsigned residue_bits)
{
if (exp >= 128 || exp < -100000)
throw std::overflow_error("nmod: power-of-two reciprocal does not fit");
if (exp >= 0)
return nmod(x_raw, x_bits, nmod_detail::u128{1} << static_cast<unsigned>(exp), 0, residue_bits);
return nmod(x_raw, x_bits, 1, static_cast<unsigned>(-exp), residue_bits);
}
} // namespace grotto
#endif // LIBDPF_INCLUDE_GROTTO_NMOD_HPP__