libdpf/include/grotto/nmod.hpp

206 lines
7.9 KiB
C++

/// @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`.
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.
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__