libdpf/include/dpf/field128.hpp

493 lines
15 KiB
C++
Raw Permalink Normal View History

/// @file dpf/field128.hpp
/// @brief Prime field of libprio's `Field128`, as a DPF output.
/// @details The modulus is 340282366920938462946865773367900766209. Pass
/// `dpf::field128{n}` as a `make_dpf` payload. Leaf addition and
/// scaling are the field operations. A raw PRG block is reduced into
/// the field on the first leaf operation. This is a point-function
/// output, not a comparison payload.
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details.
#ifndef LIBDPF_INCLUDE_DPF_FIELD128_HPP__
#define LIBDPF_INCLUDE_DPF_FIELD128_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <iomanip>
#include <ostream>
#include <type_traits>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "dpf/leaf_arithmetic.hpp"
#include "dpf/random.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// @brief Element of the 128-bit libprio field.
class field128
{
public:
/// @brief Low limb of the modulus `2^128 - 0x1bffffffffffffffff`.
static constexpr std::uint64_t mod_lo = 0x0000000000000001ull;
/// @brief High limb of the modulus.
static constexpr std::uint64_t mod_hi = 0xffffffffffffffe4ull;
static constexpr bool dpf_point_group = true;
/// @brief The zero element.
HEDLEY_ALWAYS_INLINE
constexpr field128() noexcept = default;
/// @brief Reduce `v` into the field. A negative value is negated in the field.
/// @tparam T integral type, at most 128 bits
/// @param v the integer to reduce
template <typename T, typename = std::enable_if_t<std::is_integral_v<T>>>
HEDLEY_ALWAYS_INLINE
constexpr field128(T v) noexcept
{
assign_integer(v);
}
/// @brief Reduce a PRG block into the field. Uses up to 16 bytes.
/// @param bytes the PRG output
/// @param n the number of bytes available
/// @return the field element
HEDLEY_ALWAYS_INLINE
static field128 from_seed(const void * bytes, std::size_t n) noexcept
{
unsigned char buf[16]{};
if (n > sizeof(buf))
n = sizeof(buf);
std::memcpy(buf, bytes, n);
std::uint64_t w[2]{};
std::memcpy(w, buf, sizeof(w));
return from_lane(w[0], w[1]);
}
/// @brief Reduce an unreduced 128-bit lane `(hi << 64) + lo`.
/// @param lo the low limb, not necessarily reduced
/// @param hi the high limb, not necessarily reduced
/// @return the field element
HEDLEY_ALWAYS_INLINE
static constexpr field128 from_lane(std::uint64_t lo, std::uint64_t hi) noexcept
{
const std::uint64_t words[4] = {lo, hi, 0, 0};
return reduce_words(words);
}
/// @brief Canonical representative in `[0, mod)`.
HEDLEY_ALWAYS_INLINE
static constexpr field128 canonicalize(field128 a) noexcept
{
return from_lane(a.lo_, a.hi_);
}
/// @brief Low 64 bits of the reduced representative.
/// @return the low limb
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
constexpr std::uint64_t lo() const noexcept { return lo_; }
/// @brief High 64 bits of the reduced representative.
/// @return the high limb
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
constexpr std::uint64_t hi() const noexcept { return hi_; }
/// @brief Field addition.
/// @param a left addend
/// @param b right addend
/// @return `a + b` in the field
HEDLEY_ALWAYS_INLINE
friend constexpr field128 operator+(field128 a, field128 b) noexcept
{
const unsigned __int128 low =
static_cast<unsigned __int128>(a.lo_) + b.lo_;
const unsigned __int128 high =
static_cast<unsigned __int128>(a.hi_) + b.hi_ + (low >> 64);
const std::uint64_t words[4] = {
static_cast<std::uint64_t>(low),
static_cast<std::uint64_t>(high),
static_cast<std::uint64_t>(high >> 64),
0};
return reduce_words(words);
}
/// @brief Field subtraction.
/// @param a minuend
/// @param b subtrahend
/// @return `a - b` in the field
HEDLEY_ALWAYS_INLINE
friend constexpr field128 operator-(field128 a, field128 b) noexcept
{
a = canonicalize(a);
b = canonicalize(b);
if (!less(a.hi_, a.lo_, b.hi_, b.lo_))
{
std::uint64_t lo = 0, hi = 0;
sub_words(a.lo_, a.hi_, b.lo_, b.hi_, lo, hi);
return from_reduced(lo, hi);
}
std::uint64_t lo = 0, hi = 0;
sub_words(mod_lo, mod_hi, b.lo_, b.hi_, lo, hi);
return from_reduced(lo, hi) + a;
}
/// @brief Field negation.
/// @param a the element to negate
/// @return `-a`, with `-0 = 0`
HEDLEY_ALWAYS_INLINE
friend constexpr field128 operator-(field128 a) noexcept
{
a = canonicalize(a);
if (a.lo_ == 0 && a.hi_ == 0)
return a;
std::uint64_t lo = 0, hi = 0;
sub_words(mod_lo, mod_hi, a.lo_, a.hi_, lo, hi);
return from_reduced(lo, hi);
}
/// @brief Field multiplication.
/// @param a left factor
/// @param b right factor
/// @return `a * b` in the field
HEDLEY_ALWAYS_INLINE
friend constexpr field128 operator*(field128 a, field128 b) noexcept
{
std::uint64_t words[4]{};
mul128(a.lo_, a.hi_, b.lo_, b.hi_, words);
return reduce_words(words);
}
/// @brief Field equality.
/// @param a left element
/// @param b right element
/// @return `true` when the reduced values match
HEDLEY_ALWAYS_INLINE
friend constexpr bool operator==(field128 a, field128 b) noexcept
{
a = canonicalize(a);
b = canonicalize(b);
return a.lo_ == b.lo_ && a.hi_ == b.hi_;
}
/// @brief Field inequality.
/// @param a left element
/// @param b right element
/// @return `true` when the reduced values differ
HEDLEY_ALWAYS_INLINE
friend constexpr bool operator!=(field128 a, field128 b) noexcept
{
return !(a == b);
}
/// @brief Write the reduced representative in hexadecimal.
/// @param os the output stream
/// @param a the element to write
/// @return `os`
friend std::ostream & operator<<(std::ostream & os, field128 a)
{
const auto flags = os.flags();
os << "0x" << std::hex << a.hi_ << std::setfill('0') << std::setw(16) << a.lo_;
os.flags(flags);
return os;
}
private:
std::uint64_t lo_{};
std::uint64_t hi_{};
HEDLEY_ALWAYS_INLINE
static constexpr field128 from_reduced(std::uint64_t lo, std::uint64_t hi) noexcept
{
field128 out;
out.lo_ = lo;
out.hi_ = hi;
return out;
}
HEDLEY_ALWAYS_INLINE
static constexpr bool less(std::uint64_t hi, std::uint64_t lo,
std::uint64_t ohi, std::uint64_t olo) noexcept
{
return hi < ohi || (hi == ohi && lo < olo);
}
HEDLEY_ALWAYS_INLINE
static constexpr bool ge_mod(std::uint64_t hi, std::uint64_t lo) noexcept
{
return hi > mod_hi || (hi == mod_hi && lo >= mod_lo);
}
HEDLEY_ALWAYS_INLINE
static constexpr void sub_words(std::uint64_t al, std::uint64_t ah,
std::uint64_t bl, std::uint64_t bh,
std::uint64_t &ol, std::uint64_t &oh) noexcept
{
const unsigned borrow = al < bl ? 1u : 0u;
ol = static_cast<std::uint64_t>(al - bl);
oh = static_cast<std::uint64_t>(ah - bh - borrow);
}
HEDLEY_ALWAYS_INLINE
static constexpr void add_c(std::uint64_t &lo, std::uint64_t &hi,
std::uint64_t &top) noexcept
{
const unsigned __int128 low =
static_cast<unsigned __int128>(lo) + 0xffffffffffffffffull;
lo = static_cast<std::uint64_t>(low);
const unsigned __int128 high =
static_cast<unsigned __int128>(hi) + 0x1bull + static_cast<std::uint64_t>(low >> 64);
hi = static_cast<std::uint64_t>(high);
top = static_cast<std::uint64_t>(high >> 64);
}
/// @brief 128×128 → 256 multiply, little-endian limbs.
HEDLEY_ALWAYS_INLINE
static constexpr void mul128(std::uint64_t a0, std::uint64_t a1,
std::uint64_t b0, std::uint64_t b1, std::uint64_t out[4]) noexcept
{
const unsigned __int128 p00 = static_cast<unsigned __int128>(a0) * b0;
const unsigned __int128 p01 = static_cast<unsigned __int128>(a0) * b1;
const unsigned __int128 p10 = static_cast<unsigned __int128>(a1) * b0;
const unsigned __int128 p11 = static_cast<unsigned __int128>(a1) * b1;
out[0] = static_cast<std::uint64_t>(p00);
const unsigned __int128 mid = (p00 >> 64)
+ static_cast<std::uint64_t>(p01) + static_cast<std::uint64_t>(p10);
out[1] = static_cast<std::uint64_t>(mid);
const unsigned __int128 high = (mid >> 64) + (p01 >> 64) + (p10 >> 64)
+ static_cast<std::uint64_t>(p11);
out[2] = static_cast<std::uint64_t>(high);
out[3] = static_cast<std::uint64_t>((high >> 64) + (p11 >> 64));
}
/// @brief Reduce a little-endian 256-bit integer. `2^128 ≡ C`.
HEDLEY_ALWAYS_INLINE
static constexpr field128 reduce_words(const std::uint64_t z[4]) noexcept
{
std::uint64_t lo = 0, hi = 0, top = 0;
for (int bit = 255; bit >= 0; --bit)
{
top = hi >> 63;
hi = (hi << 1) | (lo >> 63);
lo <<= 1;
const unsigned limb = static_cast<unsigned>(bit >> 6);
const unsigned off = static_cast<unsigned>(bit & 63);
if ((z[limb] >> off) & 1u)
{
const unsigned __int128 low =
static_cast<unsigned __int128>(lo) + 1;
lo = static_cast<std::uint64_t>(low);
const unsigned __int128 high =
static_cast<unsigned __int128>(hi) + static_cast<std::uint64_t>(low >> 64);
hi = static_cast<std::uint64_t>(high);
top += static_cast<std::uint64_t>(high >> 64);
}
while (top)
add_c(lo, hi, top);
if (ge_mod(hi, lo))
sub_words(lo, hi, mod_lo, mod_hi, lo, hi);
}
return from_reduced(lo, hi);
}
template <typename T>
HEDLEY_ALWAYS_INLINE
constexpr void assign_integer(T v) noexcept
{
bool neg = false;
unsigned __int128 mag = 0;
if constexpr (std::is_signed_v<T>)
{
if (v < 0)
{
neg = true;
using U = std::make_unsigned_t<T>;
mag = static_cast<U>(0) - static_cast<U>(v);
}
else
{
mag = static_cast<std::make_unsigned_t<T>>(v);
}
}
else
{
mag = static_cast<unsigned __int128>(v);
}
const std::uint64_t words[4] = {
static_cast<std::uint64_t>(mag),
static_cast<std::uint64_t>(mag >> 64),
0, 0};
*this = reduce_words(words);
if (neg)
*this = -*this;
}
};
namespace utils
{
template <>
struct bitlength_of<field128>
: std::integral_constant<std::size_t, 128>
{ };
template <>
struct has_characteristic_two<field128> : std::false_type
{ };
} // namespace utils
namespace leaf_arithmetic
{
namespace detail
{
HEDLEY_ALWAYS_INLINE
field128 load_field128(const void *lane) noexcept
{
std::uint64_t w[2]{};
std::memcpy(w, lane, sizeof(w));
return field128::from_lane(w[0], w[1]);
}
HEDLEY_ALWAYS_INLINE
void store_field128(void *lane, field128 v) noexcept
{
const std::uint64_t w[2] = {v.lo(), v.hi()};
std::memcpy(lane, w, sizeof(w));
}
template <std::size_t Lanes>
HEDLEY_ALWAYS_INLINE
void field128_lanes(const void *a, const void *b, void *out,
field128 (*op)(field128, field128)) noexcept
{
auto *aa = static_cast<const unsigned char *>(a);
auto *bb = static_cast<const unsigned char *>(b);
auto *cc = static_cast<unsigned char *>(out);
for (std::size_t i = 0; i < Lanes; ++i)
{
const field128 y = op(load_field128(aa + i * 16), load_field128(bb + i * 16));
store_field128(cc + i * 16, y);
}
}
template <std::size_t Lanes>
HEDLEY_ALWAYS_INLINE
void field128_scale(const void *a, field128 b, void *out) noexcept
{
auto *aa = static_cast<const unsigned char *>(a);
auto *cc = static_cast<unsigned char *>(out);
for (std::size_t i = 0; i < Lanes; ++i)
store_field128(cc + i * 16, load_field128(aa + i * 16) * b);
}
} // namespace detail
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
template <>
struct add_t<field128, simde__m128i>
{
auto operator()(const simde__m128i &a, const simde__m128i &b) const
{
simde__m128i out;
detail::field128_lanes<1>(&a, &b, &out, [](field128 x, field128 y) {
return x + y;
});
return out;
}
};
template <>
struct subtract_t<field128, simde__m128i>
{
auto operator()(const simde__m128i &a, const simde__m128i &b) const
{
simde__m128i out;
detail::field128_lanes<1>(&a, &b, &out, [](field128 x, field128 y) {
return x - y;
});
return out;
}
};
template <>
struct multiply_t<field128, simde__m128i>
{
auto operator()(const simde__m128i &a, field128 b) const
{
simde__m128i out;
detail::field128_scale<1>(&a, b, &out);
return out;
}
};
template <>
struct add_t<field128, simde__m256i>
{
auto operator()(const simde__m256i &a, const simde__m256i &b) const
{
simde__m256i out;
detail::field128_lanes<2>(&a, &b, &out, [](field128 x, field128 y) {
return x + y;
});
return out;
}
};
template <>
struct subtract_t<field128, simde__m256i>
{
auto operator()(const simde__m256i &a, const simde__m256i &b) const
{
simde__m256i out;
detail::field128_lanes<2>(&a, &b, &out, [](field128 x, field128 y) {
return x - y;
});
return out;
}
};
template <>
struct multiply_t<field128, simde__m256i>
{
auto operator()(const simde__m256i &a, field128 b) const
{
simde__m256i out;
detail::field128_scale<2>(&a, b, &out);
return out;
}
};
HEDLEY_PRAGMA(GCC diagnostic pop)
} // namespace leaf_arithmetic
/// @brief Sample a uniform field element by rejection.
/// @return an element of the field
template <>
HEDLEY_NO_THROW
inline auto uniform_sample<field128>() noexcept
{
for (;;)
{
const auto lo = uniform_sample<std::uint64_t>();
const auto hi = uniform_sample<std::uint64_t>();
if (hi < field128::mod_hi || (hi == field128::mod_hi && lo < field128::mod_lo))
return field128::from_lane(lo, hi);
}
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_FIELD128_HPP__