Document the new DPF surfaces in one command set, and test the field, half-tree, and multipoint edges.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Ryan Henry 2026-09-24 23:18:10 -06:00
parent 0d8a5a8131
commit 0dff6df8ed
250 changed files with 12199 additions and 1981 deletions

File diff suppressed because it is too large Load diff

1
doc/assets/assets.dox Normal file
View file

@ -0,0 +1 @@
// Registers doc/assets with Doxygen so `@dir` can document the image directory.

View file

@ -116,6 +116,14 @@
#include "dpf/uint256_t.hpp" #include "dpf/uint256_t.hpp"
#include "dpf/fp61.hpp"
#include "dpf/verifiable.hpp"
#include "dpf/multipoint.hpp"
#include "dpf/vec.hpp"
#include "dpf/interval.hpp" #include "dpf/interval.hpp"
#endif // LIBDPF_INCLUDE_DPF_HPP__ #endif // LIBDPF_INCLUDE_DPF_HPP__

View file

@ -339,12 +339,15 @@ auto bit_array_from_advice_bits_simde(Iterator first, Iterator last,
static_assert(CHAR_BIT == 8, "CHAR_BIT not equal to 8"); static_assert(CHAR_BIT == 8, "CHAR_BIT not equal to 8");
auto ret = dynamic_bit_array(bits); auto ret = dynamic_bit_array(bits);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::size_t bits_per_byte = CHAR_BIT, std::size_t bits_per_byte = CHAR_BIT,
bytes = (bits-1)/bits_per_byte + 1, bytes = (bits-1)/bits_per_byte + 1,
bits_per_word = ret.bits_per_word, bits_per_word = ret.bits_per_word,
bits_per_simde = dpf::utils::bitlength_of_v<simde_type>, bits_per_simde = dpf::utils::bitlength_of_v<simde_type>,
bytes_per_simde = sizeof(simde_type), bytes_per_simde = sizeof(simde_type),
words_per_simde = bits_per_simde / bits_per_word; words_per_simde = bits_per_simde / bits_per_word;
HEDLEY_PRAGMA(GCC diagnostic pop)
std::size_t curbits = 0, pos = 0; std::size_t curbits = 0, pos = 0;
std::array<char, 32> in = {0}; std::array<char, 32> in = {0};

View file

@ -54,6 +54,7 @@ class aligned_allocator
/// @brief a `deleter` functor for use by `std::unique_ptr<T[]>` to free /// @brief a `deleter` functor for use by `std::unique_ptr<T[]>` to free
/// memory allocated when the `std::unique_ptr<T[]>` was /// memory allocated when the `std::unique_ptr<T[]>` was
/// constructed /// constructed
/// @tparam Pointer pointer type stored in the deleter
template <typename Pointer> template <typename Pointer>
struct deleter struct deleter
{ {

View file

@ -1,6 +1,5 @@
/// @file dpf/asio.hpp /// @file dpf/asio.hpp
/// @brief /// @brief ASIO helpers for shipping DPF keys and assigning wildcard inputs.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2023 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2023 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;

View file

@ -14,7 +14,9 @@
/// A polynomial is a sum of monomials in several wires. /// A polynomial is a sum of monomials in several wires.
/// `2 + 3*x + 4*y + 5*x*y + 6*pow(x, 2) + pow(x, 2)*y + x*y*z` /// `2 + 3*x + 4*y + 5*x*y + 6*pow(x, 2) + pow(x, 2)*y + x*y*z`
/// is one round. `λ_x²` is stored once whether it appears as `x²`, /// is one round. `λ_x²` is stored once whether it appears as `x²`,
/// inside `x² y`, or in a second polynomial. Wires that occur with the /// inside `x² y`, or in a second polynomial. An inner product is that
/// sum: `dot({x0,x1}, {y0,y1})` is `x0*y0 + x1*y1`, and the pair
/// products share one preprocessing value. Wires that occur with the
/// same exponents in every term, as in `a3*(x*z)^3 + a2*(x*z)^2 + a1*(x*z) + a0`, /// same exponents in every term, as in `a3*(x*z)^3 + a2*(x*z)^2 + a1*(x*z) + a0`,
/// are multiplied first and the univariate polynomial is a later round. /// are multiplied first and the univariate polynomial is a later round.
/// A factor shared by every term, such as a sign or a piecewise scale, /// A factor shared by every term, such as a sign or a piecewise scale,
@ -70,8 +72,9 @@ namespace dpf
namespace beavers namespace beavers
{ {
/// Ring operations used to build and consume triples. /// @brief Ring operations used to build and consume triples.
/// Specialize for a ring whose multiplicative identity is not `Ring{1}`. /// @details Specialize for a ring whose multiplicative identity is not `Ring{1}`.
/// @tparam Ring payload ring
template <typename Ring> template <typename Ring>
struct ring_traits struct ring_traits
{ {
@ -90,7 +93,8 @@ struct ring_traits
static Ring sample() { return dpf::uniform_sample<Ring>(); } static Ring sample() { return dpf::uniform_sample<Ring>(); }
}; };
/// Bitwise AND uses the all-ones word as its multiplicative identity. /// @brief Bitwise AND uses the all-ones word as its multiplicative identity.
/// @tparam T value type
template <typename T> template <typename T>
struct ring_traits<dpf::xor_wrapper<T>> struct ring_traits<dpf::xor_wrapper<T>>
{ {
@ -121,7 +125,8 @@ struct default_sampler
Ring operator()() const { return ring_traits<Ring>::sample(); } Ring operator()() const { return ring_traits<Ring>::sample(); }
}; };
/// Additive (2,2) split. `open()` is `p0 + p1`. /// @brief Additive (2,2) split. `open()` is `p0 + p1`.
/// @tparam Ring payload ring
template <typename Ring> template <typename Ring>
struct split struct split
{ {
@ -157,9 +162,10 @@ struct split
template <typename Ring> template <typename Ring>
class session; class session;
/// A value in a session. Copying a wire copies its id; it does not copy the /// @brief A value in a session. Copying a wire copies its id; it does not copy the
/// blind. The ring argument is on the type so `a * x * x` can build an /// blind. The ring argument is on the type so `a * x * x` can build an
/// expression without naming the session. /// expression without naming the session.
/// @tparam Ring payload ring
template <typename Ring> template <typename Ring>
class wire class wire
{ {
@ -192,15 +198,16 @@ private:
std::uint32_t id_ = 0; std::uint32_t id_ = 0;
}; };
/// Unevaluated sum of monomials. `*` distributes over `+`. A public /// @brief Unevaluated sum of monomials. `*` distributes over `+`. A public
/// coefficient scales a term. Nothing is sampled until `session::operator()`. /// coefficient scales a term. Nothing is sampled until `session::operator()`.
/// @tparam Ring payload ring
template <typename Ring> template <typename Ring>
struct expr struct expr
{ {
struct term struct term
{ {
Ring coeff{}; Ring coeff{};
/// Positive exponents, sorted by wire id. /// @brief Positive exponents, sorted by wire id.
std::vector<std::pair<std::uint32_t, std::uint8_t>> powers; std::vector<std::pair<std::uint32_t, std::uint8_t>> powers;
}; };
@ -226,10 +233,20 @@ expr<Ring> wire_expr(wire<Ring> w);
template <typename Ring> template <typename Ring>
expr<Ring> horner_expr(wire<Ring> x, std::initializer_list<Ring> coeffs); expr<Ring> horner_expr(wire<Ring> x, std::initializer_list<Ring> coeffs);
/// One PRG lane per blind role, plus the share-mask stream for that role. namespace detail
/// `blind(role, index)` and `share(role, index, value)` do not depend on {
template <typename Ring, typename ContX, typename ContY>
expr<Ring> dot_expr(const ContX & xs, const ContY & ys);
} // namespace detail
/// @brief One PRG lane per blind role, plus the share-mask stream for that role.
/// @details `blind(role, index)` and `share(role, index, value)` do not depend on
/// call order. Walking `index` forward stays inside a refilled window. /// call order. Walking `index` forward stays inside a refilled window.
/// `PRG` defaults to `dpf::prg::aes128`. /// `PRG` defaults to `dpf::prg::aes128`.
/// @tparam Ring payload ring
/// @tparam PRG pseudorandom generator
template <typename Ring, typename PRG = dpf::prg::aes128> template <typename Ring, typename PRG = dpf::prg::aes128>
class oracle class oracle
{ {
@ -241,7 +258,7 @@ public:
using seed_type = typename PRG::block_type; using seed_type = typename PRG::block_type;
using traits = ring_traits<Ring>; using traits = ring_traits<Ring>;
/// Monomial roles sit above wire ids. Dot-cross roles sit above those. /// @brief Monomial roles sit above wire ids. Dot-cross roles sit above those.
static constexpr std::uint32_t mono_role_base = 0x40000000u; static constexpr std::uint32_t mono_role_base = 0x40000000u;
static constexpr std::uint32_t dot_role_base = 0x80000000u; static constexpr std::uint32_t dot_role_base = 0x80000000u;
@ -260,7 +277,7 @@ public:
return dot_role_base + gate; return dot_role_base + gate;
} }
/// Fused within-polynomial λ combinations (Appendix E groupings). /// @brief Fused within-polynomial λ combinations (Appendix E groupings).
static constexpr std::uint32_t bundle_role_base = 0xC0000000u; static constexpr std::uint32_t bundle_role_base = 0xC0000000u;
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -290,7 +307,11 @@ public:
return lanes_.mask_at(role, index); return lanes_.mask_at(role, index);
} }
/// Additive split of `value`. The mask is the role's mask stream at `index`. /// @brief Additive split of `value`. The mask is the role's mask stream at `index`.
/// @param role the `role`
/// @param index the index
/// @param value the value to convert or store
/// @return Additive split of `value`
split<Ring> share(std::uint32_t role, std::uint64_t index, const Ring & value) const split<Ring> share(std::uint32_t role, std::uint64_t index, const Ring & value) const
{ {
Ring p0 = mask(role, index); Ring p0 = mask(role, index);
@ -311,20 +332,22 @@ private:
dpf::randomness::lane_table<Ring, PRG> lanes_; dpf::randomness::lane_table<Ring, PRG> lanes_;
}; };
/// Shares produced for one copy index of a recorded formula. /// @brief Shares produced for one copy index of a recorded formula.
/// @tparam Ring payload ring
template <typename Ring> template <typename Ring>
struct prg_material struct prg_material
{ {
std::vector<split<Ring>> lambda; std::vector<split<Ring>> lambda;
std::vector<split<Ring>> monomial; std::vector<split<Ring>> monomial;
/// Fused λ-combinations for polynomial gates, in bundle index order. /// @brief Fused λ-combinations for polynomial gates, in bundle index order.
std::vector<split<Ring>> bundles; std::vector<split<Ring>> bundles;
/// Parallel to the session's gates. Empty split when the gate is not a dot. /// @brief Parallel to the session's gates. Empty split when the gate is not a dot.
std::vector<split<Ring>> dot_cross; std::vector<split<Ring>> dot_cross;
}; };
/// Dealer session: record formulae, `sample` blinds and monomials, `bind` /// @brief Dealer session: record formulae, `sample` blinds and monomials, `bind`
/// input secrets, `evaluate` every round. /// input secrets, `evaluate` every round.
/// @tparam Ring payload ring
template <typename Ring> template <typename Ring>
class session class session
{ {
@ -338,7 +361,7 @@ public:
using exp_list = std::vector<std::pair<std::uint32_t, std::uint8_t>>; using exp_list = std::vector<std::pair<std::uint32_t, std::uint8_t>>;
using wire = ::dpf::beavers::wire<Ring>; using wire = ::dpf::beavers::wire<Ring>;
/// One factor of a monomial query: `{{x, 2}, {a, 1}}`. /// @brief One factor of a monomial query: `{{x, 2}, {a, 1}}`.
struct power struct power
{ {
wire base{}; wire base{};
@ -351,29 +374,34 @@ public:
session(session &&) = delete; session(session &&) = delete;
session & operator=(session &&) = delete; session & operator=(session &&) = delete;
/// Arithmetic input. Its blind is sampled once and then reused. /// @brief Arithmetic input. Its blind is sampled once and then reused.
/// @return Arithmetic input
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire input() wire input()
{ {
return emplace_wire(0, true, false); return emplace_wire(0, true, false);
} }
/// Sample this wire's blind even if no recorded formula opens it. /// @brief Sample this wire's blind even if no recorded formula opens it.
/// One-shot triples use this for the product wire, so it can be reused. /// @details One-shot triples use this for the product wire, so it can be reused.
/// @param w the `w`
void pin(wire w) void pin(wire w)
{ {
wires_[check(w)].pinned = true; wires_[check(w)].pinned = true;
} }
/// 0/1 wire in this ring. `bind` accepts only `zero()` or `one()` /// @brief 0/1 wire in this ring. `bind` accepts only `zero()` or `one()`
/// (`1` for integer rings, the all-ones word for `xor_wrapper`). /// (`1` for integer rings, the all-ones word for `xor_wrapper`).
/// @return 0/1 wire in this ring
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire bit() wire bit()
{ {
return emplace_wire(0, true, true); return emplace_wire(0, true, true);
} }
/// One-round product. Repeated wires share a blind. /// @brief One-round product. Repeated wires share a blind.
/// @param factors the `factors`
/// @return One-round product
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire product(std::initializer_list<wire> factors) wire product(std::initializer_list<wire> factors)
{ {
@ -396,7 +424,10 @@ public:
return commit_product({check(a), check(b), check(c)}); return commit_product({check(a), check(b), check(c)});
} }
/// Record a sum of monomials. Like terms share one blind product. /// @brief Record a sum of monomials. Like terms share one blind product.
/// @param e the `e`
/// @return Record a sum of monomials
/// @throws std::invalid_argument if `beaver expression is from a different session`
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire operator()(const expr<Ring> & e) wire operator()(const expr<Ring> & e)
{ {
@ -405,15 +436,23 @@ public:
return commit_poly(e); return commit_poly(e);
} }
/// `c[0] + c[1] x + c[2] x^2 + ...` in one round. /// @brief `c[0] + c[1] x + c[2] x^2 + ...` in one round.
/// @param x the `x`
/// @param coeffs the public coefficients
/// @return `c[0] + c[1] x + c[2] x^2 + ...` in one round
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire horner(wire x, std::initializer_list<Ring> coeffs) wire horner(wire x, std::initializer_list<Ring> coeffs)
{ {
return (*this)(horner_expr(x, coeffs)); return (*this)(horner_expr(x, coeffs));
} }
/// Sign-corrected Horner: `sign * (c[0] + c[1] x + ...)`, still one round /// @brief Sign-corrected Horner: `sign * (c[0] + c[1] x + ...)`, still one round
/// when `sign` and `x` are inputs. /// when `sign` and `x` are inputs.
/// @param sign the sign bit or sign value
/// @param x the `x`
/// @param coeffs the public coefficients
/// @return Sign-corrected Horner: `sign * (c[0] + c[1] x + ...)`, still one round when `sign`
/// and `x` are inputs
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire horner(wire sign, wire x, std::initializer_list<Ring> coeffs) wire horner(wire sign, wire x, std::initializer_list<Ring> coeffs)
{ {
@ -426,7 +465,10 @@ public:
return product(x, x); return product(x, x);
} }
/// One-round `a * x * x` (one blind for `x`). /// @brief One-round `a * x * x` (one blind for `x`).
/// @param a the `a`
/// @param x the `x`
/// @return One-round `a * x * x` (one blind for `x`)
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire mul_square(wire a, wire x) wire mul_square(wire a, wire x)
{ {
@ -437,31 +479,21 @@ public:
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire dot(const ContX & xs, const ContY & ys) wire dot(const ContX & xs, const ContY & ys)
{ {
std::vector<std::uint32_t> x; return finish_dot(detail::dot_expr<Ring>(xs, ys));
std::vector<std::uint32_t> y;
for (const auto & w : xs)
x.push_back(check(w));
for (const auto & w : ys)
y.push_back(check(w));
return commit_dot(std::move(x), std::move(y));
} }
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire dot(std::initializer_list<wire> xs, std::initializer_list<wire> ys) wire dot(std::initializer_list<wire> xs, std::initializer_list<wire> ys)
{ {
std::vector<std::uint32_t> x; return finish_dot(detail::dot_expr<Ring>(xs, ys));
std::vector<std::uint32_t> y;
x.reserve(xs.size());
y.reserve(ys.size());
for (auto w : xs)
x.push_back(check(w));
for (auto w : ys)
y.push_back(check(w));
return commit_dot(std::move(x), std::move(y));
} }
/// `z_i = scalar * lanes[i]`, one output wire per lane. The scalar blind /// @brief `z_i = scalar * lanes[i]`, one output wire per lane. The scalar blind
/// is shared. Each lane gets its own `λ_s λ_i` share. /// is shared. Each lane gets its own `λ_s λ_i` share.
/// @tparam Cont cont
/// @param scalar the `scalar`
/// @param lanes the lane values
/// @return `z_i = scalar * lanes[i]`, one output wire per lane
template <typename Cont> template <typename Cont>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
std::vector<wire> scale(wire scalar, const Cont & lanes) std::vector<wire> scale(wire scalar, const Cont & lanes)
@ -482,7 +514,11 @@ public:
return out; return out;
} }
/// `bit * scalar`. `bit` must come from `bit()`. /// @brief `bit * scalar`. `bit` must come from `bit()`.
/// @param selector the `selector`
/// @param scalar the `scalar`
/// @return `bit * scalar`
/// @throws std::invalid_argument if `bit_mul selector must come from bit()`
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire bit_mul(wire selector, wire scalar) wire bit_mul(wire selector, wire scalar)
{ {
@ -492,7 +528,11 @@ public:
return commit_product({id, check(scalar)}); return commit_product({id, check(scalar)});
} }
/// One-round `selector ? when1 : when0`, i.e. `when0 + selector * (when1 - when0)`. /// @brief One-round `selector ? when1 : when0`, i.e. `when0 + selector * (when1 - when0)`.
/// @param selector the `selector`
/// @param when1 the `when1`
/// @param when0 the `when0`
/// @return One-round `selector ? when1 : when0`, i.e
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
wire mux(wire selector, wire when1, wire when0) wire mux(wire selector, wire when1, wire when0)
{ {
@ -515,8 +555,11 @@ public:
return out; return out;
} }
/// Sample every missing wire blind and every missing monomial. /// @brief Sample every missing wire blind and every missing monomial.
/// Blinds already sampled are left alone. /// @details Blinds already sampled are left alone.
/// @tparam Sample sample
/// @param sampler the randomness sampler
/// @throws std::logic_error if `beaver blind is missing`
template <typename Sample> template <typename Sample>
void sample(Sample && sampler) void sample(Sample && sampler)
{ {
@ -572,9 +615,12 @@ public:
sample(default_sampler<Ring>{}); sample(default_sampler<Ring>{});
} }
/// Install missing blinds and product shares from copy `index` of `src`. /// @brief Install missing blinds and product shares from copy `index` of `src`.
/// Already-sampled wires keep their λ. New product shares are built from /// @details Already-sampled wires keep their λ. New product shares are built from
/// those stored blinds, then split with the oracle's share lane. /// those stored blinds, then split with the oracle's share lane.
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param index the index
template <typename PRG> template <typename PRG>
void sample_from(const oracle<Ring, PRG> & src, std::uint64_t index = 0) void sample_from(const oracle<Ring, PRG> & src, std::uint64_t index = 0)
{ {
@ -614,9 +660,13 @@ public:
} }
} }
/// Every wire, monomial, and dot cross of this formula at copy `index`. /// @brief Every wire, monomial, and dot cross of this formula at copy `index`.
/// Does not change the session. Copies are independent lanes samples, so /// @details Does not change the session. Copies are independent lanes samples, so
/// `material_at(src, 5)` does not depend on having asked for 0..4. /// `material_at(src, 5)` does not depend on having asked for 0..4.
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param index the index
/// @return Every wire, monomial, and dot cross of this formula at copy `index`
template <typename PRG> template <typename PRG>
prg_material<Ring> material_at(const oracle<Ring, PRG> & src, std::uint64_t index) const prg_material<Ring> material_at(const oracle<Ring, PRG> & src, std::uint64_t index) const
{ {
@ -665,7 +715,13 @@ public:
return out; return out;
} }
/// Split `secret` into fresh additive shares and bind them to an input. /// @brief Split `secret` into fresh additive shares and bind them to an input.
/// @tparam Sample sample
/// @param w the `w`
/// @param secret the secret value
/// @param sampler the randomness sampler
/// @throws std::invalid_argument if `only input wires can be bound`
/// @throws std::logic_error if `input wire is already bound`
template <typename Sample> template <typename Sample>
void bind(wire w, Ring secret, Sample && sampler) void bind(wire w, Ring secret, Sample && sampler)
{ {
@ -686,7 +742,12 @@ public:
bind(w, secret, default_sampler<Ring>{}); bind(w, secret, default_sampler<Ring>{});
} }
/// Bind shares the caller already holds. Their sum is the secret. /// @brief Bind shares the caller already holds. Their sum is the secret.
/// @param w the `w`
/// @param p0 the `p0`
/// @param p1 the `p1`
/// @throws std::invalid_argument if `only input wires can be bound`
/// @throws std::logic_error if `input wire is already bound`
void bind_shares(wire w, Ring p0, Ring p1) void bind_shares(wire w, Ring p0, Ring p1)
{ {
auto id = check(w); auto id = check(w);
@ -700,9 +761,10 @@ public:
wires_[id].value_ready = true; wires_[id].value_ready = true;
} }
/// Open every ready round. Inputs used by round-1 gates are opened /// @brief Open every ready round. Inputs used by round-1 gates are opened
/// together; a gate output is a later round's input and keeps the blind /// together; a gate output is a later round's input and keeps the blind
/// chosen in `sample`. /// chosen in `sample`.
/// @throws std::logic_error if `beaver wire is not ready to open`
void evaluate() void evaluate()
{ {
int max_round = 0; int max_round = 0;
@ -774,7 +836,10 @@ public:
return wires_[id].lambda; return wires_[id].lambda;
} }
/// Share of `Π λ_i^{e_i}`. A lone `λ_w` is the wire blind itself. /// @brief Share of `Π λ_i^{e_i}`. A lone `λ_w` is the wire blind itself.
/// @param spec the specification
/// @return Share of `Π λ_i^{e_i}`
/// @throws std::invalid_argument if `beaver power is zero`
split<Ring> monomial(const std::vector<power> & spec) const split<Ring> monomial(const std::vector<power> & spec) const
{ {
std::vector<std::pair<std::uint32_t, unsigned>> raw; std::vector<std::pair<std::uint32_t, unsigned>> raw;
@ -807,7 +872,10 @@ public:
return traits::add(v.p0, v.p1); return traits::add(v.p0, v.p1);
} }
/// Public ABY2.0 mask δ = x + λ, after `evaluate` has opened the wire. /// @brief Public ABY2.0 mask δ = x + λ, after `evaluate` has opened the wire.
/// @param w the `w`
/// @return Public ABY2.0 mask δ = x + λ, after `evaluate` has opened the wire
/// @throws std::logic_error if `beaver wire has not been opened`
Ring delta(wire w) const Ring delta(wire w) const
{ {
auto id = check(w); auto id = check(w);
@ -819,12 +887,35 @@ public:
split<Ring> dot_cross(wire w) const split<Ring> dot_cross(wire w) const
{ {
auto id = check(w); auto id = check(w);
int g = wires_[id].gate; if (!wires_[id].dot_output)
if (g < 0 || gates_[static_cast<std::size_t>(g)].kind != gate_kind::dot)
throw std::invalid_argument("wire is not a dot output"); throw std::invalid_argument("wire is not a dot output");
if (!gates_[static_cast<std::size_t>(g)].cross_ready) int g = wires_[id].gate;
if (g < 0)
throw std::invalid_argument("wire is not a dot output");
const auto & gate = gates_[static_cast<std::size_t>(g)];
if (gate.kind == gate_kind::dot)
{
if (!gate.cross_ready)
throw std::logic_error("call sample() before reading a dot cross term"); throw std::logic_error("call sample() before reading a dot cross term");
return gates_[static_cast<std::size_t>(g)].cross; return gate.cross;
}
Ring s0 = traits::zero();
Ring s1 = traits::zero();
bool found = false;
for (const auto & step : gate.steps)
{
if (step.bundle < 0 || !step.delta.empty())
continue;
const auto & bundle = bundles_[static_cast<std::size_t>(step.bundle)];
if (!bundle.ready)
throw std::logic_error("call sample() before reading a dot cross term");
s0 = traits::add(s0, traits::mul(step.scale, bundle.share.p0));
s1 = traits::add(s1, traits::mul(step.scale, bundle.share.p1));
found = true;
}
if (!found)
throw std::logic_error("call sample() before reading a dot cross term");
return split<Ring>{s0, s1};
} }
int round_of(wire w) const int round_of(wire w) const
@ -835,16 +926,19 @@ public:
HEDLEY_NO_THROW HEDLEY_NO_THROW
std::size_t wire_count() const noexcept { return wires_.size(); } std::size_t wire_count() const noexcept { return wires_.size(); }
/// Product shares beyond the per-wire blinds: subset monomials from /// @brief Product shares beyond the per-wire blinds: subset monomials from
/// `product` gates, plus one fused bundle per public-δ class in a /// `product` gates, plus one fused bundle per public-δ class in a
/// polynomial (Appendix E). A lone mask is not counted. /// polynomial (Appendix E). A lone mask is not counted.
/// @return Product shares beyond the per-wire blinds: subset monomials from `product` gates,
/// plus one fused bundle per public-δ class in a polynomial (Appendix E)
HEDLEY_NO_THROW HEDLEY_NO_THROW
std::size_t monomial_count() const noexcept std::size_t monomial_count() const noexcept
{ {
return monos_.size() + bundles_.size(); return monos_.size() + bundles_.size();
} }
/// Wire blinds that the recorded formulae actually open, plus product shares. /// @brief Wire blinds that the recorded formulae actually open, plus product shares.
/// @return Wire blinds that the recorded formulae actually open, plus product shares
std::size_t preprocessing_count() const std::size_t preprocessing_count() const
{ {
std::size_t n = monomial_count(); std::size_t n = monomial_count();
@ -865,7 +959,7 @@ private:
std::vector<std::uint32_t> factors; std::vector<std::uint32_t> factors;
}; };
/// One λ-monomial in a fused preprocessing share. /// @brief One λ-monomial in a fused preprocessing share.
struct bundle_part struct bundle_part
{ {
Ring coeff{}; Ring coeff{};
@ -879,7 +973,7 @@ private:
bool ready = false; bool ready = false;
}; };
/// Online: `public(δ) * scale * share`, where share is a wire mask, /// @brief Online: `public(δ) * scale * share`, where share is a wire mask,
/// a raw monomial, or a fused sum of monomials. /// a raw monomial, or a fused sum of monomials.
struct poly_step struct poly_step
{ {
@ -896,6 +990,7 @@ private:
bool is_input = false; bool is_input = false;
bool is_bit = false; bool is_bit = false;
bool pinned = false; bool pinned = false;
bool dot_output = false;
bool lambda_ready = false; bool lambda_ready = false;
bool value_ready = false; bool value_ready = false;
bool delta_ready = false; bool delta_ready = false;
@ -1369,25 +1464,10 @@ private:
return last; return last;
} }
wire commit_dot(std::vector<std::uint32_t> xs, std::vector<std::uint32_t> ys) wire finish_dot(const expr<Ring> & e)
{ {
if (xs.empty() || xs.size() != ys.size()) auto out = (*this)(e);
throw std::invalid_argument( wires_[out.id_].dot_output = true;
"beaver dot operands must have the same non-zero length");
int round = 1;
for (std::size_t i = 0; i < xs.size(); ++i)
{
round = std::max(round, wires_[xs[i]].ready_round + 1);
round = std::max(round, wires_[ys[i]].ready_round + 1);
}
auto out = emplace_wire(round, false, false);
gate g;
g.kind = gate_kind::dot;
g.out = out.id_;
g.lhs = std::move(xs);
g.rhs = std::move(ys);
gates_.push_back(std::move(g));
wires_[out.id_].gate = static_cast<int>(gates_.size() - 1);
return out; return out;
} }
@ -1572,8 +1652,10 @@ private:
return false; return false;
} }
/// Group λ-monomials that share a public δ monomial into one share. /// @brief Group λ-monomials that share a public δ monomial into one share.
/// A bucket that is only `c · λ_i` reuses the wire blind. /// @details A bucket that is only `c · λ_i` reuses the wire blind.
/// @param terms the polynomial terms
/// @return Group λ-monomials that share a public δ monomial into one share
std::vector<poly_step> compile_poly(const std::vector<poly_term> & terms) std::vector<poly_step> compile_poly(const std::vector<poly_term> & terms)
{ {
struct bucket struct bucket
@ -1665,9 +1747,20 @@ private:
steps.push_back(std::move(step)); steps.push_back(std::move(step));
continue; continue;
} }
Ring scale = parts[0].coeff;
for (const auto & part : parts)
{
if (!(part.coeff == scale))
scale = traits::one();
}
if (!(scale == traits::one()))
{
for (auto & part : parts)
part.coeff = traits::one();
}
poly_step step; poly_step step;
step.delta = delta; step.delta = delta;
step.scale = traits::one(); step.scale = scale;
step.bundle = require_bundle(std::move(parts)); step.bundle = require_bundle(std::move(parts));
steps.push_back(std::move(step)); steps.push_back(std::move(step));
} }
@ -1852,8 +1945,40 @@ Ring coeff_of(Coeff value)
return Ring{value}; return Ring{value};
} }
template <typename Ring, typename ContX, typename ContY>
expr<Ring> dot_expr(const ContX & xs, const ContY & ys)
{
std::vector<wire<Ring>> x;
std::vector<wire<Ring>> y;
for (const auto & w : xs)
x.push_back(w);
for (const auto & w : ys)
y.push_back(w);
if (x.empty() || x.size() != y.size())
throw std::invalid_argument(
"beaver dot operands must have the same non-zero length");
expr<Ring> acc = wire_expr(x[0]) * wire_expr(y[0]);
for (std::size_t i = 1; i < x.size(); ++i)
acc = add_exprs(std::move(acc), wire_expr(x[i]) * wire_expr(y[i]));
return acc;
}
} // namespace detail } // namespace detail
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> dot(std::initializer_list<wire<Ring>> xs, std::initializer_list<wire<Ring>> ys)
{
return detail::dot_expr<Ring>(xs, ys);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
expr<Ring> dot(const std::vector<wire<Ring>> & xs, const std::vector<wire<Ring>> & ys)
{
return detail::dot_expr<Ring>(xs, ys);
}
template <typename Ring> template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
expr<Ring> wire_expr(wire<Ring> w) expr<Ring> wire_expr(wire<Ring> w)
@ -1895,7 +2020,13 @@ expr<Ring> horner_expr(wire<Ring> x, std::initializer_list<Ring> coeffs)
return e; return e;
} }
/// `coeff * v0 * v1 * ...`, with repeated wires counting as a power. /// @brief `coeff * v0 * v1 * ...`, with repeated wires counting as a power.
/// @tparam Ring payload ring
/// @tparam Wires wires
/// @param coeff the public coefficient
/// @param first the first element of the range
/// @param rest the remaining arguments
/// @return `coeff * v0 * v1 * ...`, with repeated wires counting as a power
template <typename Ring, typename... Wires> template <typename Ring, typename... Wires>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
expr<Ring> monomial(Ring coeff, wire<Ring> first, Wires... rest) expr<Ring> monomial(Ring coeff, wire<Ring> first, Wires... rest)
@ -2128,8 +2259,10 @@ expr<Ring> operator+(expr<Ring> e, Coeff coeff)
// `out` is the ABY2.0 blind of the product wire, for a later round. // `out` is the ABY2.0 blind of the product wire, for a later round.
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
/// `subset[mask - 1]` is `Π λ_i` over the bits set in `mask` (bit i selects /// @brief `subset[mask - 1]` is `Π λ_i` over the bits set in `mask` (bit i selects
/// `in[i]`). Singleton masks are the wire blinds themselves. /// `in[i]`). Singleton masks are the wire blinds themselves.
/// @tparam Arity arity
/// @tparam Ring payload ring
template <std::size_t Arity, typename Ring> template <std::size_t Arity, typename Ring>
struct fresh_beaver struct fresh_beaver
{ {
@ -2259,9 +2392,14 @@ struct mux_beaver
split<Ring> out{}; split<Ring> out{};
}; };
/// One fresh Beaver pair from copy `index` of an oracle. /// @brief One fresh Beaver pair from copy `index` of an oracle.
/// Roles match a session that records `input, input, product`: wires 0 and 1, /// @details Roles match a session that records `input, input, product`: wires 0 and 1,
/// the product wire, and monomial 0. /// the product wire, and monomial 0.
/// @tparam Ring payload ring
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param index the index
/// @return One fresh Beaver pair from copy `index` of an oracle
template <typename Ring, typename PRG = dpf::prg::aes128> template <typename Ring, typename PRG = dpf::prg::aes128>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
beaver2<Ring> beaver2_at(const oracle<Ring, PRG> & src, std::uint64_t index) beaver2<Ring> beaver2_at(const oracle<Ring, PRG> & src, std::uint64_t index)
@ -2278,7 +2416,14 @@ beaver2<Ring> beaver2_at(const oracle<Ring, PRG> & src, std::uint64_t index)
src.share(oracle<Ring>::wire_role(2), index, out)}; src.share(oracle<Ring>::wire_role(2), index, out)};
} }
/// `n` copies starting at `begin`. Each role is one contiguous lane read. /// @brief `n` copies starting at `begin`. Each role is one contiguous lane read.
/// @tparam Ring payload ring
/// @tparam PRG pseudorandom generator
/// @param src the source
/// @param begin the iterator to the first query
/// @param out the output buffer
/// @param n the `n`
/// @throws std::invalid_argument if `beaver2 output is null`
template <typename Ring, typename PRG = dpf::prg::aes128> template <typename Ring, typename PRG = dpf::prg::aes128>
void fill_beaver2(const oracle<Ring, PRG> & src, std::uint64_t begin, void fill_beaver2(const oracle<Ring, PRG> & src, std::uint64_t begin,
beaver2<Ring> * out, std::size_t n) beaver2<Ring> * out, std::size_t n)

View file

@ -52,7 +52,7 @@ enum bit : bool
/// equal to `dpf::bit::one` if the *least-significant bit* of /// equal to `dpf::bit::one` if the *least-significant bit* of
/// `value` is `1` and `dpf::bit::zero` otherwise. /// `value` is `1` and `dpf::bit::zero` otherwise.
/// @param value the `int` to convert /// @param value the `int` to convert
/// @returns `static_cast<dpf::bit>(value & 1)` /// @return `static_cast<dpf::bit>(value & 1)`
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -64,6 +64,8 @@ static constexpr dpf::bit to_bit(int value) noexcept
/// @brief converts the least-significant bit of an integer literal to a `dpf::bit` /// @brief converts the least-significant bit of an integer literal to a `dpf::bit`
/// @details This overload exists so `operator""_bit` does not select the /// @details This overload exists so `operator""_bit` does not select the
/// character converter, which is an exact match for `unsigned long long`. /// character converter, which is an exact match for `unsigned long long`.
/// @param value the value to convert or store
/// @return the returned `dpf::bit`
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -77,7 +79,7 @@ static constexpr dpf::bit to_bit(unsigned long long value) noexcept
/// equal to `dpf::bit::one` if `value==true` and `dpf::bit::zero` /// equal to `dpf::bit::one` if `value==true` and `dpf::bit::zero`
/// otherwise. /// otherwise.
/// @param value the `bool` to convert /// @param value the `bool` to convert
/// @returns `static_cast<dpf::bit>(value)` /// @return `static_cast<dpf::bit>(value)`
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -90,10 +92,12 @@ static constexpr dpf::bit to_bit(bool value) noexcept
/// @details Convert a character to a `dpf::bit`. The resulting `dpf::bit` is /// @details Convert a character to a `dpf::bit`. The resulting `dpf::bit` is
/// equal to `dpf::bit::one` if `value==one` and `dpf::bit::zero` /// equal to `dpf::bit::one` if `value==one` and `dpf::bit::zero`
/// otherwise. /// otherwise.
/// @tparam CharT character type
/// @tparam Traits character traits
/// @param value the character to convert /// @param value the character to convert
/// @param zero character used to represent `0` (default: ``CharT('0')``) /// @param zero character used to represent `0` (default: ``CharT('0')``)
/// @param one character used to represent `1` (default: ``CharT('1')``) /// @param one character used to represent `1` (default: ``CharT('1')``)
/// @returns `static_cast<dpf::bit>(0)` if `value==0` or /// @return `static_cast<dpf::bit>(0)` if `value==0` or
/// `static_cast<dpf::bit>(1)` if `value==1` /// `static_cast<dpf::bit>(1)` if `value==1`
/// @throws std::domain_error if `value != zero && value != one` /// @throws std::domain_error if `value != zero && value != one`
template <typename CharT, template <typename CharT,
@ -117,6 +121,9 @@ static constexpr dpf::bit to_bit(
/// @details Converts the contents of a `dpf::bit` to a `std::string` for /// @details Converts the contents of a `dpf::bit` to a `std::string` for
/// human-friendly printing. Uses `zero` to represent the value /// human-friendly printing. Uses `zero` to represent the value
/// `0` and `one` to the value `1`. /// `0` and `one` to the value `1`.
/// @tparam CharT character type
/// @tparam Traits character traits
/// @tparam Allocator allocator type
/// @param value the `dpf::bit` to convert /// @param value the `dpf::bit` to convert
/// @param zero character to use to represent `false`/`0` (default: ``CharT('0')``) /// @param zero character to use to represent `false`/`0` (default: ``CharT('0')``)
/// @param one character to use to represent `true`/`1` (default: ``CharT('1')``) /// @param one character to use to represent `true`/`1` (default: ``CharT('1')``)
@ -144,6 +151,8 @@ static std::basic_string<CharT, Traits, Allocator> to_string(
/// characters to use for zero and one are obtained from the /// characters to use for zero and one are obtained from the
/// currently-imbued locale by calling `os.widen()` with `0` and `1` /// currently-imbued locale by calling `os.widen()` with `0` and `1`
/// as the arguments. /// as the arguments.
/// @tparam CharT character type
/// @tparam Traits character traits
/// @param os a character output stream /// @param os a character output stream
/// @param value the `dpf::bit` to insert into the output stream /// @param value the `dpf::bit` to insert into the output stream
/// @return `os` /// @return `os`
@ -162,6 +171,8 @@ operator<<(std::basic_ostream<CharT, Traits> & os, const dpf::bit & value)
/// stored in `value`. The characters to use for zero and one are /// stored in `value`. The characters to use for zero and one are
/// obtained from the currently-imbued locale by calling `is.widen()` /// obtained from the currently-imbued locale by calling `is.widen()`
/// with `0` and `1` as the arguments. /// with `0` and `1` as the arguments.
/// @tparam CharT character type
/// @tparam Traits character traits
/// @param is a character input stream /// @param is a character input stream
/// @param value the `dpf::bit` to extract from the input stream /// @param value the `dpf::bit` to extract from the input stream
/// @return `is` /// @return `is`
@ -190,6 +201,9 @@ inline constexpr dpf::bit operator+(dpf::bit lhs, dpf::bit rhs) noexcept
} }
/// @brief GF(2) subtraction. Identical to `operator+`. /// @brief GF(2) subtraction. Identical to `operator+`.
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return GF(2) subtraction
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline constexpr dpf::bit operator-(dpf::bit lhs, dpf::bit rhs) noexcept inline constexpr dpf::bit operator-(dpf::bit lhs, dpf::bit rhs) noexcept
{ {

View file

@ -1,6 +1,5 @@
/// @file dpf/bit_array.hpp /// @file dpf/bit_array.hpp
/// @brief /// @brief Packed bit arrays, static and dynamic, with bit proxies and iterators.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -68,6 +67,8 @@ class const_bit_iterator; // forward reference
/// @brief a base class for classes representing a sequence of bits /// @brief a base class for classes representing a sequence of bits
/// @details A `bit_array` represents a sequence of bits. The underlying /// @details A `bit_array` represents a sequence of bits. The underlying
/// storage is an array of integers of type `dpf::bit_array::word_type`. /// storage is an array of integers of type `dpf::bit_array::word_type`.
/// @tparam ConcreteBitArrayT concrete bit array type
/// @tparam WordT word used to pack bits
template <typename ConcreteBitArrayT, typename WordT = psnip_uint64_t> template <typename ConcreteBitArrayT, typename WordT = psnip_uint64_t>
class bit_array_base class bit_array_base
{ {
@ -143,10 +144,12 @@ class bit_array_base
~bit_array_base() = default; ~bit_array_base() = default;
/// @brief default copy assignment /// @brief default copy assignment
/// @return `*this`
inline constexpr inline constexpr
bit_array_base & operator=(const bit_array_base &) = default; bit_array_base & operator=(const bit_array_base &) = default;
/// @brief defaulted move assignment /// @brief defaulted move assignment
/// @return `*this`
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline constexpr inline constexpr
bit_array_base & operator=(bit_array_base &&) noexcept = default; bit_array_base & operator=(bit_array_base &&) noexcept = default;
@ -202,7 +205,7 @@ class bit_array_base
/// significant to most significant) /// significant to most significant)
/// @note Unlike `test` and `at`, does not throw exceptions: the behavior /// @note Unlike `test` and `at`, does not throw exceptions: the behavior
/// is undefined if `pos` is out of bounds /// is undefined if `pos` is out of bounds
/// @returns an object of type `dpf::bit_array_base::reference`, which /// @return an object of type `dpf::bit_array_base::reference`, which
/// allows writing to the requested bit /// allows writing to the requested bit
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -215,9 +218,9 @@ class bit_array_base
/// @details accesses the bit at position `pos` /// @details accesses the bit at position `pos`
/// @param pos the 0-based position of the bit to return (least /// @param pos the 0-based position of the bit to return (least
/// significant to most significant) /// significant to most significant)
/// @return the value of the requested bit
/// @note Unlike `test` and `at`, does not throw exceptions: the behavior /// @note Unlike `test` and `at`, does not throw exceptions: the behavior
/// is undefined if `pos` is out of bounds /// is undefined if `pos` is out of bounds
/// @returns the value of the requested bit
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline constexpr const_reference operator[](size_type pos) const noexcept inline constexpr const_reference operator[](size_type pos) const noexcept
@ -235,7 +238,7 @@ class bit_array_base
/// significant to most significant) /// significant to most significant)
/// @throws std::out_of_range if `pos` does not correspond to a valid /// @throws std::out_of_range if `pos` does not correspond to a valid
/// position within the `bit_array_base` /// position within the `bit_array_base`
/// @returns an object of type `dpf::bit_array_base::reference`, which /// @return an object of type `dpf::bit_array_base::reference`, which
/// allows writing to the requested bit /// allows writing to the requested bit
/// @complexity `O(1)` /// @complexity `O(1)`
constexpr reference at(size_type pos) constexpr reference at(size_type pos)
@ -249,9 +252,9 @@ class bit_array_base
/// @details accesses the bit at position `pos` /// @details accesses the bit at position `pos`
/// @param pos the 0-based position of the bit to return (least /// @param pos the 0-based position of the bit to return (least
/// significant to most significant) /// significant to most significant)
/// @return the value of the requested bit
/// @throws std::out_of_range if `pos` does not correspond to a valid /// @throws std::out_of_range if `pos` does not correspond to a valid
/// position within the `bit_array_base` /// position within the `bit_array_base`
/// @returns the value of the requested bit
/// @complexity `O(1)` /// @complexity `O(1)`
constexpr const_reference at(size_type pos) const constexpr const_reference at(size_type pos) const
{ {
@ -264,7 +267,7 @@ class bit_array_base
/// @brief returns an iterator to the first bit /// @brief returns an iterator to the first bit
/// @{ /// @{
/// @returns iterator to the first element /// @return iterator to the first element
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr iterator begin() noexcept constexpr iterator begin() noexcept
@ -273,7 +276,7 @@ class bit_array_base
if (p == nullptr) return iterator{}; if (p == nullptr) return iterator{};
return iterator{p, word_type(1)}; return iterator{p, word_type(1)};
} }
/// @returns iterator to the first element /// @return iterator to the first element
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr const_iterator begin() const noexcept constexpr const_iterator begin() const noexcept
@ -282,7 +285,7 @@ class bit_array_base
if (p == nullptr) return const_iterator{}; if (p == nullptr) return const_iterator{};
return const_iterator{p, word_type(1)}; return const_iterator{p, word_type(1)};
} }
/// @returns iterator to the first element /// @return iterator to the first element
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr const_iterator cbegin() const noexcept constexpr const_iterator cbegin() const noexcept
@ -293,7 +296,7 @@ class bit_array_base
/// @brief returns an iterator to the end (one past the last bit) /// @brief returns an iterator to the end (one past the last bit)
/// @{ /// @{
/// @returns iterator to the element following the last element /// @return iterator to the element following the last element
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr iterator end() noexcept constexpr iterator end() noexcept
@ -303,7 +306,7 @@ class bit_array_base
return iterator{p + (size() >> lg_bits_per_word), return iterator{p + (size() >> lg_bits_per_word),
static_cast<word_type>(word_type(1) << (size() % bits_per_word))}; static_cast<word_type>(word_type(1) << (size() % bits_per_word))};
} }
/// @returns iterator to the element following the last element /// @return iterator to the element following the last element
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr const_iterator end() const noexcept constexpr const_iterator end() const noexcept
@ -313,7 +316,7 @@ class bit_array_base
return const_iterator{p + (size() >> lg_bits_per_word), return const_iterator{p + (size() >> lg_bits_per_word),
static_cast<word_type>(word_type(1) << (size() % bits_per_word))}; static_cast<word_type>(word_type(1) << (size() % bits_per_word))};
} }
/// @returns iterator to the element following the last element /// @return iterator to the element following the last element
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr const_iterator cend() const noexcept constexpr const_iterator cend() const noexcept
@ -325,7 +328,7 @@ class bit_array_base
/// @brief checks if the specified bit is set to `true` /// @brief checks if the specified bit is set to `true`
/// @param pos the 0-based position of the bit to return (least /// @param pos the 0-based position of the bit to return (least
/// significant to most significant) /// significant to most significant)
/// @returns `true` if the requested bit is set, `false` otherwise /// @return `true` if the requested bit is set, `false` otherwise
/// @complexity `O(1)` /// @complexity `O(1)`
bool test(size_type pos) const bool test(size_type pos) const
{ {
@ -357,10 +360,9 @@ class bit_array_base
} }
/// @details checks if all bits in a range are set to `true` /// @details checks if all bits in a range are set to `true`
/// @param first,last the range of elements under consideration
/// @tparam Iterator an iterator type /// @tparam Iterator an iterator type
/// @return `true` if all of the bits in the given range are set to /// @param first,last the range of elements under consideration
/// `true`, otherwise `false` /// @return `true` if all of the bits in the given range are set to `true`, otherwise `false`
/// @complexity `O(last-first)` /// @complexity `O(last-first)`
template <typename Iterator> template <typename Iterator>
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -394,10 +396,9 @@ class bit_array_base
} }
/// @details checks if any bits in a range are set to `true` /// @details checks if any bits in a range are set to `true`
/// @param first,last the range of elements under consideration
/// @tparam Iterator an iterator type /// @tparam Iterator an iterator type
/// @return `true` if any of the bits in the given range are set to /// @param first,last the range of elements under consideration
/// `true`, otherwise `false` /// @return `true` if any of the bits in the given range are set to `true`, otherwise `false`
/// @complexity `O(last-first)` /// @complexity `O(last-first)`
template <typename Iterator> template <typename Iterator>
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -422,10 +423,9 @@ class bit_array_base
} }
/// @details checks if none bits in a range are set to `true` /// @details checks if none bits in a range are set to `true`
/// @param first,last the range of elements under consideration
/// @tparam Iterator an iterator type /// @tparam Iterator an iterator type
/// @return `true` if none of the bits in the given range are set to /// @param first,last the range of elements under consideration
/// `true`, otherwise `false` /// @return `true` if none of the bits in the given range are set to `true`, otherwise `false`
/// @complexity `O(last-first)` /// @complexity `O(last-first)`
template <typename Iterator> template <typename Iterator>
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -438,7 +438,7 @@ class bit_array_base
/// @brief returns the number of bits set to `true` /// @brief returns the number of bits set to `true`
/// @{ /// @{
/// @details counts the number of bits that are set to `true` /// @details counts the number of bits that are set to `true`
/// @returns the number of bits set to `true` /// @return the number of bits set to `true`
/// @complexity `O(size())` /// @complexity `O(size())`
HEDLEY_NO_THROW HEDLEY_NO_THROW
size_type count() const noexcept size_type count() const noexcept
@ -456,8 +456,8 @@ class bit_array_base
return sum; return sum;
} }
/// @details counts the number of bits in a range that are set to `true` /// @details counts the number of bits in a range that are set to `true`
/// @param first,last the range of elements under consideration
/// @tparam Iterator an iterator type /// @tparam Iterator an iterator type
/// @param first,last the range of elements under consideration
/// @return the number of bits in the given range that are set to `true` /// @return the number of bits in the given range that are set to `true`
/// @complexity `O(last-first)` /// @complexity `O(last-first)`
template <typename Iterator> template <typename Iterator>
@ -477,7 +477,7 @@ class bit_array_base
/// @brief returns the parity of all stored bits /// @brief returns the parity of all stored bits
/// @{ /// @{
/// @details counts the parity of all stored bits /// @details counts the parity of all stored bits
/// @returns the parity of all stored bits /// @return the parity of all stored bits
/// @complexity `O(size())` /// @complexity `O(size())`
HEDLEY_NO_THROW HEDLEY_NO_THROW
size_type parity() const noexcept size_type parity() const noexcept
@ -496,8 +496,8 @@ class bit_array_base
} }
/// @details counts the parity of bits in a range /// @details counts the parity of bits in a range
/// @param first,last the range of elements under consideration
/// @tparam Iterator an iterator type /// @tparam Iterator an iterator type
/// @param first,last the range of elements under consideration
/// @return the parity of all bits in the given range /// @return the parity of all bits in the given range
/// @complexity `O(last-first)` /// @complexity `O(last-first)`
template <typename Iterator> template <typename Iterator>
@ -515,7 +515,7 @@ class bit_array_base
/// @} /// @}
/// @brief returns the number of bits /// @brief returns the number of bits
/// @returns number of bits that the `bit_array_base` holds /// @return number of bits that the `bit_array_base` holds
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -633,9 +633,12 @@ class bit_array_base
/// contains `size()` characters with the first character /// contains `size()` characters with the first character
/// corresponding to the last `(size()-1th)` bit and the last /// corresponding to the last `(size()-1th)` bit and the last
/// character corresponding tot he first `(0th)` bit. /// character corresponding tot he first `(0th)` bit.
/// @tparam CharT character type
/// @tparam Traits character traits
/// @tparam Allocator allocator type
/// @param zero character to use to represent `false`/`0` (default: ``CharT('0')``) /// @param zero character to use to represent `false`/`0` (default: ``CharT('0')``)
/// @param one character to use to represent `true`/`1` (default: ``CharT('1')``) /// @param one character to use to represent `true`/`1` (default: ``CharT('1')``)
/// @returns the converted string /// @return the converted string
/// @throws May throw `std::bad_alloc` from the `std::string` constructor. /// @throws May throw `std::bad_alloc` from the `std::string` constructor.
/// @complexity `O(size())` /// @complexity `O(size())`
template <typename CharT = char, template <typename CharT = char,
@ -731,6 +734,9 @@ class bit_array_base
/// @brief XOR. Exact match so `bit_reference - bit_reference` is not /// @brief XOR. Exact match so `bit_reference - bit_reference` is not
/// ambiguous with integer subtraction of the proxy. /// ambiguous with integer subtraction of the proxy.
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return XOR
friend constexpr dpf::bit operator-(bit_reference lhs, bit_reference rhs) noexcept friend constexpr dpf::bit operator-(bit_reference lhs, bit_reference rhs) noexcept
{ {
return static_cast<dpf::bit>(static_cast<bool>(lhs) ^ static_cast<bool>(rhs)); return static_cast<dpf::bit>(static_cast<bool>(lhs) ^ static_cast<bool>(rhs));
@ -746,7 +752,7 @@ class bit_array_base
/// @{ /// @{
/// @details sets `*this` to the result of binary AND on `*this` and `b` /// @details sets `*this` to the result of binary AND on `*this` and `b`
/// @param b the other bit /// @param b the other bit
/// @returns `*this` /// @return `*this`
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -758,7 +764,7 @@ class bit_array_base
/// @details sets `*this` to the result of binary OR on `*this` and `b` /// @details sets `*this` to the result of binary OR on `*this` and `b`
/// @param b the other bit /// @param b the other bit
/// @returns `*this` /// @return `*this`
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -770,7 +776,7 @@ class bit_array_base
/// @details sets `*this` to the result of binary XOR on `*this` and `b` /// @details sets `*this` to the result of binary XOR on `*this` and `b`
/// @param b the other bit /// @param b the other bit
/// @returns `*this` /// @return `*this`
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -780,8 +786,8 @@ class bit_array_base
return *this; return *this;
} }
/// @details returns a temporary copy of `*this` with its value /// @details returns a temporary copy of `*this` with its value flipped (binary NOT)
/// flipped (binary NOT) /// @return the flipped bit
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -792,7 +798,7 @@ class bit_array_base
/// @} /// @}
/// @brief sets to the referenced bit to 1 /// @brief sets to the referenced bit to 1
/// @returns `*this` /// @return `*this`
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -805,7 +811,7 @@ class bit_array_base
} }
/// @brief unsets the referenced bit to 0 /// @brief unsets the referenced bit to 0
/// @returns `*this` /// @return `*this`
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -818,7 +824,8 @@ class bit_array_base
} }
/// @brief assigns `b ? 1 : 0` to the referenced bit /// @brief assigns `b ? 1 : 0` to the referenced bit
/// @returns `*this` /// @param b the `b`
/// @return `*this`
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -832,7 +839,7 @@ class bit_array_base
} }
/// @brief flips the referenced bit /// @brief flips the referenced bit
/// @returns `*this` /// @return `*this`
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -846,6 +853,8 @@ class bit_array_base
/// @brief Exchange the bits named by two proxies, including temporaries /// @brief Exchange the bits named by two proxies, including temporaries
/// returned from `operator[]` and `operator*`. /// returned from `operator[]` and `operator*`.
/// @param a the `a`
/// @param b the `b`
HEDLEY_NO_THROW HEDLEY_NO_THROW
friend constexpr void swap(bit_reference a, bit_reference b) noexcept friend constexpr void swap(bit_reference a, bit_reference b) noexcept
{ {
@ -897,6 +906,8 @@ class bit_array_base
static constexpr word_type sentinel = ~word_type(0); static constexpr word_type sentinel = ~word_type(0);
/// @brief Low `n` bits set. `n == 0` yields 0. `n >= bits_per_word` yields all ones. /// @brief Low `n` bits set. `n == 0` yields 0. `n >= bits_per_word` yields all ones.
/// @param n the `n`
/// @return Low `n` bits set
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr word_type low_bits_mask(size_type n) noexcept static constexpr word_type low_bits_mask(size_type n) noexcept
{ {
@ -913,6 +924,8 @@ class bit_array_base
} }
/// @brief Bits strictly below the single set bit in `mask`. /// @brief Bits strictly below the single set bit in `mask`.
/// @param mask the bit mask
/// @return Bits strictly below the single set bit in `mask`
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr word_type bits_below(word_type mask) noexcept static constexpr word_type bits_below(word_type mask) noexcept
{ {
@ -920,6 +933,8 @@ class bit_array_base
} }
/// @brief Bits at and above the single set bit in `mask`. /// @brief Bits at and above the single set bit in `mask`.
/// @param mask the bit mask
/// @return Bits at and above the single set bit in `mask`
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr word_type bits_at_and_above(word_type mask) noexcept static constexpr word_type bits_at_and_above(word_type mask) noexcept
{ {
@ -929,6 +944,11 @@ class bit_array_base
/// @brief Invoke `fn(masked_bits, relevant_mask)` for each limb touched by /// @brief Invoke `fn(masked_bits, relevant_mask)` for each limb touched by
/// `[first, last)`. Does not dereference a one-past-the-end word. /// `[first, last)`. Does not dereference a one-past-the-end word.
/// `fn` returns false to stop early. /// `fn` returns false to stop early.
/// @tparam Iterator iterator type
/// @tparam Fn fn
/// @param first the first element of the range
/// @param last the past-the-end element of the range
/// @param fn the `fn`
template <typename Iterator, typename Fn> template <typename Iterator, typename Fn>
void for_each_span(Iterator first, Iterator last, Fn fn) const void for_each_span(Iterator first, Iterator last, Fn fn) const
{ {
@ -975,6 +995,8 @@ class bit_array_base
/// @brief a base class provided to simplify the definition of /// @brief a base class provided to simplify the definition of
/// `bit_iterator` and `const_bit_iterator` /// `bit_iterator` and `const_bit_iterator`
/// @tparam ConcreteBitArrayT concrete bit array type
/// @tparam WordT word used to pack bits
template <typename ConcreteBitArrayT, template <typename ConcreteBitArrayT,
typename WordT> typename WordT>
class bit_iterator_base class bit_iterator_base
@ -1461,6 +1483,7 @@ class alignas(utils::max_align_v) static_bit_array final
} }
/// @brief constructs a `static_bit_array` from the low bits of `val` /// @brief constructs a `static_bit_array` from the low bits of `val`
/// @param val the `val`
inline constexpr explicit static_bit_array(std::size_t val) inline constexpr explicit static_bit_array(std::size_t val)
: data_{} : data_{}
{ {
@ -1503,7 +1526,7 @@ class alignas(utils::max_align_v) static_bit_array final
} }
/// @brief returns the number of bits /// @brief returns the number of bits
/// @returns number of bits that the `static_bit_array` holds /// @return number of bits that the `static_bit_array` holds
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -1534,6 +1557,8 @@ class dynamic_bit_array
using unique_ptr = typename allocator::unique_ptr; using unique_ptr = typename allocator::unique_ptr;
public: public:
/// @brief constructs a zeroed `dynamic_bit_array` that holds `nbits` bits /// @brief constructs a zeroed `dynamic_bit_array` that holds `nbits` bits
/// @param nbits the width in bits
/// @param alloc the `alloc`
/// @throws std::bad_alloc if allocating storage fails /// @throws std::bad_alloc if allocating storage fails
inline explicit dynamic_bit_array(std::size_t nbits, inline explicit dynamic_bit_array(std::size_t nbits,
allocator alloc = allocator{}) allocator alloc = allocator{})
@ -1605,6 +1630,7 @@ class dynamic_bit_array
} }
/// @brief direct access to the underlying data array /// @brief direct access to the underlying data array
/// @param i the `i`
/// @return a pointer to the start of the data array /// @return a pointer to the start of the data array
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -1625,6 +1651,7 @@ class dynamic_bit_array
} }
/// @brief direct access to the underlying data array /// @brief direct access to the underlying data array
/// @param i the `i`
/// @return a pointer to the start of the data array /// @return a pointer to the start of the data array
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -1643,7 +1670,7 @@ class dynamic_bit_array
} }
/// @brief returns the number of bits /// @brief returns the number of bits
/// @returns number of bits that the `dynamic_bit_array` holds /// @return number of bits that the `dynamic_bit_array` holds
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -1654,7 +1681,7 @@ class dynamic_bit_array
} }
private: private:
/// Store zeros through `volatile` so the wipe is not deleted as a dead store. /// @brief Store zeros through `volatile` so the wipe is not deleted as a dead store.
HEDLEY_NO_THROW HEDLEY_NO_THROW
void wipe() noexcept void wipe() noexcept
{ {
@ -1671,7 +1698,10 @@ class dynamic_bit_array
unique_ptr data_; unique_ptr data_;
}; };
/// @brief /// @brief Exchanges the bits named by two `dynamic_bit_array` proxies.
/// @tparam WordT word used to pack bits
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
template <typename WordT> template <typename WordT>
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline constexpr void swap(typename dynamic_bit_array<WordT>::reference lhs, inline constexpr void swap(typename dynamic_bit_array<WordT>::reference lhs,
@ -1682,6 +1712,11 @@ inline constexpr void swap(typename dynamic_bit_array<WordT>::reference lhs,
rhs = tmp; rhs = tmp;
} }
/// @brief Exchanges the bits named by two `static_bit_array` proxies.
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
template <std::size_t Nbits, template <std::size_t Nbits,
typename WordT> typename WordT>
HEDLEY_NO_THROW HEDLEY_NO_THROW

View file

@ -68,6 +68,7 @@ namespace dpf
/// `dpf::bit_array_base` and is parametrized on `Nbits`, which is /// `dpf::bit_array_base` and is parametrized on `Nbits`, which is
/// the length of the bitstring. /// the length of the bitstring.
/// @tparam Nbits the bitlength of the string /// @tparam Nbits the bitlength of the string
/// @tparam WordT word used to pack bits
template <std::size_t Nbits, template <std::size_t Nbits,
typename WordT = utils::integral_type_from_bitlength_t<Nbits, 8, 64>> typename WordT = utils::integral_type_from_bitlength_t<Nbits, 8, 64>>
class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT> class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
@ -81,6 +82,7 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
using const_pointer = typename base::const_pointer; using const_pointer = typename base::const_pointer;
using size_type = typename base::size_type; using size_type = typename base::size_type;
static constexpr auto bits_per_word = base::bits_per_word; static constexpr auto bits_per_word = base::bits_per_word;
static constexpr bool dpf_bitstring = true;
private: private:
/// @brief the number of `word_type`s are being used to represent the /// @brief the number of `word_type`s are being used to represent the
/// `num_bits_` bits /// `num_bits_` bits
@ -135,11 +137,15 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
/// and length `len` can be provided, as well as characters /// and length `len` can be provided, as well as characters
/// denoting alternate values for set (`one`) and unset (`zero`) /// denoting alternate values for set (`one`) and unset (`zero`)
/// bits. /// bits.
/// @tparam CharT character type
/// @tparam Traits character traits
/// @tparam Alloc allocator type
/// @param str `string` used to initialize the `dpf::bitstring` /// @param str `string` used to initialize the `dpf::bitstring`
/// @param pos a starting offset into `str` /// @param pos a starting offset into `str`
/// @param len number of characters to use from `str` /// @param len number of characters to use from `str`
/// @param zero character used to represent `0` (default: `CharT('0')`) /// @param zero character used to represent `0` (default: `CharT('0')`)
/// @param one character used to represent `1` (default: `CharT('1')`) /// @param one character used to represent `1` (default: `CharT('1')`)
/// @throws std::out_of_range
template <typename CharT, template <typename CharT,
typename Traits, typename Traits,
typename Alloc> typename Alloc>
@ -166,10 +172,12 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
/// `CharT *` `str`. An optional starting position `pos` and length /// `CharT *` `str`. An optional starting position `pos` and length
/// `len` can be provided, as well as characters denoting alternate /// `len` can be provided, as well as characters denoting alternate
/// values for set (`one`) and unset (`zero`) bits. /// values for set (`one`) and unset (`zero`) bits.
/// @tparam CharT character type
/// @param str string used to initialize the `dpf::bitstring` /// @param str string used to initialize the `dpf::bitstring`
/// @param len number of characters to use from `str` /// @param len number of characters to use from `str`
/// @param zero character used to represent `false`/`0` (default: ``CharT('0')``) /// @param zero character used to represent `false`/`0` (default: ``CharT('0')``)
/// @param one character used to represent `true`/`1` (default: ``CharT('1')``) /// @param one character used to represent `true`/`1` (default: ``CharT('1')``)
/// @throws std::invalid_argument if `null string`
template <typename CharT> template <typename CharT>
explicit bitstring(const CharT * str, explicit bitstring(const CharT * str,
typename std::basic_string<CharT>::size_type len typename std::basic_string<CharT>::size_type len
@ -293,6 +301,7 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
/// @brief shifts the bit mask to the right by the given number of /// @brief shifts the bit mask to the right by the given number of
/// bits /// bits
/// @param shift_by number of bits to shift the mask to the right /// @param shift_by number of bits to shift the mask to the right
/// @param mask the bit mask
/// @return a reference to the modified `dpf::bitstring::bit_mask` /// @return a reference to the modified `dpf::bitstring::bit_mask`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -306,6 +315,7 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
/// @brief shifts the bit mask to the left by the given number of /// @brief shifts the bit mask to the left by the given number of
/// bits /// bits
/// @param shift_by number of bits to shift the mask to the right /// @param shift_by number of bits to shift the mask to the right
/// @param mask the bit mask
/// @return a reference to the modified `dpf::bitstring::bit_mask` /// @return a reference to the modified `dpf::bitstring::bit_mask`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -342,6 +352,8 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
} }
/// @brief Inequality of the defined bits. /// @brief Inequality of the defined bits.
/// @param rhs the right-hand operand
/// @return Inequality of the defined bits
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr bool operator!=(const bitstring & rhs) const constexpr bool operator!=(const bitstring & rhs) const
{ {
@ -349,6 +361,8 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
} }
/// @brief Less than, most-significant bit first. /// @brief Less than, most-significant bit first.
/// @param rhs the right-hand operand
/// @return Less than, most-significant bit first
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr bool operator<(const bitstring & rhs) const constexpr bool operator<(const bitstring & rhs) const
{ {
@ -356,6 +370,8 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
} }
/// @brief Less than or equal, most-significant bit first. /// @brief Less than or equal, most-significant bit first.
/// @param rhs the right-hand operand
/// @return Less than or equal, most-significant bit first
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr bool operator<=(const bitstring & rhs) const constexpr bool operator<=(const bitstring & rhs) const
{ {
@ -363,6 +379,8 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
} }
/// @brief Greater than, most-significant bit first. /// @brief Greater than, most-significant bit first.
/// @param rhs the right-hand operand
/// @return Greater than, most-significant bit first
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr bool operator>(const bitstring & rhs) const constexpr bool operator>(const bitstring & rhs) const
{ {
@ -501,7 +519,7 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
} }
/// @brief returns the number of bits /// @brief returns the number of bits
/// @returns number of bits that the `bitstring` holds /// @return number of bits that the `bitstring` holds
/// @complexity `O(1)` /// @complexity `O(1)`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -517,6 +535,7 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
std::array<word_type, data_length_> data_{}; std::array<word_type, data_length_> data_{};
/// @brief Mask of the bits that belong to this string in the high word. /// @brief Mask of the bits that belong to this string in the high word.
/// @return the returned `word_type`
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr word_type defined_high_mask() noexcept static constexpr word_type defined_high_mask() noexcept
{ {
@ -537,6 +556,8 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
/// @brief Most-significant word first. Unused high bits are ignored. /// @brief Most-significant word first. Unused high bits are ignored.
/// Every limb is visited, so the time does not depend on where /// Every limb is visited, so the time does not depend on where
/// the strings differ. /// the strings differ.
/// @param rhs the right-hand operand
/// @return Most-significant word first
constexpr int compare(const bitstring & rhs) const constexpr int compare(const bitstring & rhs) const
{ {
if constexpr (data_length_ == 0) return 0; if constexpr (data_length_ == 0) return 0;
@ -558,6 +579,12 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
} }
/// @brief Last character is bit 0. `len` must be at most `Nbits`. /// @brief Last character is bit 0. `len` must be at most `Nbits`.
/// @tparam CharT character type
/// @param str the source string
/// @param len the number of bytes
/// @param zero the character used for 0
/// @param one the character used for 1
/// @throws std::out_of_range if `string longer than Nbits`
template <typename CharT> template <typename CharT>
void assign_msb_string(const CharT * str, std::size_t len, CharT zero, CharT one) void assign_msb_string(const CharT * str, std::size_t len, CharT zero, CharT one)
{ {
@ -615,6 +642,8 @@ namespace utils
{ {
/// @brief specializes `dpf::utils::bitlength_of` for `dpf::bitstring` /// @brief specializes `dpf::utils::bitlength_of` for `dpf::bitstring`
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
template <std::size_t Nbits, template <std::size_t Nbits,
typename WordT> typename WordT>
struct bitlength_of<dpf::bitstring<Nbits, WordT>> struct bitlength_of<dpf::bitstring<Nbits, WordT>>
@ -622,6 +651,8 @@ struct bitlength_of<dpf::bitstring<Nbits, WordT>>
{ }; { };
/// @brief specializes `dpf::utils::msb_of` for `dpf::bitstring` /// @brief specializes `dpf::utils::msb_of` for `dpf::bitstring`
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
template <std::size_t Nbits, template <std::size_t Nbits,
typename WordT> typename WordT>
struct msb_of<dpf::bitstring<Nbits, WordT>> struct msb_of<dpf::bitstring<Nbits, WordT>>
@ -633,6 +664,8 @@ struct msb_of<dpf::bitstring<Nbits, WordT>>
/// @brief specializes `dpf::utils::countl_zero_symmetric_difference` for /// @brief specializes `dpf::utils::countl_zero_symmetric_difference` for
/// `dpf::bitstring` /// `dpf::bitstring`
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
template <std::size_t Nbits, template <std::size_t Nbits,
typename WordT> typename WordT>
struct countl_zero_symmetric_difference<dpf::bitstring<Nbits, WordT>> struct countl_zero_symmetric_difference<dpf::bitstring<Nbits, WordT>>
@ -909,6 +942,9 @@ namespace bitstrings
/// @brief Build a bitstring from characters. The first character is the /// @brief Build a bitstring from characters. The first character is the
/// most significant bit of the digit string (same order as `0b...`). /// most significant bit of the digit string (same order as `0b...`).
/// Digits shorter than `Bitstring::size()` occupy the low bits. /// Digits shorter than `Bitstring::size()` occupy the low bits.
/// @tparam Bitstring bitstring
/// @tparam bits bits
/// @return the returned `Bitstring`
template <typename Bitstring, char... bits> template <typename Bitstring, char... bits>
constexpr Bitstring bitstring_literal() constexpr Bitstring bitstring_literal()
{ {
@ -925,6 +961,8 @@ constexpr Bitstring bitstring_literal()
/// @details The leftmost character is the most significant bit, matching /// @details The leftmost character is the most significant bit, matching
/// `0b` integer literals. `10101001_bitstring` equals /// `0b` integer literals. `10101001_bitstring` equals
/// `dpf::bitstring<8>(0b10101001)`. /// `dpf::bitstring<8>(0b10101001)`.
/// @tparam bits bits
/// @return user-defined numeric literal for creating `dpf::bitstring` objects
template <char ...bits> template <char ...bits>
constexpr static auto operator "" _bitstring() constexpr static auto operator "" _bitstring()
{ {
@ -937,7 +975,9 @@ constexpr static auto operator "" _bitstring_u8()
return bitstring_literal<dpf::bitstring<sizeof...(bits), psnip_uint8_t>, bits...>(); return bitstring_literal<dpf::bitstring<sizeof...(bits), psnip_uint8_t>, bits...>();
} }
/// Alias used by the test suite: word type `uint8_t`, not length 8. /// @brief Alias used by the test suite: word type `uint8_t`, not length 8.
/// @tparam bits bits
/// @return Alias used by the test suite: word type `uint8_t`, not length 8
template <char ...bits> template <char ...bits>
constexpr static auto operator "" _bitstring_8() constexpr static auto operator "" _bitstring_8()
{ {
@ -1123,6 +1163,8 @@ namespace std
/// @{ /// @{
/// @details specializes `std::numeric_limits` for `dpf::bitstring<Nbits, WordT>` /// @details specializes `std::numeric_limits` for `dpf::bitstring<Nbits, WordT>`
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
template<std::size_t Nbits, template<std::size_t Nbits,
typename WordT> typename WordT>
class numeric_limits<dpf::bitstring<Nbits, WordT>> class numeric_limits<dpf::bitstring<Nbits, WordT>>
@ -1174,20 +1216,26 @@ class numeric_limits<dpf::bitstring<Nbits, WordT>>
}; };
/// @details specializes `std::numeric_limits` for `dpf::bitstring<Nbits, WordT> const` /// @details specializes `std::numeric_limits` for `dpf::bitstring<Nbits, WordT> const`
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
template<std::size_t Nbits, template<std::size_t Nbits,
typename WordT> typename WordT>
class numeric_limits<dpf::bitstring<Nbits, WordT> const> class numeric_limits<dpf::bitstring<Nbits, WordT> const>
: public numeric_limits<dpf::bitstring<Nbits, WordT>> {}; : public numeric_limits<dpf::bitstring<Nbits, WordT>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `dpf::bitstring<Nbits, WordT> volatile` /// @brief `dpf::bitstring<Nbits, WordT> volatile`
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
template<std::size_t Nbits, template<std::size_t Nbits,
typename WordT> typename WordT>
class numeric_limits<dpf::bitstring<Nbits, WordT> volatile> class numeric_limits<dpf::bitstring<Nbits, WordT> volatile>
: public numeric_limits<dpf::bitstring<Nbits, WordT>> {}; : public numeric_limits<dpf::bitstring<Nbits, WordT>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `dpf::bitstring<Nbits, WordT> const volatile` /// @brief `dpf::bitstring<Nbits, WordT> const volatile`
/// @tparam Nbits width in bits
/// @tparam WordT word used to pack bits
template<std::size_t Nbits, template<std::size_t Nbits,
typename WordT> typename WordT>
class numeric_limits<dpf::bitstring<Nbits, WordT> const volatile> class numeric_limits<dpf::bitstring<Nbits, WordT> const volatile>

View file

@ -22,6 +22,7 @@
#include "dpf/dcf.hpp" #include "dpf/dcf.hpp"
#include "dpf/path_memoizer.hpp" #include "dpf/path_memoizer.hpp"
#include "dpf/tree_traits.hpp"
#include "dpf/twiddle.hpp" #include "dpf/twiddle.hpp"
#include "dpf/utils.hpp" #include "dpf/utils.hpp"
@ -79,11 +80,11 @@ struct schedule
}; };
template <typename PRG, typename Node> template <typename PRG, typename Node>
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
uint64_t rho_of(const Node & node, uint64_t mask) noexcept uint64_t rho_of(const Node & node, uint64_t mask) noexcept
{ {
auto kids = PRG::eval01(dpf::unset_lo_2bits(node)); auto kids = dpf::tree_traits<PRG>::expand_value(node);
return dcf_impl::convert_node(kids[0], mask); return dcf_impl::convert_node(kids[0], mask);
} }
@ -107,7 +108,10 @@ constexpr uint64_t mul_sgn(int sgn, uint64_t v, uint64_t mask) noexcept
return 0; return 0;
} }
/// Group element `sgn` (`+1`, `-1`, or `0`) used as an `assign_cmp` coefficient. /// @brief Group element `sgn` (`+1`, `-1`, or `0`) used as an `assign_cmp` coefficient.
/// @param sgn the `sgn`
/// @param mask the bit mask
/// @return Group element `sgn` (`+1`, `-1`, or `0`) used as an `assign_cmp` coefficient
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -136,6 +140,7 @@ uint64_t checkpoint_word(const Node & n0, const Node & n1, uint64_t beta,
} }
template <typename Node> template <typename Node>
HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t checkpoint_coeff(const Node & n0, const Node & n1, uint64_t checkpoint_coeff(const Node & n0, const Node & n1,
uint64_t mask) noexcept uint64_t mask) noexcept
@ -160,7 +165,7 @@ void suffix_masks(const Node & seed, std::size_t q, uint64_t mask,
std::size_t m = 0; std::size_t m = 0;
for (std::size_t i = 0; i < n; ++i) for (std::size_t i = 0; i < n; ++i)
{ {
auto kids = PRG::eval01(dpf::unset_lo_2bits(cur[i])); auto kids = dpf::tree_traits<PRG>::expand_value(cur[i]);
nxt[m++] = kids[0]; nxt[m++] = kids[0];
nxt[m++] = kids[1]; nxt[m++] = kids[1];
} }
@ -219,8 +224,11 @@ uint64_t add_frontier(uint64_t acc, const typename KeyT::interior_node & seed,
uint64_t mask, int party) uint64_t mask, int party)
{ {
using node = typename KeyT::interior_node; using node = typename KeyT::interior_node;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::vector<node> cur; std::vector<node> cur;
std::vector<node> nxt; std::vector<node> nxt;
HEDLEY_PRAGMA(GCC diagnostic pop)
cur.push_back(seed); cur.push_back(seed);
for (std::size_t lvl = from_depth; lvl < to_depth; ++lvl) for (std::size_t lvl = from_depth; lvl < to_depth; ++lvl)
{ {
@ -228,9 +236,10 @@ uint64_t add_frontier(uint64_t acc, const typename KeyT::interior_node & seed,
nxt.reserve(cur.size() * 2); nxt.reserve(cur.size() * 2);
const node cw0 = dpf.correction_word(lvl, false); const node cw0 = dpf.correction_word(lvl, false);
const node cw1 = dpf.correction_word(lvl, true); const node cw1 = dpf.correction_word(lvl, true);
const bool is_last = KeyT::tree::is_last_level(lvl, KeyT::depth);
for (const node & fs : cur) for (const node & fs : cur)
{ {
auto kids = KeyT::traverse_interior01(fs, cw0, cw1); auto kids = KeyT::traverse_interior01(fs, cw0, cw1, is_last);
nxt.push_back(kids[0]); nxt.push_back(kids[0]);
nxt.push_back(kids[1]); nxt.push_back(kids[1]);
} }
@ -334,7 +343,8 @@ uint64_t eval_share(const KeyT & dpf, InputT tx, PathMemoizer & path)
const bool xi = !!(bit_mask & tx); const bool xi = !!(bit_mask & tx);
const node & parent = path[level]; const node & parent = path[level];
const node right = KeyT::traverse_interior(parent, const node right = KeyT::traverse_interior(parent,
dpf.correction_word(level, true), true); dpf.correction_word(level, true), true,
KeyT::tree::is_last_level(level, KeyT::depth));
if (!xi) if (!xi)
{ {
pend[npend].seed = right; pend[npend].seed = right;
@ -358,6 +368,7 @@ uint64_t eval_share(const KeyT & dpf, InputT tx, PathMemoizer & path)
} }
template <typename KeyT, typename Integral, typename Memo> template <typename KeyT, typename Integral, typename Memo>
HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
bool memo_has(const Memo & memo, Integral prefix, std::size_t depth, bool memo_has(const Memo & memo, Integral prefix, std::size_t depth,
Integral from_lane, Integral to_excl) noexcept Integral from_lane, Integral to_excl) noexcept
@ -422,7 +433,8 @@ uint64_t eval_share_memo(const KeyT & dpf, Integral lane,
if (!xi) if (!xi)
{ {
const node right = KeyT::traverse_interior(parent, const node right = KeyT::traverse_interior(parent,
dpf.correction_word(level, true), true); dpf.correction_word(level, true), true,
KeyT::tree::is_last_level(level, KeyT::depth));
pend[npend].seed = right; pend[npend].seed = right;
pend[npend].prefix = sib; pend[npend].prefix = sib;
pend[npend].depth = level + 1; pend[npend].depth = level + 1;

View file

@ -65,7 +65,12 @@ struct lane_codec
return out; return out;
} }
/// Element `index` is the packed byte range `[index * sizeof(T), ...)`. /// @brief Element `index` is the packed byte range `[index * sizeof(T), ...)`.
/// @param seed the PRG seed
/// @param index the index
/// @param out the output buffer
/// @param count the number of blocks
/// @throws std::invalid_argument if `prg lane index is out of range`
static void fill(block_type seed, std::uint64_t index, T * out, std::size_t count) static void fill(block_type seed, std::uint64_t index, T * out, std::size_t count)
{ {
if (count == 0) if (count == 0)
@ -85,8 +90,8 @@ struct lane_codec
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
dpf::aligned_allocator<block_type> alloc; dpf::aligned_allocator<block_type> alloc;
auto blocks = alloc.allocate_unique_ptr(static_cast<std::size_t>(nblocks));
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
auto blocks = alloc.allocate_unique_ptr(static_cast<std::size_t>(nblocks));
PRG::eval(seed, blocks.get(), static_cast<psnip_uint32_t>(nblocks), PRG::eval(seed, blocks.get(), static_cast<psnip_uint32_t>(nblocks),
static_cast<psnip_uint32_t>(start)); static_cast<psnip_uint32_t>(start));
auto * bytes = reinterpret_cast<const unsigned char *>(blocks.get()); auto * bytes = reinterpret_cast<const unsigned char *>(blocks.get());
@ -167,13 +172,15 @@ typename PRG::block_type sample_master_seed()
return dpf::uniform_sample<typename PRG::block_type>(); return dpf::uniform_sample<typename PRG::block_type>();
} }
/// Forward cursor over one PRG stream per value type. /// @brief Forward cursor over one PRG stream per value type.
/// ///
/// `get<I>()` and `fill<I>()` consume the cursor. `at<I>(index)` reads an /// `get<I>()` and `fill<I>()` consume the cursor. `at<I>(index)` reads an
/// absolute index and leaves the cursor where it is. `sampled<I>()` reports /// absolute index and leaves the cursor where it is. `sampled<I>()` reports
/// how far `get` and `fill` have advanced. `per_stream_buffer_elems` is at /// how far `get` and `fill` have advanced. `per_stream_buffer_elems` is at
/// least 1. /// least 1.
/// @snippet evaluation/buffered_prg.cpp buffered-prg /// @snippet evaluation/buffered_prg.cpp buffered-prg
/// @tparam PRG pseudorandom generator
/// @tparam Ts ts
template <typename PRG, typename... Ts> template <typename PRG, typename... Ts>
class buffered_prg class buffered_prg
{ {
@ -248,12 +255,14 @@ private:
template <typename... Ts> template <typename... Ts>
using aes_buffered_prg = buffered_prg<dpf::prg::aes128, Ts...>; using aes_buffered_prg = buffered_prg<dpf::prg::aes128, Ts...>;
/// Seekable value and mask streams for a runtime set of roles. /// @brief Seekable value and mask streams for a runtime set of roles.
/// ///
/// `value_at(role, index)` and `mask_at(role, index)` are independent of /// `value_at(role, index)` and `mask_at(role, index)` are independent of
/// call order. A repeated index returns the same element. `window` is at /// call order. A repeated index returns the same element. `window` is at
/// least 1. /// least 1.
/// @snippet evaluation/buffered_prg.cpp lane-table /// @snippet evaluation/buffered_prg.cpp lane-table
/// @tparam T value type
/// @tparam PRG pseudorandom generator
template <typename T, typename PRG = dpf::prg::aes128> template <typename T, typename PRG = dpf::prg::aes128>
class lane_table class lane_table
{ {

548
include/dpf/cmp_group.hpp Normal file
View file

@ -0,0 +1,548 @@
/// @file dpf/cmp_group.hpp
/// @brief Comparison-payload group for types that do not fit in a masked `uint64_t`.
/// @details Payloads of at most 64 bits that already convert to `uint64_t`
/// stay on that path. Everything else — wider integers, `modint`,
/// `fixedpoint`, `xor_wrapper`, `bitstring`, and `dpf::vec` — is a
/// little-endian limb vector. Lanes of a `vec` add (or XOR) apart,
/// with no carry from one lane into the next. A PRG stretch fills a
/// group element from one GGM node, so the element is uniform even
/// when it is wider than 64 bits.
/// @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_DPF_CMP_GROUP_HPP__
#define LIBDPF_INCLUDE_DPF_CMP_GROUP_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <type_traits>
#include "hedley/hedley.h"
#include "dpf/twiddle.hpp"
#include "dpf/utils.hpp"
#include "dpf/wildcard.hpp"
namespace dpf
{
namespace detail
{
struct group_elem
{
static constexpr std::size_t cap = 4;
uint64_t limb[cap]{};
std::uint16_t lane_bits = 64;
std::uint16_t lanes = 1;
bool xor_group = false;
};
template <typename T, typename = void>
struct is_modint_tag : std::false_type {};
template <typename T>
struct is_modint_tag<T, std::void_t<decltype(T::dpf_modint)>>
: std::bool_constant<T::dpf_modint> {};
template <typename T, typename = void>
struct is_bitstring_tag : std::false_type {};
template <typename T>
struct is_bitstring_tag<T, std::void_t<decltype(T::dpf_bitstring)>>
: std::bool_constant<T::dpf_bitstring> {};
template <typename T, typename = void>
struct is_vec_tag : std::false_type {};
template <typename T>
struct is_vec_tag<T, std::void_t<decltype(T::dpf_vec)>>
: std::bool_constant<T::dpf_vec> {};
template <typename T, typename = void>
struct has_integral_representation : std::false_type {};
template <typename T>
struct has_integral_representation<T, std::void_t<decltype(
std::declval<const T &>().integral_representation())>>
: std::true_type {};
template <typename T, bool IsVec>
struct cmp_lane_of { using type = T; };
template <typename T>
struct cmp_lane_of<T, true> { using type = typename T::lane_type; };
template <typename Beta>
struct cmp_group_info
{
using type = concrete_type_t<std::decay_t<Beta>>;
static constexpr bool is_vec = is_vec_tag<type>::value;
using lane = typename cmp_lane_of<type, is_vec>::type;
static constexpr bool lane_xor =
utils::is_xor_wrapper_v<lane> || is_bitstring_tag<lane>::value;
static constexpr std::size_t lanes = []() constexpr {
if constexpr (is_vec)
return type::lane_count;
else
return std::size_t{1};
}();
static constexpr std::size_t lane_bits = utils::bitlength_of_v<lane>;
static constexpr std::size_t total_bits = lanes * lane_bits;
/// @brief `uint64_t` ring, including a fixed-point value whose raw word fits.
static constexpr bool narrow_ring =
!is_vec && !lane_xor && lane_bits <= 64;
static constexpr bool custom = !narrow_ring;
static_assert(!custom || total_bits <= 256,
"comparison payload exceeds 256 bits");
static_assert(!custom || total_bits > 0,
"comparison payload has no bits");
};
template <typename Concrete>
HEDLEY_ALWAYS_INLINE
group_elem group_layout()
{
using info = cmp_group_info<Concrete>;
group_elem g;
g.lane_bits = static_cast<std::uint16_t>(info::lane_bits);
g.lanes = static_cast<std::uint16_t>(info::lanes);
g.xor_group = info::lane_xor;
return g;
}
HEDLEY_ALWAYS_INLINE
group_elem group_zero(const group_elem & layout)
{
group_elem g;
g.lane_bits = layout.lane_bits;
g.lanes = layout.lanes;
g.xor_group = layout.xor_group;
return g;
}
inline bool bit_at(const uint64_t limb[4], std::size_t bit) noexcept
{
return ((limb[bit / 64u] >> (bit % 64u)) & 1ull) != 0;
}
inline void set_bit(uint64_t limb[4], std::size_t bit, bool on) noexcept
{
const std::size_t i = bit / 64u;
const uint64_t m = 1ull << (bit % 64u);
if (on)
limb[i] |= m;
else
limb[i] &= ~m;
}
inline void mask_limbs(uint64_t limb[4], std::size_t bits) noexcept
{
if (bits >= 256)
return;
for (std::size_t b = bits; b < 256; ++b)
set_bit(limb, b, false);
}
inline void read_lane(const group_elem & g, std::size_t lane, uint64_t out[4]) noexcept
{
std::memset(out, 0, 4 * sizeof(uint64_t));
const std::size_t base = lane * g.lane_bits;
for (std::size_t b = 0; b < g.lane_bits; ++b)
set_bit(out, b, bit_at(g.limb, base + b));
}
inline void write_lane(group_elem & g, std::size_t lane, const uint64_t in[4]) noexcept
{
const std::size_t base = lane * g.lane_bits;
for (std::size_t b = 0; b < g.lane_bits; ++b)
set_bit(g.limb, base + b, bit_at(in, b));
}
inline void add_lane(uint64_t a[4], const uint64_t b[4], std::size_t bits) noexcept
{
unsigned carry = 0;
const std::size_t n = (bits + 63u) / 64u;
for (std::size_t i = 0; i < n; ++i)
{
const unsigned __int128 sum =
static_cast<unsigned __int128>(a[i]) + b[i] + carry;
a[i] = static_cast<uint64_t>(sum);
carry = static_cast<unsigned>(sum >> 64);
}
mask_limbs(a, bits);
}
inline void neg_lane(uint64_t a[4], std::size_t bits) noexcept
{
uint64_t one[4] = {1, 0, 0, 0};
for (std::size_t i = 0; i < 4; ++i)
a[i] = ~a[i];
mask_limbs(a, bits);
add_lane(a, one, bits);
}
/// @brief Bit-serial product. Used when a wildcard coefficient scales δ, and when
/// interval containment multiplies δ by a small public integer.
/// @param a the `a`
/// @param b the `b`
/// @param bits the packed bits
/// @param out the output buffer
inline void mul_lane(const uint64_t a[4], const uint64_t b[4], std::size_t bits,
uint64_t out[4]) noexcept
{
std::memset(out, 0, 4 * sizeof(uint64_t));
for (std::size_t bit = 0; bit < bits; ++bit)
{
if (!bit_at(b, bit))
continue;
uint64_t shifted[4]{};
for (std::size_t s = 0; s < bits; ++s)
{
if (s + bit < bits && bit_at(a, s))
set_bit(shifted, s + bit, true);
}
add_lane(out, shifted, bits);
}
}
inline group_elem group_apply(const group_elem & a, const group_elem & b,
void (*lane_op)(uint64_t *, const uint64_t *, std::size_t))
{
group_elem out = group_zero(a);
for (std::size_t i = 0; i < a.lanes; ++i)
{
uint64_t la[4]{}, lb[4]{};
read_lane(a, i, la);
read_lane(b, i, lb);
if (a.xor_group)
{
for (std::size_t k = 0; k < 4; ++k)
la[k] ^= lb[k];
mask_limbs(la, a.lane_bits);
}
else
{
lane_op(la, lb, a.lane_bits);
}
write_lane(out, i, la);
}
return out;
}
inline group_elem group_add(const group_elem & a, const group_elem & b)
{
return group_apply(a, b, add_lane);
}
inline group_elem group_neg(const group_elem & a)
{
group_elem out = group_zero(a);
if (a.xor_group)
return a;
for (std::size_t i = 0; i < a.lanes; ++i)
{
uint64_t la[4]{};
read_lane(a, i, la);
neg_lane(la, a.lane_bits);
write_lane(out, i, la);
}
return out;
}
inline group_elem group_sub(const group_elem & a, const group_elem & b)
{
if (a.xor_group)
return group_add(a, b);
return group_add(a, group_neg(b));
}
inline group_elem group_mul(const group_elem & a, const group_elem & b)
{
group_elem out = group_zero(a);
for (std::size_t i = 0; i < a.lanes; ++i)
{
uint64_t la[4]{}, lb[4]{}, lc[4]{};
read_lane(a, i, la);
read_lane(b, i, lb);
if (a.xor_group)
{
for (std::size_t k = 0; k < 4; ++k)
lc[k] = la[k] & lb[k];
mask_limbs(lc, a.lane_bits);
}
else
{
mul_lane(la, lb, a.lane_bits, lc);
}
write_lane(out, i, lc);
}
return out;
}
inline group_elem group_sgn(bool t, const group_elem & a)
{
return t ? group_neg(a) : a;
}
/// @brief Multiplicative identity: `1` in each additive lane, all-ones in an XOR lane.
/// @param layout the `layout`
/// @return Multiplicative identity: `1` in each additive lane, all-ones in an XOR lane
inline group_elem group_one(const group_elem & layout)
{
group_elem g = group_zero(layout);
for (std::size_t i = 0; i < g.lanes; ++i)
{
uint64_t lane[4]{};
if (g.xor_group)
{
for (std::size_t b = 0; b < g.lane_bits; ++b)
set_bit(lane, b, true);
}
else
{
lane[0] = 1;
}
write_lane(g, i, lane);
}
return g;
}
/// @brief Integer `s` in every lane. Negative `s` is the group negation of `|s|`.
/// @param s the `s`
/// @param layout the `layout`
/// @return Integer `s` in every lane
inline group_elem group_scalar(int s, const group_elem & layout)
{
if (layout.xor_group)
{
if ((s & 1) == 0)
return group_zero(layout);
return group_one(layout);
}
const bool neg = s < 0;
const auto mag = static_cast<unsigned>(neg ? -s : s);
group_elem g = group_zero(layout);
for (std::size_t i = 0; i < g.lanes; ++i)
{
uint64_t lane[4] = {mag, 0, 0, 0};
mask_limbs(lane, g.lane_bits);
if (neg)
neg_lane(lane, g.lane_bits);
write_lane(g, i, lane);
}
return g;
}
inline group_elem group_from_bytes(const unsigned char * bytes, std::size_t nbytes,
const group_elem & layout)
{
group_elem g = group_zero(layout);
const std::size_t need = (static_cast<std::size_t>(layout.lanes) * layout.lane_bits
+ 7u) / 8u;
if (nbytes < need)
throw std::invalid_argument("comparison group stretch was short");
std::size_t bit = 0;
for (std::size_t lane = 0; lane < layout.lanes; ++lane)
{
uint64_t raw[4]{};
for (std::size_t b = 0; b < layout.lane_bits; ++b, ++bit)
{
const unsigned char byte = bytes[bit / 8u];
const bool on = ((byte >> (bit % 8u)) & 1u) != 0;
set_bit(raw, b, on);
}
write_lane(g, lane, raw);
}
return g;
}
template <typename PRG, typename Node>
group_elem group_from_node(Node node, const group_elem & layout)
{
auto seed = dpf::unset_lo_2bits(node);
auto kids = PRG::eval01(seed);
unsigned char bytes[64]{};
constexpr std::size_t nb = sizeof(kids[0]);
static_assert(nb <= 32, "comparison stretch expects a 128- or 256-bit block");
std::memcpy(bytes, &kids[0], nb);
std::memcpy(bytes + nb, &kids[1], nb);
return group_from_bytes(bytes, nb * 2, layout);
}
template <typename Word>
Word group_to_word(const group_elem & g)
{
Word w{};
static_assert(sizeof(Word) <= sizeof(g.limb),
"comparison word is wider than 256 bits");
std::memcpy(&w, g.limb, sizeof(Word));
return w;
}
template <typename Word>
group_elem group_from_word(const Word & w, const group_elem & layout)
{
group_elem g = group_zero(layout);
std::memcpy(g.limb, &w, sizeof(Word) < sizeof(g.limb) ? sizeof(Word) : sizeof(g.limb));
mask_limbs(g.limb, static_cast<std::size_t>(layout.lanes) * layout.lane_bits);
return g;
}
template <typename T>
void store_raw_integer(group_elem & g, std::size_t lane, const T & value)
{
uint64_t tmp[4]{};
if constexpr (has_integral_representation<T>::value)
{
auto raw = value.integral_representation();
std::memcpy(tmp, &raw, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
}
else if constexpr (is_modint_tag<T>::value)
{
auto raw = static_cast<typename T::integral_type>(value);
std::memcpy(tmp, &raw, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
}
else if constexpr (is_bitstring_tag<T>::value)
{
auto raw = utils::to_integral_type<T>{}(value);
std::memcpy(tmp, &raw, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
}
else if constexpr (utils::is_xor_wrapper_v<T>)
{
auto raw = value.data();
std::memcpy(tmp, &raw, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
}
else
{
static_assert(std::is_trivially_copyable_v<T>,
"comparison payload must be trivially copyable");
static_assert(sizeof(T) <= sizeof(tmp),
"comparison payload exceeds 256 bits");
std::memcpy(tmp, &value, sizeof(T));
}
mask_limbs(tmp, g.lane_bits);
write_lane(g, lane, tmp);
}
template <typename Beta>
group_elem group_from_beta(const Beta & value)
{
using C = std::decay_t<Beta>;
group_elem g = group_layout<C>();
if constexpr (is_vec_tag<C>::value)
{
for (std::size_t i = 0; i < C::lane_count; ++i)
{
auto lane = group_from_beta(value.lanes[i]);
uint64_t raw[4]{};
read_lane(lane, 0, raw);
write_lane(g, i, raw);
}
}
else
{
store_raw_integer(g, 0, value);
}
return g;
}
template <typename Beta>
Beta group_to_beta(const group_elem & g)
{
using C = std::decay_t<Beta>;
if constexpr (is_vec_tag<C>::value)
{
C out{};
for (std::size_t i = 0; i < C::lane_count; ++i)
{
group_elem lane = group_zero(g);
lane.lanes = 1;
lane.lane_bits = g.lane_bits;
lane.xor_group = g.xor_group;
uint64_t raw[4]{};
read_lane(g, i, raw);
write_lane(lane, 0, raw);
out.lanes[i] = group_to_beta<typename C::lane_type>(lane);
}
return out;
}
else
{
uint64_t tmp[4]{};
read_lane(g, 0, tmp);
if constexpr (has_integral_representation<C>::value)
{
using integral = typename C::integral_type;
integral raw{};
std::memcpy(&raw, tmp, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
return C::from_raw(raw);
}
else if constexpr (is_modint_tag<C>::value)
{
using integral = typename C::integral_type;
integral raw{};
std::memcpy(&raw, tmp, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
return C{raw};
}
else if constexpr (is_bitstring_tag<C>::value)
{
using integral = typename utils::make_from_integral_value<C>::integral_type;
integral raw{};
std::memcpy(&raw, tmp, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
return utils::make_from_integral_value<C>{}(raw);
}
else if constexpr (utils::is_xor_wrapper_v<C>)
{
using raw_type = typename C::value_type;
raw_type raw{};
std::memcpy(&raw, tmp, sizeof(raw) < sizeof(tmp) ? sizeof(raw) : sizeof(tmp));
return C{raw};
}
else
{
C raw{};
std::memcpy(&raw, tmp, sizeof(C) < sizeof(tmp) ? sizeof(C) : sizeof(tmp));
return raw;
}
}
}
template <typename PRG>
group_elem group_value_cw(const typename PRG::block_type & c0L,
const typename PRG::block_type & c0R,
const typename PRG::block_type & c1L,
const typename PRG::block_type & c1R,
uint8_t t0, uint8_t t1, int ai, group_elem & Va, const group_elem & beta)
{
(void)t0;
const group_elem v0L = group_from_node<PRG>(c0L, Va);
const group_elem v0R = group_from_node<PRG>(c0R, Va);
const group_elem v1L = group_from_node<PRG>(c1L, Va);
const group_elem v1R = group_from_node<PRG>(c1R, Va);
const group_elem & v0K = ai == 0 ? v0L : v0R;
const group_elem & v1K = ai == 0 ? v1L : v1R;
const group_elem & v0Lo = ai == 0 ? v0R : v0L;
const group_elem & v1Lo = ai == 0 ? v1R : v1L;
group_elem vcw = group_sgn(t1 != 0,
group_add(group_add(v1Lo, group_neg(v0Lo)), group_neg(Va)));
// Lose-left plants β. A planted recipe passes the plant in `beta` for both directions.
if (ai == 1)
vcw = group_add(vcw, group_sgn(t1 != 0, beta));
Va = group_add(group_add(group_add(Va, group_neg(v1K)), v0K),
group_sgn(t1 != 0, vcw));
return vcw;
}
template <typename PRG>
group_elem group_final_cw(const typename PRG::block_type & s0,
const typename PRG::block_type & s1, uint8_t t1, const group_elem & Va,
const group_elem & on_path)
{
const group_elem c0 = group_from_node<PRG>(s0, Va);
const group_elem c1 = group_from_node<PRG>(s1, Va);
return group_sgn(t1 != 0,
group_add(group_add(group_add(c1, group_neg(c0)), group_neg(Va)), on_path));
}
} // namespace detail
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_CMP_GROUP_HPP__

View file

@ -0,0 +1,110 @@
/// @file dpf/constrained_cmp.hpp
/// @brief Constrained integer comparison Π_CCMP (NDSS 2025 Alg. 1).
/// @details Given two positive integers that differ by exactly one,
/// `1{x0 < x1}` is computed with a single AND on two derived bits.
/// Local joint simulation opens the AND clearly; an MPC backend would
/// replace that open with the existing Beaver AND tape.
/// @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_CONSTRAINED_CMP_HPP__
#define LIBDPF_INCLUDE_DPF_CONSTRAINED_CMP_HPP__
#include <cstdint>
#include <type_traits>
#include "hedley/hedley.h"
namespace dpf
{
namespace detail
{
/// @brief Last two bits of `x`: high = bit 1, low = bit 0.
/// @param x the `x`
/// @return Last two bits of `x`: high = bit 1, low = bit 0
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
constexpr uint8_t ccmp_lo_bit(std::uint64_t x) noexcept
{
return static_cast<uint8_t>(x & 1u);
}
/// @brief Bit 1 of `x`.
/// @param x the integer
/// @return `(x >> 1) & 1`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
constexpr uint8_t ccmp_hi_bit(std::uint64_t x) noexcept
{
return static_cast<uint8_t>((x >> 1) & 1u);
}
/// @brief Party `b`'s local share inputs for the AND: `z0 = h`, `z1 = h ⊕ l ⊕ b`.
/// @param x the integer whose low two bits are split
/// @param party party index, `0` or `1`
/// @param z0 first AND input, `h`
/// @param z1 second AND input, `h ⊕ l ⊕ party`
/// @param l bit 0 of `x`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr void ccmp_party_terms(std::uint64_t x, uint8_t party,
uint8_t & z0, uint8_t & z1, uint8_t & l) noexcept
{
const uint8_t h = ccmp_hi_bit(x);
l = ccmp_lo_bit(x);
z0 = h;
z1 = static_cast<uint8_t>(h ^ l ^ (party & 1u));
}
/// @brief Opened result of Π_CCMP when both inputs are known (local joint sim).
/// @details Precondition: `|x0 - x1| = 1`.
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @return Opened result of Π_CCMP when both inputs are known (local joint sim)
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
constexpr uint8_t local_ccmp(std::uint64_t x0, std::uint64_t x1) noexcept
{
uint8_t z00 = 0, z01 = 0, l0 = 0;
uint8_t z10 = 0, z11 = 0, l1 = 0;
ccmp_party_terms(x0, 0, z00, z01, l0);
ccmp_party_terms(x1, 1, z10, z11, l1);
const uint8_t z0 = static_cast<uint8_t>(z00 ^ z10);
const uint8_t z1 = static_cast<uint8_t>(z01 ^ z11);
const uint8_t t = static_cast<uint8_t>(z0 & z1);
// Party shares: y0 = t0, y1 = t1 ⊕ (l1 ∧ 1). Opened y = t ⊕ l1.
return static_cast<uint8_t>(t ^ l1);
}
/// @brief Same as `local_ccmp` for any unsigned or enum-convertible integer.
/// @tparam T0 integral type of the first operand
/// @tparam T1 integral type of the second operand
/// @param x0 the first integer
/// @param x1 the second integer
/// @return Same as `local_ccmp` for any unsigned or enum-convertible integer
template <typename T0, typename T1>
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
constexpr uint8_t local_ccmp_int(T0 x0, T1 x1) noexcept
{
static_assert(std::is_integral_v<T0> && std::is_integral_v<T1>,
"local_ccmp_int: integral operands");
return local_ccmp(static_cast<std::uint64_t>(x0),
static_cast<std::uint64_t>(x1));
}
} // namespace detail
/// @brief Constrained comparison: `1{x0 < x1}` when `|x0 − x1| = 1`.
using detail::local_ccmp;
using detail::local_ccmp_int;
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_CONSTRAINED_CMP_HPP__

View file

@ -22,6 +22,7 @@
#include "dpf/bit.hpp" #include "dpf/bit.hpp"
#include "dpf/xor_wrapper.hpp" #include "dpf/xor_wrapper.hpp"
#include "dpf/twiddle.hpp" #include "dpf/twiddle.hpp"
#include "dpf/cmp_group.hpp"
namespace dpf namespace dpf
{ {
@ -39,8 +40,8 @@ template <typename ...Ts>
inline constexpr bool no_ic_pack_v = inline constexpr bool no_ic_pack_v =
(!is_ic_pack<std::decay_t<Ts>>::value && ...); (!is_ic_pack<std::decay_t<Ts>>::value && ...);
/// Comparison kind for the optional DCF channel on a key. /// @brief Comparison kind for the optional DCF channel on a key.
/// `lt`/`leq`/`gt`/`geq` are the comparison predicates. The later kinds are /// @details `lt`/`leq`/`gt`/`geq` are the comparison predicates. The later kinds are
/// path paints: one constant on each sibling subtree of the secret point, /// path paints: one constant on each sibling subtree of the secret point,
/// evaluated by the same value-correction walk. /// evaluated by the same value-correction walk.
enum class cmp_kind : uint8_t enum class cmp_kind : uint8_t
@ -58,7 +59,9 @@ enum class cmp_kind : uint8_t
paint = 10 // caller-supplied unit plant paint = 10 // caller-supplied unit plant
}; };
/// True for the path-paint kinds. Comparisons stay `lt`/`leq`/`gt`/`geq`. /// @brief True for the path-paint kinds. Comparisons stay `lt`/`leq`/`gt`/`geq`.
/// @param kind the comparison or paint kind
/// @return True for the path-paint kinds
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline constexpr bool is_paint_kind(cmp_kind kind) noexcept inline constexpr bool is_paint_kind(cmp_kind kind) noexcept
{ {
@ -77,7 +80,7 @@ inline constexpr bool is_paint_kind(cmp_kind kind) noexcept
} }
} }
/// Unit plant for `path_paint`. `prefix` is the in-lane matched prefix. /// @brief Unit plant for `path_paint`. `prefix` is the in-lane matched prefix.
using paint_callback = uint64_t (*)(std::size_t matched, uint64_t prefix, using paint_callback = uint64_t (*)(std::size_t matched, uint64_t prefix,
bool leaf, const void * ctx); bool leaf, const void * ctx);
@ -119,6 +122,16 @@ uint64_t beta_delta_u64(const Beta & if_true, const Beta & if_false,
return (static_cast<uint64_t>(if_true) ^ static_cast<uint64_t>(if_false)) return (static_cast<uint64_t>(if_true) ^ static_cast<uint64_t>(if_false))
& mask; & mask;
} }
else if constexpr (detail::has_integral_representation<Beta>::value)
{
using raw_type = typename Beta::integral_type;
using unsigned_type = utils::make_unsigned_t<raw_type>;
const auto t = static_cast<uint64_t>(
static_cast<unsigned_type>(if_true.integral_representation()));
const auto f = static_cast<uint64_t>(
static_cast<unsigned_type>(if_false.integral_representation()));
return (t - f) & mask;
}
else else
{ {
return (static_cast<uint64_t>(if_true) return (static_cast<uint64_t>(if_true)
@ -132,6 +145,13 @@ uint64_t beta_to_u64_simple(const Beta & beta, uint64_t mask) noexcept
{ {
if constexpr (std::is_same_v<Beta, dpf::bit>) if constexpr (std::is_same_v<Beta, dpf::bit>)
return (static_cast<bool>(beta) ? 1ULL : 0ULL) & mask; return (static_cast<bool>(beta) ? 1ULL : 0ULL) & mask;
else if constexpr (detail::has_integral_representation<Beta>::value)
{
using raw_type = typename Beta::integral_type;
using unsigned_type = utils::make_unsigned_t<raw_type>;
return static_cast<uint64_t>(static_cast<unsigned_type>(
beta.integral_representation())) & mask;
}
else else
return static_cast<uint64_t>(beta) & mask; return static_cast<uint64_t>(beta) & mask;
} }
@ -154,6 +174,8 @@ Beta u64_to_beta(uint64_t v) noexcept
{ {
if constexpr (std::is_same_v<Beta, dpf::bit>) if constexpr (std::is_same_v<Beta, dpf::bit>)
return dpf::bit{static_cast<bool>(v & 1u)}; return dpf::bit{static_cast<bool>(v & 1u)};
else if constexpr (detail::has_integral_representation<Beta>::value)
return Beta::from_raw(static_cast<typename Beta::integral_type>(v));
else else
return static_cast<Beta>(v); return static_cast<Beta>(v);
} }
@ -186,7 +208,10 @@ constexpr uint64_t sgn_m(uint8_t t1, uint64_t x, uint64_t mask) noexcept
return t1 ? neg_m(x, mask) : x; return t1 ? neg_m(x, mask) : x;
} }
/// Convert a GGM node to a group element (low 64 bits, control bits cleared). /// @brief Convert a GGM node to a group element (low 64 bits, control bits cleared).
/// @param n the `n`
/// @param mask the bit mask
/// @return the returned `uint64_t`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
uint64_t convert_node(simde__m128i n, uint64_t mask) noexcept uint64_t convert_node(simde__m128i n, uint64_t mask) noexcept
{ {
@ -194,11 +219,15 @@ uint64_t convert_node(simde__m128i n, uint64_t mask) noexcept
simde_mm_cvtsi128_si64(dpf::unset_lo_2bits(n))) & mask; simde_mm_cvtsi128_si64(dpf::unset_lo_2bits(n))) & mask;
} }
/// Draw the group-width blind `r` used to split the `cmp_addend` share. /// @brief Draw the group-width blind `r` used to split the `cmp_addend` share.
/// `sample` yields one interior block; only `popcount(mask)` live bits are /// @details `sample` yields one interior block; only `popcount(mask)` live bits are
/// kept, so the blind (and thus the addend share) never needs a full padded /// kept, so the blind (and thus the addend share) never needs a full padded
/// `uint64_t` on the wire. Dealer and Doerner–Shelat gen call this with the /// `uint64_t` on the wire. Dealer and Doerner–Shelat gen call this with the
/// same block source so their keys stay byte-identical (matched tapes). /// same block source so their keys stay byte-identical (matched tapes).
/// @tparam BlockSampler block sampler
/// @param mask the bit mask
/// @param sample the `sample`
/// @return the returned `uint64_t`
template <typename BlockSampler> template <typename BlockSampler>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -207,8 +236,19 @@ uint64_t sample_addend_blind(uint64_t mask, BlockSampler && sample) noexcept
return convert_node(dpf::unset_lo_2bits(sample()), mask); return convert_node(dpf::unset_lo_2bits(sample()), mask);
} }
/// One level of value CW on GGM children. Updates running `Va`. /// @brief One level of value CW on GGM children. Updates running `Va`.
/// `ai` is the keep-path bit of the (effective) threshold. /// @details `ai` is the keep-path bit of the (effective) threshold.
/// @param c0L the `c0L`
/// @param c0R the `c0R`
/// @param c1L the `c1L`
/// @param c1R the `c1R`
/// @param t0 the `t0`
/// @param t1 the `t1`
/// @param ai the `ai`
/// @param Va the `Va`
/// @param beta the payload
/// @param mask the bit mask
/// @return One level of value CW on GGM children
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline uint64_t make_value_cw(simde__m128i c0L, simde__m128i c0R, inline uint64_t make_value_cw(simde__m128i c0L, simde__m128i c0R,
simde__m128i c1L, simde__m128i c1R, uint8_t t0, uint8_t t1, int ai, simde__m128i c1L, simde__m128i c1R, uint8_t t0, uint8_t t1, int ai,
@ -239,8 +279,20 @@ inline uint64_t make_value_cw(simde__m128i c0L, simde__m128i c0R,
return vcw; return vcw;
} }
/// Same recurrence as `make_value_cw`, planting `plant` on the lose child /// @brief Same recurrence as `make_value_cw`, planting `plant` on the lose child
/// in both directions. `plant == 0` leaves the correction unchanged. /// in both directions. `plant == 0` leaves the correction unchanged.
/// @param c0L the `c0L`
/// @param c0R the `c0R`
/// @param c1L the `c1L`
/// @param c1R the `c1R`
/// @param t0 the `t0`
/// @param t1 the `t1`
/// @param ai the `ai`
/// @param Va the `Va`
/// @param plant the unit plant
/// @param mask the bit mask
/// @return Same recurrence as `make_value_cw`, planting `plant` on the lose child in both
/// directions
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline uint64_t make_value_cw_planted(simde__m128i c0L, simde__m128i c0R, inline uint64_t make_value_cw_planted(simde__m128i c0L, simde__m128i c0R,
simde__m128i c1L, simde__m128i c1R, uint8_t t0, uint8_t t1, int ai, simde__m128i c1L, simde__m128i c1R, uint8_t t0, uint8_t t1, int ai,
@ -280,7 +332,11 @@ inline unsigned __int128 paint_lane_mask(std::size_t nbits) noexcept
return (u128{1} << nbits) - 1; return (u128{1} << nbits) - 1;
} }
/// High `d` bits of an `nbits`-wide lane, in that lane's own positions. /// @brief High `d` bits of an `nbits`-wide lane, in that lane's own positions.
/// @param alpha the secret input point
/// @param nbits the width in bits
/// @param d the `d`
/// @return High `d` bits of an `nbits`-wide lane, in that lane's own positions
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline unsigned __int128 paint_high_bits(unsigned __int128 alpha, inline unsigned __int128 paint_high_bits(unsigned __int128 alpha,
std::size_t nbits, std::size_t d) noexcept std::size_t nbits, std::size_t d) noexcept
@ -306,9 +362,18 @@ inline unsigned __int128 paint_low_aligned(unsigned __int128 alpha,
return alpha >> (nbits - d); return alpha >> (nbits - d);
} }
/// Unit (β = 1) lose-subtree or leaf plant. The caller scales by δ. /// @brief Unit (β = 1) lose-subtree or leaf plant. The caller scales by δ.
/// `matched` is the number of leading bits already shared with α. A lose /// @details `matched` is the number of leading bits already shared with α. A lose
/// subtree at that depth reconstructs to this value; `leaf` is the full match. /// subtree at that depth reconstructs to this value; `leaf` is the full match.
/// @param kind the comparison or paint kind
/// @param matched the number of leading bits shared with the secret point
/// @param alpha the secret input point
/// @param nbits the width in bits
/// @param length_bits the bits used to store the prefix length
/// @param leaf the leaf value
/// @param fn the `fn`
/// @param ctx the `ctx`
/// @return Unit (β = 1) lose-subtree or leaf plant
inline uint64_t paint_unit(cmp_kind kind, std::size_t matched, inline uint64_t paint_unit(cmp_kind kind, std::size_t matched,
unsigned __int128 alpha, std::size_t nbits, std::size_t length_bits, unsigned __int128 alpha, std::size_t nbits, std::size_t length_bits,
bool leaf, paint_callback fn, const void * ctx) bool leaf, paint_callback fn, const void * ctx)
@ -357,8 +422,15 @@ inline uint64_t scale_plant(uint64_t unit, uint64_t scale, uint64_t mask) noexce
return (unit * scale) & mask; return (unit * scale) & mask;
} }
/// Final leaf value CW. `on_path` is the payload reconstructed when the query /// @brief Final leaf value CW. `on_path` is the payload reconstructed when the query
/// stays on α's path through all levels (0 for strict lt/geq; β for leq/gt). /// stays on α's path through all levels (0 for strict lt/geq; β for leq/gt).
/// @param s0 the `s0`
/// @param s1 the `s1`
/// @param t1 the `t1`
/// @param Va the `Va`
/// @param mask the bit mask
/// @param on_path the value reconstructed on the secret path
/// @return Final leaf value CW
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline uint64_t make_final_cw(simde__m128i s0, simde__m128i s1, uint8_t t1, inline uint64_t make_final_cw(simde__m128i s0, simde__m128i s1, uint8_t t1,
uint64_t Va, uint64_t mask, uint64_t on_path = 0) noexcept uint64_t Va, uint64_t mask, uint64_t on_path = 0) noexcept
@ -371,8 +443,8 @@ inline uint64_t make_final_cw(simde__m128i s0, simde__m128i s1, uint8_t t1,
} // namespace dcf_impl } // namespace dcf_impl
/// Comparison metadata on an incremental key (value CWs live on the key). /// @brief Comparison metadata on an incremental key (value CWs live on the key).
/// Payload δ = if_true − if_false is dealer-known and baked into `value_cw` / /// @details Payload δ = if_true − if_false is dealer-known and baked into `value_cw` /
/// `cw_last` only — never stored clear on the key (traditional DPF hiding). /// `cw_last` only — never stored clear on the key (traditional DPF hiding).
/// The second output value (`if_false`) is held as a per-party additive share /// The second output value (`if_false`) is held as a per-party additive share
/// on the key (`cmp_addend`), not as a public constant. /// on the key (`cmp_addend`), not as a public constant.
@ -393,7 +465,7 @@ struct cmp_meta
bool empty() const noexcept { return !active; } bool empty() const noexcept { return !active; }
}; };
/// Backward-compatible alias while call sites migrate. /// @brief Backward-compatible alias while call sites migrate.
using cmp_channel = cmp_meta; using cmp_channel = cmp_meta;
} // namespace detail } // namespace detail
@ -591,8 +663,13 @@ inline auto break_bit_at(Beta t = Beta{1},
return paint_at_pack<N, cmp_kind::break_bit, Beta>(std::move(t), std::move(f)); return paint_at_pack<N, cmp_kind::break_bit, Beta>(std::move(t), std::move(f));
} }
/// Low `LengthBits` hold the common-prefix length. Above them sits the /// @brief Low `LengthBits` hold the common-prefix length. Above them sits the
/// matched prefix packed into the low bits of the lane (`α >> (N − d)`). /// matched prefix packed into the low bits of the lane (`α >> (N − d)`).
/// @tparam LengthBits length bits
/// @tparam Beta payload type
/// @param t the `t`
/// @param f the `f`
/// @return Low `LengthBits` hold the common-prefix length
template <std::size_t LengthBits = 8, typename Beta = uint64_t> template <std::size_t LengthBits = 8, typename Beta = uint64_t>
inline auto prefix_with_length(Beta t = Beta{1}, inline auto prefix_with_length(Beta t = Beta{1},
Beta f = detail::dcf_impl::default_false<Beta>()) Beta f = detail::dcf_impl::default_false<Beta>())
@ -608,9 +685,11 @@ inline auto prefix_with_length_at(Beta t = Beta{1},
std::move(t), std::move(f)); std::move(t), std::move(f));
} }
/// Arbitrary unit plant. `fn(matched, in_lane_prefix, leaf)` returns the β = 1 /// @brief Arbitrary unit plant. `fn(matched, in_lane_prefix, leaf)` returns the β = 1
/// value of that sibling subtree (`leaf` is the full match, `matched == N`). /// value of that sibling subtree (`leaf` is the full match, `matched == N`).
/// The result is scaled by `if_true − if_false` like the canned recipes. /// @details The result is scaled by `if_true − if_false` like the canned recipes.
/// @tparam Beta payload type
/// @tparam Fn fn
template <typename Beta, typename Fn> template <typename Beta, typename Fn>
struct paint_fn_pack struct paint_fn_pack
{ {
@ -667,8 +746,9 @@ inline auto path_paint_at(Fn fn, Beta t,
std::move(t), std::move(f), std::move(fn)); std::move(t), std::move(f), std::move(fn));
} }
/// Incremental comparison: the same predicate, correct at every prefix length. /// @brief Incremental comparison: the same predicate, correct at every prefix length.
/// Evaluate the full point with `cmp`, and a prefix with `cmp_prefix<L>`. /// @details Evaluate the full point with `cmp`, and a prefix with `cmp_prefix<L>`.
/// @tparam Spec comparison or interval specification
template <typename Spec> template <typename Spec>
struct idcf_pack struct idcf_pack
{ {

View file

@ -25,11 +25,12 @@
#include "dpf/dpf_key.hpp" #include "dpf/dpf_key.hpp"
#include "dpf/random.hpp" #include "dpf/random.hpp"
#include "dpf/dcf.hpp" #include "dpf/dcf.hpp"
#include "dpf/constrained_cmp.hpp"
namespace dpf namespace dpf
{ {
/// Tag: Doerner–Shelat / geneval takes additive shares of the point /// @brief Tag: Doerner–Shelat / geneval takes additive shares of the point
/// (`x0 + x1` in the input ring). Default calls take XOR shares. /// (`x0 + x1` in the input ring). Default calls take XOR shares.
struct arith_input_t struct arith_input_t
{ {
@ -37,9 +38,60 @@ struct arith_input_t
inline constexpr arith_input_t arith_input{}; inline constexpr arith_input_t arith_input{};
/// Roots and the Beaver-pad stream for one Doerner–Shelat generation. /// @brief Tag: payload β is additively shared (`y0 + y1`). Leaf CW is opened via
/// `root` is called twice, same as `make_dpf`: party 0 clears the low bit of /// Π_CCMP on the on-path control bits (see `open_arith_leaf`).
/// @see `open_arith_leaf`
struct arith_output_t
{
};
inline constexpr arith_output_t arith_output{};
/// @brief Additive (or XOR) shares of one concrete payload for dealerless leaf open.
/// @details Use as a placed value / `at<>` element when several outputs are shared.
/// @tparam T value type
template <typename T>
struct arith_beta
{
using payload_type = T;
T y0{};
T y1{};
};
template <typename T>
struct is_arith_beta : std::false_type
{
};
template <typename T>
struct is_arith_beta<arith_beta<T>> : std::true_type
{
};
template <typename T>
inline constexpr bool is_arith_beta_v = is_arith_beta<std::decay_t<T>>::value;
namespace detail
{
namespace incr
{
/// @brief `placed<N, arith_beta<T>>::output_type` is `T` (see placement.hpp).
/// @tparam T value type
template <typename T>
struct unwrap_placed_output<arith_beta<T>>
{
using type = T;
};
} // namespace incr
} // namespace detail
/// @brief Roots and the Beaver-pad stream for one Doerner–Shelat generation.
/// @details `root` is called twice, same as `make_dpf`: party 0 clears the low bit of
/// the first sample, party 1 sets the low bit of the second. /// the first sample, party 1 sets the low bit of the second.
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
template <typename RootSampler, typename PadRng> template <typename RootSampler, typename PadRng>
struct ds_randomness struct ds_randomness
{ {
@ -227,7 +279,7 @@ inline simde__m128i ds_deliver(uint8_t b_exp, simde__m128i base, simde__m128i M,
return ds_xor(ds_xor(local, z.z0), z.z1); return ds_xor(ds_xor(local, z.z0), z.z1);
} }
/// Per-level messages prepared before the CW protocol runs (blinds + pads). /// @brief Per-level messages prepared before the CW protocol runs (blinds + pads).
struct ds_level_blinds struct ds_level_blinds
{ {
ds_cw_pads cwp; ds_cw_pads cwp;
@ -240,7 +292,7 @@ struct ds_level_blinds
uint8_t bit1; uint8_t bit1;
}; };
/// Opened CW, advice, and AND products delivered by a `CwProtocol`. /// @brief Opened CW, advice, and AND products delivered by a `CwProtocol`.
struct ds_level_open struct ds_level_open
{ {
simde__m128i cw; simde__m128i cw;
@ -250,8 +302,8 @@ struct ds_level_open
uint64_t value_cw = 0; // public after open when cmp is active at this level uint64_t value_cw = 0; // public after open when cmp is active at this level
}; };
/// Running comparison-gen state shared across DS levels (Va residual). /// @brief Running comparison-gen state shared across DS levels (Va residual).
/// When `track_coeff` is set (wildcard cmp payload), a parallel β = 1 /// @details When `track_coeff` is set (wildcard cmp payload), a parallel β = 1
/// accumulator `Va1` is advanced alongside `Va` so the gen can stash /// accumulator `Va1` is advanced alongside `Va` so the gen can stash
/// `value_cw(1) − value_cw(β)` coefficients for a later `assign_cmp`. /// `value_cw(1) − value_cw(β)` coefficients for a later `assign_cmp`.
struct ds_cmp_gen_state struct ds_cmp_gen_state
@ -274,8 +326,9 @@ struct ds_cmp_gen_state
const void * paint_ctx = nullptr; const void * paint_ctx = nullptr;
}; };
/// Local joint simulation: today's `ds_cw_outs` / `ds_open_advice` / `ds_and_open`. /// @brief Local joint simulation: today's `ds_cw_outs` / `ds_open_advice` / `ds_and_open`.
/// An MPC backend would send `blinds` and return the same `ds_level_open` shape. /// @details An MPC backend would send `blinds` and return the same `ds_level_open` shape.
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
template <typename PadRng> template <typename PadRng>
struct local_cw_protocol struct local_cw_protocol
{ {
@ -310,7 +363,11 @@ struct local_cw_protocol
return out; return out;
} }
/// Open CW + advice only (AND pads stay in `blinds` for a later open). /// @brief Open CW + advice only (AND pads stay in `blinds` for a later open).
/// @param b the `b`
/// @return the returned `std::pair<simde__m128i, uint8_t>`
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
HEDLEY_NO_THROW HEDLEY_NO_THROW
std::pair<simde__m128i, uint8_t> open_cw(const ds_level_blinds & b) noexcept std::pair<simde__m128i, uint8_t> open_cw(const ds_level_blinds & b) noexcept
{ {
@ -318,9 +375,18 @@ struct local_cw_protocol
b.b0, b.b1), b.b0, b.b1),
ds_open_advice(b.L0, b.R0, b.bit0, b.L1, b.R1, b.bit1)}; ds_open_advice(b.L0, b.R0, b.bit0, b.L1, b.R1, b.bit1)};
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
/// Open the public value CW for this level (local: clear convert+make_value_cw). /// @brief Open the public value CW for this level (local: clear convert+make_value_cw).
/// MPC backends open additive shares of the same word. /// @details MPC backends open additive shares of the same word.
/// @param b the `b`
/// @param adv0 the `adv0`
/// @param adv1 the `adv1`
/// @param ai the `ai`
/// @param Va the `Va`
/// @param beta the payload
/// @param mask the bit mask
/// @return the returned `uint64_t`
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t open_value_cw(const ds_level_blinds & b, uint8_t adv0, uint8_t adv1, uint64_t open_value_cw(const ds_level_blinds & b, uint8_t adv0, uint8_t adv1,
int ai, uint64_t & Va, uint64_t beta, uint64_t mask) noexcept int ai, uint64_t & Va, uint64_t beta, uint64_t mask) noexcept
@ -329,7 +395,15 @@ struct local_cw_protocol
adv1, ai, Va, beta, mask); adv1, ai, Va, beta, mask);
} }
/// Open a path-paint value CW. `plant` is the scaled lose-subtree constant. /// @brief Open a path-paint value CW. `plant` is the scaled lose-subtree constant.
/// @param b the `b`
/// @param adv0 the `adv0`
/// @param adv1 the `adv1`
/// @param ai the `ai`
/// @param Va the `Va`
/// @param plant the unit plant
/// @param mask the bit mask
/// @return the returned `uint64_t`
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t open_planted_cw(const ds_level_blinds & b, uint8_t adv0, uint8_t adv1, uint64_t open_planted_cw(const ds_level_blinds & b, uint8_t adv0, uint8_t adv1,
int ai, uint64_t & Va, uint64_t plant, uint64_t mask) noexcept int ai, uint64_t & Va, uint64_t plant, uint64_t mask) noexcept
@ -345,9 +419,16 @@ struct local_cw_protocol
return ds_and_open(p, M, b_recv); return ds_and_open(p, M, b_recv);
} }
/// Open the final comparison leaf CW. Wraps `make_final_cw` so the /// @brief Open the final comparison leaf CW. Wraps `make_final_cw` so the
/// Doerner–Shelat gen does not call it directly on reconstructed seeds; /// Doerner–Shelat gen does not call it directly on reconstructed seeds;
/// an MPC backend would open additive shares of the same word. /// an MPC backend would open additive shares of the same word.
/// @param s0 the `s0`
/// @param s1 the `s1`
/// @param t1 the `t1`
/// @param Va the `Va`
/// @param mask the bit mask
/// @param on_path the value reconstructed on the secret path
/// @return the returned `uint64_t`
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t open_final_cw(simde__m128i s0, simde__m128i s1, uint8_t t1, uint64_t open_final_cw(simde__m128i s0, simde__m128i s1, uint8_t t1,
uint64_t Va, uint64_t mask, uint64_t on_path) noexcept uint64_t Va, uint64_t mask, uint64_t on_path) noexcept
@ -355,9 +436,13 @@ struct local_cw_protocol
return dcf_impl::make_final_cw(s0, s1, t1, Va, mask, on_path); return dcf_impl::make_final_cw(s0, s1, t1, Va, mask, on_path);
} }
/// Draw the group-width `cmp_addend` blind. Local joint simulation reuses /// @brief Draw the group-width `cmp_addend` blind. Local joint simulation reuses
/// the shared root sampler so the blind matches the dealer's; an MPC /// the shared root sampler so the blind matches the dealer's; an MPC
/// backend would instead pull a group-width element from the pad stream. /// backend would instead pull a group-width element from the pad stream.
/// @tparam BlockSampler block sampler
/// @param mask the bit mask
/// @param sample the `sample`
/// @return the returned `uint64_t`
template <typename BlockSampler> template <typename BlockSampler>
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t sample_addend_blind(uint64_t mask, BlockSampler && sample) noexcept uint64_t sample_addend_blind(uint64_t mask, BlockSampler && sample) noexcept
@ -366,14 +451,23 @@ struct local_cw_protocol
std::forward<BlockSampler>(sample)); std::forward<BlockSampler>(sample));
} }
/// Majority of three bits (next carry of a full adder). /// @brief Majority of three bits (next carry of a full adder).
/// @param a the `a`
/// @param b the `b`
/// @param c the `c`
/// @return Majority of three bits (next carry of a full adder)
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr uint8_t majority(uint8_t a, uint8_t b, uint8_t c) noexcept static constexpr uint8_t majority(uint8_t a, uint8_t b, uint8_t c) noexcept
{ {
return static_cast<uint8_t>((a & b) | (a & c) | (b & c)); return static_cast<uint8_t>((a & b) | (a & c) | (b & c));
} }
/// One additive digit: sum bit `a XOR b XOR cin`, carry out = majority. /// @brief One additive digit: sum bit `a XOR b XOR cin`, carry out = majority.
/// @param a the `a`
/// @param b the `b`
/// @param cin the `cin`
/// @param cout the `cout`
/// @return One additive digit: sum bit `a XOR b XOR cin`, carry out = majority
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr uint8_t open_sum_bit(uint8_t a, uint8_t b, uint8_t cin, static constexpr uint8_t open_sum_bit(uint8_t a, uint8_t b, uint8_t cin,
uint8_t & cout) noexcept uint8_t & cout) noexcept
@ -382,9 +476,14 @@ struct local_cw_protocol
return static_cast<uint8_t>(a ^ b ^ cin); return static_cast<uint8_t>(a ^ b ^ cin);
} }
/// Reconstruct `a0 + a1` via a LSB→MSB carry chain, then flip the MSB /// @brief Reconstruct `a0 + a1` via a LSB→MSB carry chain, then flip the MSB
/// when the domain is signed — matching `make_dpf` on the sum. The call /// when the domain is signed — matching `make_dpf` on the sum. The call
/// site never forms the sum; an MPC backend would open the same bits. /// site never forms the sum; an MPC backend would open the same bits.
/// @tparam InputT input domain type
/// @param a0 the `a0`
/// @param a1 the `a1`
/// @return Reconstruct `a0 + a1` via a LSB→MSB carry chain, then flip the MSB when the domain
/// is signed — matching `make_dpf` on the sum
template <typename InputT> template <typename InputT>
InputT open_arith_point(InputT a0, InputT a1) const InputT open_arith_point(InputT a0, InputT a1) const
{ {
@ -409,9 +508,13 @@ struct local_cw_protocol
return out; return out;
} }
/// Encode shares for the XOR-style CW walk. XOR mode flips party 0's MSB /// @brief Encode shares for the XOR-style CW walk. XOR mode flips party 0's MSB
/// (linear over XOR). Arithmetic mode opens the sum (carry + signed MSB) /// (linear over XOR). Arithmetic mode opens the sum (carry + signed MSB)
/// and returns `(alpha, 0)` so the walk matches `make_dpf(alpha)`. /// and returns `(alpha, 0)` so the walk matches `make_dpf(alpha)`.
/// @tparam InputT input domain type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param arith the `arith`
template <typename InputT> template <typename InputT>
void encode_walk_shares(InputT & x0, InputT & x1, bool arith) const void encode_walk_shares(InputT & x0, InputT & x1, bool arith) const
{ {
@ -427,13 +530,84 @@ struct local_cw_protocol
} }
} }
/// Open a group of leaf correction words for one prefix group. In this /// @brief Constrained comparison Π_CCMP: open `1{x0 < x1}` when `|x0−x1|=1`.
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @return Constrained comparison Π_CCMP: open `1{x0 < x1}` when `|x0−x1|=1`
HEDLEY_NO_THROW
uint8_t open_ccmp(std::uint64_t x0, std::uint64_t x1) noexcept
{
(void)pads; // MPC backend would consume an AND pad here
return dpf::local_ccmp(x0, x1);
}
/// @brief Open a public leaf CW for a shared payload.
/// @details Ring: `β = y0 + y1`; `g = CCMP(t0,t1)` selects `β − M` vs `M − β`
/// (matches `make_leaf` with `sign = t0`). Characteristic 2: `β = y0 ⊕ y1`
/// and CW = `β ⊕ M` (sign mux is a no-op under XOR).
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam I output index
/// @tparam OutputsTuple outputs tuple
/// @tparam InteriorBlock interior block
/// @tparam OutputT output type
/// @param seed0 the `seed0`
/// @param seed1 the `seed1`
/// @param t0 the `t0`
/// @param t1 the `t1`
/// @param y0 the `y0`
/// @param y1 the `y1`
/// @param pos_base the `pos_base`
/// @param lane_x lane of the shared payload
/// @return the opened leaf correction word
template <typename ExteriorPRG, std::size_t I = 0, typename OutputsTuple,
typename InteriorBlock, typename OutputT>
auto open_arith_leaf(const InteriorBlock & seed0, const InteriorBlock & seed1,
uint8_t t0, uint8_t t1, OutputT y0, OutputT y1, std::size_t pos_base,
std::size_t lane_x) -> dpf::leaf_node_t<typename ExteriorPRG::block_type,
OutputT>
{
using output_type = OutputT;
using node_type = typename ExteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using leaf_type = dpf::leaf_node_t<node_type, output_type>;
HEDLEY_PRAGMA(GCC diagnostic pop)
const auto M = dpf::make_leaf_mask<ExteriorPRG, I, OutputsTuple, InteriorBlock>(
seed0, seed1, pos_base);
output_type beta{};
if constexpr (utils::has_characteristic_two_v<output_type>)
{
(void)t0;
(void)t1;
beta = static_cast<output_type>(y0 ^ y1);
const leaf_type naked = dpf::make_naked_leaf<node_type>(lane_x, beta);
return dpf::subtract_leaf<output_type>(naked, M);
}
else
{
const uint8_t g = open_ccmp(t0, t1);
beta = static_cast<output_type>(y0 + y1);
const leaf_type naked = dpf::make_naked_leaf<node_type>(lane_x, beta);
// CW = (−1)^{t1}(β − M): g=0 → β−M; g=1 → M−β. Matches make_leaf(sign=t0).
if (g & 1u)
return dpf::subtract_leaf<output_type>(M, naked);
return dpf::subtract_leaf<output_type>(naked, M);
}
}
/// @brief Open a group of leaf correction words for one prefix group. In this
/// local joint simulation both XOR shares of the point are present, so the /// local joint simulation both XOR shares of the point are present, so the
/// point is reconstructed *inside* the protocol and handed to `leaf_fn` /// point is reconstructed *inside* the protocol and handed to `leaf_fn`
/// (which runs `make_leaves` for the group). The Doerner–Shelat gen never /// (which runs `make_leaves` for the group). The Doerner–Shelat gen never
/// forms `x = x0 ^ x1` at its own call site; an MPC backend would instead /// forms `x = x0 ^ x1` at its own call site; an MPC backend would instead
/// run a per-group leaf CW exchange that never reveals `x`. After /// run a per-group leaf CW exchange that never reveals `x`. After
/// `encode_walk_shares`, arithmetic inputs are already `(alpha, 0)`. /// `encode_walk_shares`, arithmetic inputs are already `(alpha, 0)`.
/// @tparam InputT input domain type
/// @tparam LeafFn leaf fn
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param leaf_fn the `leaf_fn`
template <typename InputT, typename LeafFn> template <typename InputT, typename LeafFn>
void open_leaf_group(InputT x0, InputT x1, LeafFn && leaf_fn) void open_leaf_group(InputT x0, InputT x1, LeafFn && leaf_fn)
{ {
@ -441,7 +615,8 @@ struct local_cw_protocol
} }
}; };
/// Generation-side level state (seeds / home bits). Not an eval path memoizer. /// @brief Generation-side level state (seeds / home bits). Not an eval path memoizer.
/// @tparam NodeT GGM node type
template <typename NodeT> template <typename NodeT>
struct ds_gen_state struct ds_gen_state
{ {
@ -471,16 +646,34 @@ struct ds_gen_state
const NodeT & seed1() const noexcept { return inbox[home[1]]; } const NodeT & seed1() const noexcept { return inbox[home[1]]; }
}; };
/// One interior level: expand, protocol open, advance both party seeds. /// @brief One interior level: expand, protocol open, advance both party seeds.
/// When `cmp` is non-null and active for `level`, also opens `value_cw` via /// @details When `cmp` is non-null and active for `level`, also opens `value_cw` via
/// the protocol (no second PRG expand outside). /// the protocol (no second PRG expand outside).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam CwProtocol correction-word protocol
/// @tparam NodeT GGM node type
/// @tparam InputT input domain type
/// @tparam MaskT mask type
/// @tparam AdviceT advice type
/// @param st the `st`
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param mask the bit mask
/// @param level the tree level
/// @param depth the tree depth
/// @param proto the `proto`
/// @param cw_out the `cw_out`
/// @param advice_out the `advice_out`
/// @param value_cw_out the `value_cw_out`
/// @param cmp the comparison specification
template <typename InteriorPRG, typename CwProtocol, typename NodeT, template <typename InteriorPRG, typename CwProtocol, typename NodeT,
typename InputT, typename MaskT, typename AdviceT> typename InputT, typename MaskT, typename AdviceT>
void ds_advance_level(ds_gen_state<NodeT> & st, InputT x0, InputT x1, void ds_advance_level(ds_gen_state<NodeT> & st, InputT x0, InputT x1,
MaskT mask, std::size_t level, CwProtocol & proto, NodeT & cw_out, MaskT mask, std::size_t level, std::size_t depth, CwProtocol & proto,
AdviceT & advice_out, uint64_t * value_cw_out = nullptr, NodeT & cw_out, AdviceT & advice_out, uint64_t * value_cw_out = nullptr,
ds_cmp_gen_state * cmp = nullptr) ds_cmp_gen_state * cmp = nullptr)
{ {
using tree = dpf::tree_traits<InteriorPRG>;
// Integral bridge so bit extraction works for `keyword` / `modint` / // Integral bridge so bit extraction works for `keyword` / `modint` /
// signed / bitstring the same way dealer gen does via `mask & x`. // signed / bitstring the same way dealer gen does via `mask & x`.
// `msb_mask` is the unsigned bit pattern; a signed input must not be // `msb_mask` is the unsigned bit pattern; a signed input must not be
@ -490,21 +683,28 @@ void ds_advance_level(ds_gen_state<NodeT> & st, InputT x0, InputT x1,
const auto mi = to_mask(mask); const auto mi = to_mask(mask);
const uint8_t bit0 = static_cast<uint8_t>(!!(mi & to_int(x0))); const uint8_t bit0 = static_cast<uint8_t>(!!(mi & to_int(x0)));
const uint8_t bit1 = static_cast<uint8_t>(!!(mi & to_int(x1))); const uint8_t bit1 = static_cast<uint8_t>(!!(mi & to_int(x1)));
const bool is_last = tree::is_last_level(level, depth);
NodeT s0 = st.seed0(); NodeT s0 = st.seed0();
NodeT s1 = st.seed1(); NodeT s1 = st.seed1();
const uint8_t adv0 = static_cast<uint8_t>( const uint8_t adv0 = static_cast<uint8_t>(dpf::get_lo_bit(s0));
dpf::get_lo_bit_and_clear_lo_2bits(s0)); const uint8_t adv1 = static_cast<uint8_t>(dpf::get_lo_bit(s1));
const uint8_t adv1 = static_cast<uint8_t>( const auto c0 = tree::expand(s0, is_last);
dpf::get_lo_bit_and_clear_lo_2bits(s1)); const auto c1 = tree::expand(s1, is_last);
const auto c0 = InteriorPRG::eval01(s0);
const auto c1 = InteriorPRG::eval01(s1);
auto blinds = proto.prepare_level(c0[0], c0[1], bit0, c1[0], c1[1], bit1); auto blinds = proto.prepare_level(c0[0], c0[1], bit0, c1[0], c1[1], bit1);
if (value_cw_out != nullptr && cmp != nullptr && cmp->active if (value_cw_out != nullptr && cmp != nullptr && cmp->active
&& cmp->trivial == cmp_trivial::none && level < cmp->nbits) && cmp->trivial == cmp_trivial::none && level < cmp->nbits)
{ {
// Convert uses expand_value (HT: always two-tweak); seed walk used expand.
const auto v0 = tree::expand_value(s0);
const auto v1 = tree::expand_value(s1);
auto vblinds = blinds;
vblinds.L0 = v0[0];
vblinds.R0 = v0[1];
vblinds.L1 = v1[0];
vblinds.R1 = v1[1];
const int ai = static_cast<int>( const int ai = static_cast<int>(
(cmp->thresh >> (cmp->nbits - 1 - level)) & 1); (cmp->thresh >> (cmp->nbits - 1 - level)) & 1);
if (cmp->paint) if (cmp->paint)
@ -514,34 +714,47 @@ void ds_advance_level(ds_gen_state<NodeT> & st, InputT x0, InputT x1,
cmp->paint_cb, cmp->paint_ctx); cmp->paint_cb, cmp->paint_ctx);
const uint64_t plant = dcf_impl::scale_plant(unit, cmp->beta, const uint64_t plant = dcf_impl::scale_plant(unit, cmp->beta,
cmp->mask); cmp->mask);
*value_cw_out = proto.open_planted_cw(blinds, adv0, adv1, ai, *value_cw_out = proto.open_planted_cw(vblinds, adv0, adv1, ai,
cmp->Va, plant, cmp->mask); cmp->Va, plant, cmp->mask);
if (cmp->track_coeff) if (cmp->track_coeff)
{ {
const uint64_t plant1 = dcf_impl::scale_plant(unit, 1ULL, const uint64_t plant1 = dcf_impl::scale_plant(unit, 1ULL,
cmp->mask); cmp->mask);
const uint64_t v1 = proto.open_planted_cw(blinds, adv0, adv1, const uint64_t v1w = proto.open_planted_cw(vblinds, adv0, adv1,
ai, cmp->Va1, plant1, cmp->mask); ai, cmp->Va1, plant1, cmp->mask);
cmp->last_vcw_coeff = cmp->last_vcw_coeff =
(v1 + dcf_impl::neg_m(*value_cw_out, cmp->mask)) & cmp->mask; (v1w + dcf_impl::neg_m(*value_cw_out, cmp->mask)) & cmp->mask;
} }
} }
else else
{ {
*value_cw_out = proto.open_value_cw(blinds, adv0, adv1, ai, cmp->Va, *value_cw_out = proto.open_value_cw(vblinds, adv0, adv1, ai, cmp->Va,
cmp->beta, cmp->mask); cmp->beta, cmp->mask);
if (cmp->track_coeff) if (cmp->track_coeff)
{ {
// Affine coefficient: same level with β = 1 on a parallel Va. // Affine coefficient: same level with β = 1 on a parallel Va.
const uint64_t v1 = proto.open_value_cw(blinds, adv0, adv1, ai, const uint64_t v1w = proto.open_value_cw(vblinds, adv0, adv1, ai,
cmp->Va1, 1ULL, cmp->mask); cmp->Va1, 1ULL, cmp->mask);
cmp->last_vcw_coeff = cmp->last_vcw_coeff =
(v1 + dcf_impl::neg_m(*value_cw_out, cmp->mask)) & cmp->mask; (v1w + dcf_impl::neg_m(*value_cw_out, cmp->mask)) & cmp->mask;
} }
} }
} }
auto [cw, tpack] = proto.open_cw(blinds); auto [cw, tpack] = proto.open_cw(blinds);
// Half-Tree mid levels store no advice; last level keeps BGI packing.
if constexpr (tree::is_half_tree)
{
if (!is_last)
tpack = 0;
}
else
{
// BGI: opened advice stands.
}
// Dealer-equivalent CW for Half-Tree mid: off-path children XOR already
// matches H(s0)⊕H(s1)⊕ᾱΔ via the open. For last/BGI, open matches Gen.
const uint8_t exp0 = st.home[0] == 0 ? bit0 : bit1; const uint8_t exp0 = st.home[0] == 0 ? bit0 : bit1;
const uint8_t rec0 = st.home[0] == 0 ? bit1 : bit0; const uint8_t rec0 = st.home[0] == 0 ? bit1 : bit0;
@ -549,8 +762,29 @@ void ds_advance_level(ds_gen_state<NodeT> & st, InputT x0, InputT x1,
const uint8_t rec1 = st.home[1] == 0 ? bit1 : bit0; const uint8_t rec1 = st.home[1] == 0 ? bit1 : bit0;
NodeT M0, base0, M1, base1; NodeT M0, base0, M1, base1;
if constexpr (tree::is_half_tree)
{
if (!is_last)
{
// Mid: next = child[bit] ⊕ (t ? full_cw : 0).
const NodeT D0 = ds_xor(c0[0], c0[1]);
const NodeT D1 = ds_xor(c1[0], c1[1]);
M0 = D0;
base0 = (adv0 & 1u) ? ds_xor(c0[0], cw) : c0[0];
M1 = D1;
base1 = (adv1 & 1u) ? ds_xor(c1[0], cw) : c1[0];
}
else
{
ds_next_terms(c0[0], c0[1], adv0, cw, tpack, M0, base0); ds_next_terms(c0[0], c0[1], adv0, cw, tpack, M0, base0);
ds_next_terms(c1[0], c1[1], adv1, cw, tpack, M1, base1); ds_next_terms(c1[0], c1[1], adv1, cw, tpack, M1, base1);
}
}
else
{
ds_next_terms(c0[0], c0[1], adv0, cw, tpack, M0, base0);
ds_next_terms(c1[0], c1[1], adv1, cw, tpack, M1, base1);
}
const NodeT nxt0 = const NodeT nxt0 =
ds_deliver(exp0, base0, M0, proto.open_and(blinds.and0, M0, rec0)); ds_deliver(exp0, base0, M0, proto.open_and(blinds.and0, M0, rec0));
const NodeT nxt1 = const NodeT nxt1 =
@ -598,9 +832,124 @@ template <typename InteriorPRG,
typename ExteriorPRG, typename ExteriorPRG,
typename InputT, typename InputT,
typename OutputT, typename OutputT,
typename ...OutputTs,
typename RootSampler, typename RootSampler,
typename CwProtocol> typename CwProtocol>
auto make_dpf_doerner_shelat_impl(bool arith, bool arith_out, InputT x0,
InputT x1, RootSampler & root_sampler, CwProtocol & proto, OutputT y0,
OutputT y1 = OutputT{})
{
static_assert(!dpf::is_wildcard_v<InputT>,
"Doerner–Shelat gen takes shares of a concrete point");
static_assert(!dpf::is_secret_share_v<InputT>,
"Doerner–Shelat: pass additive_share of xor_wrapper, or raw shares");
static_assert(!dpf::is_wildcard_v<OutputT>,
"arith_output / classic DS leaf expects a concrete payload");
static_assert(sizeof(typename InteriorPRG::block_type) == sizeof(simde__m128i),
"Doerner–Shelat gen uses the AES-block interior node");
using dpf_type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, InputT, OutputT>;
using node = typename dpf_type::interior_node;
using input_type = typename dpf_type::input_type;
using leaf_tuple = typename dpf_type::leaf_tuple;
using beaver_tuple = typename dpf_type::beaver_tuple;
using outputs_tuple = std::tuple<OutputT>;
constexpr auto depth = dpf_type::depth;
proto.encode_walk_shares(x0, x1, arith);
using tree = dpf::tree_traits<InteriorPRG>;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
node roots[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(roots, [&]() -> node {
return static_cast<node>(root_sampler());
});
const node root0 = roots[0];
const node root1 = roots[1];
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_gen_state<node> st;
HEDLEY_PRAGMA(GCC diagnostic pop)
st.init(root0, root1);
typename dpf_type::correction_words_array correction_words{};
typename dpf_type::correction_advice_array correction_advice{};
auto mask = dpf_type::msb_mask;
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{
ds_advance_level<InteriorPRG>(st, x0, x1, mask, level, depth, proto,
correction_words[level], correction_advice[level]);
}
const node parent0 = st.seed0();
const node parent1 = st.seed1();
const bool sign0 = dpf::get_lo_bit(parent0);
const uint8_t t0 = static_cast<uint8_t>(sign0);
const uint8_t t1 = static_cast<uint8_t>(dpf::get_lo_bit(parent1));
input_type x = utils::xor_input_shares(x0, x1);
leaf_tuple leaves0{};
leaf_tuple leaves1{};
beaver_tuple beavers0{};
beaver_tuple beavers1{};
if (arith_out)
{
constexpr auto to_int = utils::to_integral_type<input_type>{};
const std::size_t lane = static_cast<std::size_t>(to_int(x));
auto cw = proto.template open_arith_leaf<ExteriorPRG, 0, outputs_tuple>(
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), t0, t1,
y0, y1, std::size_t{0}, lane);
std::get<0>(leaves0) = cw;
std::get<0>(leaves1) = cw;
}
else
{
auto built = dpf::make_leaves<ExteriorPRG>(x,
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), sign0,
std::size_t{0}, y0);
leaves0 = std::move(built.first.first);
beavers0 = std::move(built.first.second);
leaves1 = std::move(built.second.first);
beavers1 = std::move(built.second.second);
(void)y1;
}
input_type off0{};
input_type off1{};
return dpf::make_party_key_pair(
dpf_type{root0, correction_words, correction_advice,
leaves0, beavers0, off0},
dpf_type{root1, correction_words, correction_advice,
leaves1, beavers1, off1});
}
/// @brief Plaintext-β multi-output classic path (unchanged).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputTs output ts
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam CwProtocol correction-word protocol
/// @param arith the `arith`
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param root_sampler the `root_sampler`
/// @param proto the `proto`
/// @param y the `y`
/// @param ys the `ys`
/// @return Plaintext-β multi-output classic path (unchanged)
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename RootSampler,
typename CwProtocol,
typename = std::enable_if_t<(sizeof...(OutputTs) > 0)>>
auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1, auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1,
RootSampler & root_sampler, CwProtocol & proto, OutputT && y, RootSampler & root_sampler, CwProtocol & proto, OutputT && y,
OutputTs && ...ys) OutputTs && ...ys)
@ -620,10 +969,21 @@ auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1,
proto.encode_walk_shares(x0, x1, arith); proto.encode_walk_shares(x0, x1, arith);
const node root0 = dpf::unset_lo_bit(static_cast<node>(root_sampler())); using tree = dpf::tree_traits<InteriorPRG>;
const node root1 = dpf::set_lo_bit(static_cast<node>(root_sampler())); HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
node roots[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(roots, [&]() -> node {
return static_cast<node>(root_sampler());
});
const node root0 = roots[0];
const node root1 = roots[1];
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_gen_state<node> st; ds_gen_state<node> st;
HEDLEY_PRAGMA(GCC diagnostic pop)
st.init(root0, root1); st.init(root0, root1);
typename dpf_type::correction_words_array correction_words{}; typename dpf_type::correction_words_array correction_words{};
@ -632,7 +992,7 @@ auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1,
auto mask = dpf_type::msb_mask; auto mask = dpf_type::msb_mask;
for (std::size_t level = 0; level < depth; ++level, mask >>= 1) for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{ {
ds_advance_level<InteriorPRG>(st, x0, x1, mask, level, proto, ds_advance_level<InteriorPRG>(st, x0, x1, mask, level, depth, proto,
correction_words[level], correction_advice[level]); correction_words[level], correction_advice[level]);
} }
@ -654,9 +1014,37 @@ auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1,
built.second.first, built.second.second, off1}); built.second.first, built.second.second, off1});
} }
/// @brief Single-output plaintext β (disambiguates from arith_out overload).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam CwProtocol correction-word protocol
/// @param arith the `arith`
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param root_sampler the `root_sampler`
/// @param proto the `proto`
/// @param y the `y`
/// @return Single-output plaintext β (disambiguates from arith_out overload)
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename CwProtocol>
auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1,
RootSampler & root_sampler, CwProtocol & proto, OutputT && y)
{
return make_dpf_doerner_shelat_impl<InteriorPRG, ExteriorPRG>(arith, false,
std::move(x0), std::move(x1), root_sampler, proto,
std::forward<OutputT>(y), OutputT{});
}
} // namespace detail } // namespace detail
/// Local CW protocol (pads cancel; same keys as dealer when roots match). /// @brief Local CW protocol (pads cancel; same keys as dealer when roots match).
template <typename PadRng> template <typename PadRng>
using local_cw_protocol = detail::local_cw_protocol<PadRng>; using local_cw_protocol = detail::local_cw_protocol<PadRng>;

View file

@ -1,6 +1,5 @@
/// @file dpf/dpf_key.hpp /// @file dpf/dpf_key.hpp
/// @brief /// @brief The DPF key, its correction words, and interior traversal.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -19,6 +18,7 @@
#include <atomic> #include <atomic>
#include "dpf/prg_aes.hpp" #include "dpf/prg_aes.hpp"
#include "dpf/tree_traits.hpp"
#include "dpf/wildcard.hpp" #include "dpf/wildcard.hpp"
#include "dpf/twiddle.hpp" #include "dpf/twiddle.hpp"
#include "dpf/leaf_node.hpp" #include "dpf/leaf_node.hpp"
@ -26,6 +26,7 @@
#include "dpf/leaf_wrapper.hpp" #include "dpf/leaf_wrapper.hpp"
#include "dpf/emplace.hpp" #include "dpf/emplace.hpp"
#include "dpf/placement.hpp" #include "dpf/placement.hpp"
#include "dpf/verifiable.hpp"
#include "dpf/dcf.hpp" #include "dpf/dcf.hpp"
namespace dpf namespace dpf
@ -88,16 +89,25 @@ auto make_dpfargs(InputT && x, OutputT && y, OutputTs && ...ys)
std::forward<OutputTs>(ys)...) }; std::forward<OutputTs>(ys)...) };
} }
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
template <typename InteriorPRG> template <typename InteriorPRG>
using root_sampler_t = std::add_pointer_t<typename InteriorPRG::block_type()>; using root_sampler_t = std::add_pointer_t<typename InteriorPRG::block_type()>;
HEDLEY_PRAGMA(GCC diagnostic pop)
namespace detail namespace detail
{ {
/// Classic single-level DPF key body (all outputs bare, at full input width, /// @brief Classic single-level DPF key body (all outputs bare, at full input width,
/// equal widths, no comparison channel). `Derived` is the public `dpf_key` /// equal widths, no comparison channel). `Derived` is the public `dpf_key`
/// specialization that inherits this body — threaded through only so that /// specialization that inherits this body — threaded through only so that
/// `emplace`/`emplace_back` construct the public key type. /// `emplace`/`emplace_back` construct the public key type.
/// @tparam Derived derived
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputTs output ts
template <typename Derived, template <typename Derived,
typename InteriorPRG, typename InteriorPRG,
typename ExteriorPRG, typename ExteriorPRG,
@ -108,6 +118,7 @@ struct classic_dpf_key_impl
{ {
public: public:
using interior_prg = InteriorPRG; using interior_prg = InteriorPRG;
using tree = dpf::tree_traits<InteriorPRG>;
using interior_node = typename InteriorPRG::block_type; using interior_node = typename InteriorPRG::block_type;
using exterior_prg = ExteriorPRG; using exterior_prg = ExteriorPRG;
@ -156,9 +167,18 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
static constexpr std::size_t cmp_h = 0; static constexpr std::size_t cmp_h = 0;
static constexpr std::size_t cmp_checkpoints = 0; static constexpr std::size_t cmp_checkpoints = 0;
static constexpr std::size_t cmp_tail = 0; static constexpr std::size_t cmp_tail = 0;
/// Classic keys are single-level; the unified eval surface keeps routing /// @brief Classic keys are single-level; the unified eval surface keeps routing
/// them through the classic `eval_*` fast paths (see `is_multilevel_key`). /// them through the classic `eval_*` fast paths (see `is_multilevel_key`).
/// @see `is_multilevel_key`
static constexpr bool is_multilevel = false; static constexpr bool is_multilevel = false;
static constexpr bool is_verifiable = false;
static constexpr bool is_extractable = false;
using correction_seeds_array = std::array<cs_block, 0>;
const correction_seeds_array & correction_seeds() const
{
static const correction_seeds_array empty{};
return empty;
}
private: private:
using meta_placed_tuple = std::tuple< using meta_placed_tuple = std::tuple<
@ -178,7 +198,10 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
static constexpr std::size_t outputs_per_leaf_of = static constexpr std::size_t outputs_per_leaf_of =
std::size_t{1} << lg_outputs_per_leaf_of<I>; std::size_t{1} << lg_outputs_per_leaf_of<I>;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using correction_words_array = std::array<interior_node, depth>; using correction_words_array = std::array<interior_node, depth>;
HEDLEY_PRAGMA(GCC diagnostic pop)
using correction_advice_array = std::array<psnip_uint8_t, depth>; using correction_advice_array = std::array<psnip_uint8_t, depth>;
template <typename Emplaceable> template <typename Emplaceable>
@ -267,8 +290,9 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
auto correction_word(std::size_t level, bool direction) const auto correction_word(std::size_t level, bool direction) const
{ {
return set_lo_bit(correction_word(level), return tree::pack_cw(correction_word(level),
(correction_advice_[level] >> direction) & 1); correction_advice_[level], direction,
tree::is_last_level(level, depth));
} }
template <std::size_t I = 0> template <std::size_t I = 0>
@ -318,58 +342,46 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
static auto traverse_interior(const interior_node & node, static auto traverse_interior(const interior_node & node,
const interior_node & cw, bool dir) noexcept const interior_node & cw, bool dir, bool is_last = false) noexcept
{ {
return dpf::xor_if_lo_bit( return tree::traverse(node, cw, dir, is_last);
interior_prg::eval(unset_lo_2bits(node), dir), cw, node);
} }
/// Expand both children of `node` with one pipelined `eval01`. /// @brief Expand both children of `node` with one pipelined expand.
/// Equivalent to `traverse_interior(node, cw0, 0)` and /// @details Equivalent to `traverse_interior(node, cw0, 0)` and
/// `traverse_interior(node, cw1, 1)`, but the two AES-128 blocks share /// `traverse_interior(node, cw1, 1)`. Full-domain interval eval uses this
/// a round loop. Full-domain interval eval uses this at almost every /// at almost every interior parent.
/// interior parent. /// @param node the GGM node
/// @param cw0 correction word for the left child
/// @param cw1 correction word for the right child
/// @param is_last whether this is the last interior level
/// @return both children of `node`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
static auto traverse_interior01(const interior_node & node, static auto traverse_interior01(const interior_node & node,
const interior_node & cw0, const interior_node & cw1) noexcept const interior_node & cw0, const interior_node & cw1,
bool is_last = false) noexcept
{ {
HEDLEY_PRAGMA(GCC diagnostic push) return tree::traverse01(node, cw0, cw1, is_last);
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto kids = interior_prg::eval01(unset_lo_2bits(node));
return std::array<interior_node, 2>{
dpf::xor_if_lo_bit(kids[0], cw0, node),
dpf::xor_if_lo_bit(kids[1], cw1, node)
};
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
/// Four independent `traverse_interior01` via `InteriorPRG::eval01_x4`. /// @brief Four independent `traverse_interior01` via traits batched expand.
/// `left[i]` / `right[i]` are the children of `parents[i]`. /// @details `left[i]` / `right[i]` are the children of `parents[i]`.
/// @param parents the `parents`
/// @param cw0 the `cw0`
/// @param cw1 the `cw1`
/// @param left the `left`
/// @param right the `right`
/// @param is_last the `is_last`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
static void traverse_interior01_x4(const interior_node * HEDLEY_RESTRICT parents, static void traverse_interior01_x4(const interior_node * HEDLEY_RESTRICT parents,
const interior_node & cw0, const interior_node & cw1, const interior_node & cw0, const interior_node & cw1,
interior_node * HEDLEY_RESTRICT left, interior_node * HEDLEY_RESTRICT left,
interior_node * HEDLEY_RESTRICT right) noexcept interior_node * HEDLEY_RESTRICT right, bool is_last = false) noexcept
{ {
HEDLEY_PRAGMA(GCC diagnostic push) tree::traverse01_x4(parents, cw0, cw1, left, right, is_last);
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
alignas(interior_node) interior_node seeds[4];
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
seeds[i] = unset_lo_2bits(parents[i]);
}
interior_prg::eval01_x4(seeds, left, right);
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
left[i] = dpf::xor_if_lo_bit(left[i], cw0, parents[i]);
right[i] = dpf::xor_if_lo_bit(right[i], cw1, parents[i]);
}
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <std::size_t I = 0, template <std::size_t I = 0,
@ -383,11 +395,11 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using output_type = std::tuple_element_t<I, concrete_outputs_tuple>; using output_type = std::tuple_element_t<I, concrete_outputs_tuple>;
HEDLEY_PRAGMA(GCC diagnostic pop)
// Subtractive share: CW_if_t − mask so reconstruct(y0, y1) = y0 − y1 = β. // Subtractive share: CW_if_t − mask so reconstruct(y0, y1) = y0 − y1 = β.
return dpf::subtract_leaf<output_type>( return dpf::subtract_leaf<output_type>(
dpf::get_if_lo_bit(correction_word, node), dpf::get_if_lo_bit(correction_word, node),
make_leaf_mask_inner<exterior_prg, I, concrete_outputs_tuple>(unset_lo_2bits(node))); make_leaf_mask_inner<exterior_prg, I, concrete_outputs_tuple>(unset_lo_2bits(node)));
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <std::size_t I = 0> template <std::size_t I = 0>
@ -414,9 +426,12 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
{ {
return std::apply([&leaf..., &beaver...](auto & ...foo) return std::apply([&leaf..., &beaver...](auto & ...foo)
{ {
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::make_tuple( return std::make_tuple(
dpf::leaf_wrapper<std::decay_t<decltype(foo)>, exterior_node>(leaf, beaver)... dpf::leaf_wrapper<std::decay_t<decltype(foo)>, exterior_node>(leaf, beaver)...
); );
HEDLEY_PRAGMA(GCC diagnostic pop)
}, tmp); }, tmp);
}, beavers); }, beavers);
}, leaves); }, leaves);
@ -439,16 +454,21 @@ namespace detail
namespace incr namespace incr
{ {
/// Comparison-channel storage. Value CWs, `cw_last`, and the `cmp_addend` /// @brief Comparison-channel storage. Value CWs, `cw_last`, and the `cmp_addend`
/// share are held at the comparison group width (`ValueCwWord`), not a full /// share are held at the comparison group width (`ValueCwWord`), not a full
/// padded `uint64_t` per level: a bit comparison carries 1 byte/level, a /// padded `uint64_t` per level: a bit comparison carries 1 byte/level, a
/// `uint16_t` payload 2 bytes/level, etc. Arithmetic still runs in `uint64_t` /// `uint16_t` payload 2 bytes/level, etc. Arithmetic still runs in `uint64_t`
/// (masked); the narrow word is only the on-key / on-wire representation. /// (masked); the narrow word is only the on-key / on-wire representation.
/// Extra per-level δ-coefficients kept only for wildcard comparison payloads /// @details Extra per-level δ-coefficients kept only for wildcard comparison payloads
/// (empty for concrete cmp keys, so their layout is unchanged). The value CWs /// (empty for concrete cmp keys, so their layout is unchanged). The value CWs
/// / `cw_last` are affine in the payload δ, so after keygen with δ = 0 the /// / `cw_last` are affine in the payload δ, so after keygen with δ = 0 the
/// concrete values are `base[i] + coeff[i]·δ`; `assign_cmp` patches them in /// concrete values are `base[i] + coeff[i]·δ`; `assign_cmp` patches them in
/// place with no tree re-walk / re-PRG. /// place with no tree re-walk / re-PRG.
/// @tparam Depth depth
/// @tparam ValueCwWord value cw word
/// @tparam Wild whether the payload is a wildcard
/// @tparam TailLen tail len
/// @tparam Idcf idcf
template <std::size_t Depth, typename ValueCwWord, bool Wild, template <std::size_t Depth, typename ValueCwWord, bool Wild,
std::size_t TailLen = 0, bool Idcf = false> std::size_t TailLen = 0, bool Idcf = false>
struct cmp_wild_state { }; struct cmp_wild_state { };
@ -520,6 +540,10 @@ struct cmp_storage
return static_cast<uint64_t>(value_cw_[level]); return static_cast<uint64_t>(value_cw_[level]);
} }
HEDLEY_NO_THROW HEDLEY_NO_THROW
value_cw_word cw_last_word() const noexcept { return cw_last_; }
HEDLEY_NO_THROW
value_cw_word cmp_addend_word() const noexcept { return cmp_addend_; }
HEDLEY_NO_THROW
uint64_t cw_last() const noexcept { return static_cast<uint64_t>(cw_last_); } uint64_t cw_last() const noexcept { return static_cast<uint64_t>(cw_last_); }
HEDLEY_NO_THROW HEDLEY_NO_THROW
const tail_array & tail_cw() const noexcept { return tail_; } const tail_array & tail_cw() const noexcept { return tail_; }
@ -553,9 +577,11 @@ struct cmp_storage
return true; return true;
} }
/// Patch the (public) value CWs / `cw_last` in place for a resolved δ and /// @brief Patch the (public) value CWs / `cw_last` in place for a resolved δ and
/// install this party's `cmp_addend` share. No-op on the CWs when there is /// install this party's `cmp_addend` share. No-op on the CWs when there is
/// no wildcard coefficient table (trivial domain-edge cmp). /// no wildcard coefficient table (trivial domain-edge cmp).
/// @param delta the payload difference `if_true - if_false`
/// @param addend_share the `addend_share`
void assign_cmp_delta(uint64_t delta, uint64_t addend_share) void assign_cmp_delta(uint64_t delta, uint64_t addend_share)
{ {
static_assert(Wild, static_assert(Wild,
@ -595,6 +621,117 @@ struct cmp_storage
} }
} }
/// @brief Overwrite the public final correction and this party's addend. Used when
/// the group element does not fit in the `uint64_t` the constructor takes.
/// @param last the past-the-end element of the range
/// @param addend the additive share of the off-point payload
/// @param last_coeff the `last_coeff`
void set_scalars(value_cw_word last, value_cw_word addend,
value_cw_word last_coeff)
{
cw_last_ = last;
cmp_addend_ = addend;
if constexpr (Wild)
wild_.cw_last_coeff = last_coeff;
else
(void)last_coeff;
}
/// @brief `base + coeff · δ` in `delta`'s group, then install `addend`.
/// @param delta the payload difference `if_true - if_false`
/// @param addend the additive share of the off-point payload
void assign_group(const detail::group_elem & delta, value_cw_word addend)
{
static_assert(Wild,
"assign_cmp on a key whose comparison payload is not a wildcard");
if constexpr (Wild)
{
auto mix = [&](value_cw_word base_w, value_cw_word coeff_w) {
const auto base = detail::group_from_word(base_w, delta);
const auto coeff = detail::group_from_word(coeff_w, delta);
return detail::group_to_word<value_cw_word>(
detail::group_add(base, detail::group_mul(coeff, delta)));
};
for (std::size_t i = 0; i < Depth; ++i)
value_cw_[i] = mix(value_cw_[i], wild_.value_cw_coeff[i]);
if constexpr (Blocked)
{
for (std::size_t i = 0; i < TailLen; ++i)
tail_[i] = mix(tail_[i], wild_.tail_coeff[i]);
}
cw_last_ = mix(cw_last_, wild_.cw_last_coeff);
if constexpr (Idcf)
{
for (std::size_t i = 0; i < prefix_cw_len; ++i)
prefix_cw_[i] = mix(prefix_cw_[i], wild_.prefix_cw_coeff[i]);
}
cmp_addend_ = addend;
wild_.assigned = true;
}
}
/// @brief Per-level δ coefficients for a wildcard comparison. Empty when the
/// payload is concrete.
/// @return the coefficient table
const value_cw_array & value_cw_coeff() const noexcept
{
if constexpr (Wild)
return wild_.value_cw_coeff;
else
{
static const value_cw_array empty{};
return empty;
}
}
/// @brief Coefficient of δ in `cw_last`. Zero when the payload is concrete.
/// @return the final coefficient word
HEDLEY_NO_THROW
value_cw_word cw_last_coeff_word() const noexcept
{
if constexpr (Wild)
return wild_.cw_last_coeff;
else
return value_cw_word{};
}
/// @brief Tail δ coefficients for a blocked wildcard comparison.
/// @return the tail coefficient table
const tail_array & tail_coeff() const noexcept
{
if constexpr (Wild)
return wild_.tail_coeff;
else
{
static const tail_array empty{};
return empty;
}
}
/// @brief Prefix δ coefficients for an iDCF wildcard comparison.
/// @return the prefix coefficient table
const prefix_cw_array & prefix_cw_coeff() const noexcept
{
if constexpr (Wild)
return wild_.prefix_cw_coeff;
else
{
static const prefix_cw_array empty{};
return empty;
}
}
/// @brief Mark whether a wildcard comparison payload has been assigned.
/// @param assigned true once `assign_cmp` has run
HEDLEY_NO_THROW
void set_assigned(bool assigned) noexcept
{
if constexpr (Wild)
wild_.assigned = assigned;
else
(void)assigned;
}
private: private:
detail::cmp_meta cmp_{}; detail::cmp_meta cmp_{};
value_cw_array value_cw_{}; value_cw_array value_cw_{};
@ -605,33 +742,46 @@ struct cmp_storage
cmp_wild_state<Depth, value_cw_word, Wild, TailLen, Idcf> wild_{}; cmp_wild_state<Depth, value_cw_word, Wild, TailLen, Idcf> wild_{};
}; };
/// Multi-level / comparison DPF key body. `PlacedTuple` is a tuple of /// @brief Multi-level / comparison DPF key body. `PlacedTuple` is a tuple of
/// `placed<N, T>` slots; `CmpDepth > 0` activates the comparison channel. /// `placed<N, T>` slots; `CmpDepth > 0` activates the comparison channel.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam PlacedTuple placed tuple
/// @tparam CmpDepth cmp depth
/// @tparam CmpOutBits cmp out bits
/// @tparam CmpWild cmp wild
/// @tparam CmpBlock cmp block
/// @tparam CmpIdcf cmp idcf
template <typename InteriorPRG, typename ExteriorPRG, typename InputT, template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, std::size_t CmpDepth = 0, typename PlacedTuple, std::size_t CmpDepth = 0,
std::size_t CmpOutBits = 0, bool CmpWild = false, std::size_t CmpOutBits = 0, bool CmpWild = false,
std::size_t CmpBlock = 0, bool CmpIdcf = false> std::size_t CmpBlock = 0, bool CmpIdcf = false,
bool IsVerifiable = false, bool IsExtractable = false>
struct incr_key_base struct incr_key_base
{ {
public: public:
using interior_prg = InteriorPRG; using interior_prg = InteriorPRG;
using exterior_prg = ExteriorPRG; using exterior_prg = ExteriorPRG;
using tree = dpf::tree_traits<InteriorPRG>;
using interior_node = typename InteriorPRG::block_type; using interior_node = typename InteriorPRG::block_type;
using exterior_node = typename ExteriorPRG::block_type; using exterior_node = typename ExteriorPRG::block_type;
using input_type = dpf::concrete_type_t<InputT>; using input_type = dpf::concrete_type_t<InputT>;
using placed_tuple = PlacedTuple; using placed_tuple = PlacedTuple;
using node_type = exterior_node; using node_type = exterior_node;
static constexpr std::size_t cmp_depth = CmpDepth; static constexpr std::size_t cmp_depth = CmpDepth;
/// Comparison output group width in bits (0 when there is no cmp channel). /// @brief Comparison output group width in bits (0 when there is no cmp channel).
static constexpr std::size_t cmp_out_bits = CmpOutBits; static constexpr std::size_t cmp_out_bits = CmpOutBits;
/// True when the comparison payload is an unassigned wildcard. /// @brief True when the comparison payload is an unassigned wildcard.
static constexpr bool cmp_is_wildcard = CmpWild; static constexpr bool cmp_is_wildcard = CmpWild;
/// 0 = per-level path-sum. `B >= 1` = blocked checkpoints of width `B`. /// @brief 0 = per-level path-sum. `B >= 1` = blocked checkpoints of width `B`.
static constexpr std::size_t cmp_block = CmpBlock; static constexpr std::size_t cmp_block = CmpBlock;
static constexpr bool cmp_idcf = CmpIdcf; static constexpr bool cmp_idcf = CmpIdcf;
static constexpr bool is_verifiable = IsVerifiable;
static constexpr bool is_extractable = IsExtractable;
static constexpr std::size_t max_output_level = static constexpr std::size_t max_output_level =
detail::incr::max_tree_level_v<node_type, PlacedTuple>; detail::incr::max_tree_level_v<node_type, PlacedTuple>;
/// Residual tail width. 2 only when dropping those levels does not cut an /// @brief Residual tail width. 2 only when dropping those levels does not cut an
/// output and the comparison itself is what sets the tree height. /// output and the comparison itself is what sets the tree height.
static constexpr std::size_t cmp_q = [] { static constexpr std::size_t cmp_q = [] {
if (CmpBlock == 0 || CmpDepth <= 2) if (CmpBlock == 0 || CmpDepth <= 2)
@ -646,9 +796,7 @@ struct incr_key_base
(CmpBlock == 0 || cmp_h == 0) ? 0 : (cmp_h + CmpBlock - 1) / CmpBlock; (CmpBlock == 0 || cmp_h == 0) ? 0 : (cmp_h + CmpBlock - 1) / CmpBlock;
static constexpr std::size_t cmp_tail = static constexpr std::size_t cmp_tail =
(CmpBlock == 0 || cmp_q == 0) ? 0 : (std::size_t{1} << cmp_q); (CmpBlock == 0 || cmp_q == 0) ? 0 : (std::size_t{1} << cmp_q);
/// Multi-level / comparison keys route through the slot-aware eval path. /// @brief Narrowest unsigned word that holds `cmp_out_bits` bits (1 byte for a
static constexpr bool is_multilevel = true;
/// Narrowest unsigned word that holds `cmp_out_bits` bits (1 byte for a
/// bit / ≤8-bit payload, 2 for ≤16, 4 for ≤32, 8 for ≤64). Value CWs and /// bit / ≤8-bit payload, 2 for ≤16, 4 for ≤32, 8 for ≤64). Value CWs and
/// the addend share are stored in this word. /// the addend share are stored in this word.
using value_cw_word = utils::integral_type_from_bitlength_t< using value_cw_word = utils::integral_type_from_bitlength_t<
@ -667,11 +815,15 @@ struct incr_key_base
static_assert(num_outputs > 0 || CmpDepth > 0, static_assert(num_outputs > 0 || CmpDepth > 0,
"incremental DPF needs at least one output or a comparison channel"); "incremental DPF needs at least one output or a comparison channel");
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
static_assert(detail::incr::all_prefixes_ok_v<node_type, PlacedTuple>, static_assert(detail::incr::all_prefixes_ok_v<node_type, PlacedTuple>,
"at<N> is shorter than the packing lanes required by an output"); "at<N> is shorter than the packing lanes required by an output");
using correction_words_array = std::array<interior_node, depth>; using correction_words_array = std::array<interior_node, depth>;
HEDLEY_PRAGMA(GCC diagnostic pop)
using correction_advice_array = std::array<psnip_uint8_t, depth>; using correction_advice_array = std::array<psnip_uint8_t, depth>;
using correction_seeds_array = std::array<cs_block, IsVerifiable ? depth : 0>;
using value_cw_array = std::array<value_cw_word, value_cw_len>; using value_cw_array = std::array<value_cw_word, value_cw_len>;
using tail_array = std::array<value_cw_word, cmp_tail>; using tail_array = std::array<value_cw_word, cmp_tail>;
static constexpr std::size_t prefix_cw_len = CmpIdcf ? depth + 1 : 0; static constexpr std::size_t prefix_cw_len = CmpIdcf ? depth + 1 : 0;
@ -680,6 +832,26 @@ struct incr_key_base
static constexpr meta_array meta = static constexpr meta_array meta =
detail::incr::build_meta<node_type, PlacedTuple>(); detail::incr::build_meta<node_type, PlacedTuple>();
/// @brief Multi-level / comparison keys route through the slot-aware eval path.
/// Classic-shaped packs (every slot at full input width, no cmp) keep the
/// classic `eval_*` fast path even when verifiable/extractable phantoms are
/// present.
static constexpr bool is_multilevel = [] {
if constexpr (CmpDepth > 0)
return true;
if constexpr (num_outputs == 0)
return true;
else
{
for (std::size_t i = 0; i < num_outputs; ++i)
{
if (meta[i].prefix != input_bits)
return true;
}
return false;
}
}();
template <std::size_t I, typename = void> template <std::size_t I, typename = void>
struct output_type_at struct output_type_at
{ {
@ -714,7 +886,7 @@ struct incr_key_base
public: public:
using leaf_wrapper_tuple = decltype(wrapper_tuple_t( using leaf_wrapper_tuple = decltype(wrapper_tuple_t(
std::make_index_sequence<num_outputs>{})); std::make_index_sequence<num_outputs>{}));
/// Raw leaf shares (pre-wrapper), matching classic `leaf_tuple` for asio. /// @brief Raw leaf shares (pre-wrapper), matching classic `leaf_tuple` for asio.
using leaf_tuple = decltype(leaf_tuple_type( using leaf_tuple = decltype(leaf_tuple_type(
std::make_index_sequence<num_outputs>{})); std::make_index_sequence<num_outputs>{}));
using offset_type = offset_wrapper<InputT>; using offset_type = offset_wrapper<InputT>;
@ -739,7 +911,7 @@ struct incr_key_base
} }
}(); }();
/// First output (source order) whose prefix equals `deepest_prefix`. /// @brief First output (source order) whose prefix equals `deepest_prefix`.
static constexpr std::size_t deepest_output = [] { static constexpr std::size_t deepest_output = [] {
if constexpr (num_outputs == 0) if constexpr (num_outputs == 0)
return std::size_t{0}; return std::size_t{0};
@ -754,7 +926,7 @@ struct incr_key_base
} }
}(); }();
/// Classic-shaped packing traits for deepest-group interval/sequence APIs. /// @brief Classic-shaped packing traits for deepest-group interval/sequence APIs.
static constexpr std::size_t outputs_per_leaf = static constexpr std::size_t outputs_per_leaf =
(num_outputs > 0) ? outputs_per_leaf_of<deepest_output> : 1; (num_outputs > 0) ? outputs_per_leaf_of<deepest_output> : 1;
static constexpr std::size_t lg_outputs_per_leaf = static constexpr std::size_t lg_outputs_per_leaf =
@ -775,7 +947,8 @@ struct incr_key_base
addend_tuple addends = {}, value_cw_array value_cw_coeff = {}, addend_tuple addends = {}, value_cw_array value_cw_coeff = {},
uint64_t cw_last_coeff_in = 0, tail_array tail_in = {}, uint64_t cw_last_coeff_in = 0, tail_array tail_in = {},
tail_array tail_coeff_in = {}, prefix_cw_array prefix_in = {}, tail_array tail_coeff_in = {}, prefix_cw_array prefix_in = {},
prefix_cw_array prefix_coeff_in = {}) prefix_cw_array prefix_coeff_in = {},
correction_seeds_array correction_seeds = {})
: leaf_nodes{std::move(leaves)}, : leaf_nodes{std::move(leaves)},
offset_x{offset_share}, offset_x{offset_share},
cmp_store_{cmp, value_cws, cmp_store_{cmp, value_cws,
@ -788,8 +961,9 @@ struct incr_key_base
root_{root}, root_{root},
correction_words_{correction_words}, correction_words_{correction_words},
correction_advice_{correction_advice}, correction_advice_{correction_advice},
correction_seeds_{correction_seeds},
common_part_hash_{utils::get_common_part_hash(correction_words_, common_part_hash_{utils::get_common_part_hash(correction_words_,
correction_advice_, leaf_nodes, wildcard_mask)} correction_advice_, leaf_nodes, wildcard_mask, correction_seeds_)}
{ } { }
incr_key_base(const incr_key_base &) = default; incr_key_base(const incr_key_base &) = default;
@ -806,17 +980,63 @@ struct incr_key_base
{ {
return correction_advice_; return correction_advice_;
} }
const correction_seeds_array & correction_seeds() const
{
return correction_seeds_;
}
const value_cw_array & value_cw() const { return cmp_store_.value_cw(); } const value_cw_array & value_cw() const { return cmp_store_.value_cw(); }
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t cw_last() const noexcept { return cmp_store_.cw_last(); } uint64_t cw_last() const noexcept { return cmp_store_.cw_last(); }
HEDLEY_NO_THROW HEDLEY_NO_THROW
value_cw_word cw_last_word() const noexcept { return cmp_store_.cw_last_word(); }
HEDLEY_NO_THROW
value_cw_word cmp_addend_word() const noexcept
{
return cmp_store_.cmp_addend_word();
}
void set_cmp_scalars(value_cw_word last, value_cw_word addend,
value_cw_word last_coeff)
{
cmp_store_.set_scalars(last, addend, last_coeff);
}
/// @brief Mark whether a wildcard comparison payload has been assigned.
/// @param assigned true once `assign_cmp` has run
HEDLEY_NO_THROW
void set_cmp_assigned(bool assigned) noexcept
{
cmp_store_.set_assigned(assigned);
}
const value_cw_array & value_cw_coeff() const noexcept
{
return cmp_store_.value_cw_coeff();
}
HEDLEY_NO_THROW
value_cw_word cw_last_coeff_word() const noexcept
{
return cmp_store_.cw_last_coeff_word();
}
const tail_array & tail_coeff() const noexcept
{
return cmp_store_.tail_coeff();
}
const prefix_cw_array & prefix_cw_coeff() const noexcept
{
return cmp_store_.prefix_cw_coeff();
}
void assign_cmp_group(const detail::group_elem & delta, value_cw_word addend)
{
cmp_store_.assign_group(delta, addend);
}
HEDLEY_NO_THROW
const prefix_cw_array & prefix_cws() const noexcept const prefix_cw_array & prefix_cws() const noexcept
{ {
return cmp_store_.prefix_cws(); return cmp_store_.prefix_cws();
} }
uint64_t prefix_cw(std::size_t i) const { return cmp_store_.prefix_cw(i); } uint64_t prefix_cw(std::size_t i) const { return cmp_store_.prefix_cw(i); }
/// Party-local share of the constant absorb (`if_false`, or /// @brief Party-local share of the constant absorb (`if_false`, or
/// `δ + if_false` when `eval_as_ge`). Reconstructs with the peer share. /// `δ + if_false` when `eval_as_ge`). Reconstructs with the peer share.
/// @return Party-local share of the constant absorb (`if_false`, or `δ + if_false` when
/// `eval_as_ge`)
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t cmp_addend() const noexcept { return cmp_store_.cmp_addend(); } uint64_t cmp_addend() const noexcept { return cmp_store_.cmp_addend(); }
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -834,8 +1054,9 @@ struct incr_key_base
} }
auto correction_word(std::size_t level, bool direction) const auto correction_word(std::size_t level, bool direction) const
{ {
return set_lo_bit(correction_word(level), return tree::pack_cw(correction_word(level),
(correction_advice_[level] >> direction) & 1); correction_advice_[level], direction,
tree::is_last_level(level, depth));
} }
uint64_t value_cw(std::size_t level) const { return cmp_store_.value_cw(level); } uint64_t value_cw(std::size_t level) const { return cmp_store_.value_cw(level); }
const tail_array & tail_cw() const { return cmp_store_.tail_cw(); } const tail_array & tail_cw() const { return cmp_store_.tail_cw(); }
@ -879,25 +1100,19 @@ struct incr_key_base
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
static auto traverse_interior(const interior_node & node, static auto traverse_interior(const interior_node & node,
const interior_node & cw, bool dir) noexcept const interior_node & cw, bool dir, bool is_last = false) noexcept
{ {
return dpf::xor_if_lo_bit( return tree::traverse(node, cw, dir, is_last);
interior_prg::eval(unset_lo_2bits(node), dir), cw, node);
} }
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
static auto traverse_interior01(const interior_node & node, static auto traverse_interior01(const interior_node & node,
const interior_node & cw0, const interior_node & cw1) noexcept const interior_node & cw0, const interior_node & cw1,
bool is_last = false) noexcept
{ {
HEDLEY_PRAGMA(GCC diagnostic push) return tree::traverse01(node, cw0, cw1, is_last);
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto kids = interior_prg::eval01(unset_lo_2bits(node));
return std::array<interior_node, 2>{
dpf::xor_if_lo_bit(kids[0], cw0, node),
dpf::xor_if_lo_bit(kids[1], cw1, node)};
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -905,22 +1120,9 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
static void traverse_interior01_x4(const interior_node * HEDLEY_RESTRICT parents, static void traverse_interior01_x4(const interior_node * HEDLEY_RESTRICT parents,
const interior_node & cw0, const interior_node & cw1, const interior_node & cw0, const interior_node & cw1,
interior_node * HEDLEY_RESTRICT left, interior_node * HEDLEY_RESTRICT left,
interior_node * HEDLEY_RESTRICT right) noexcept interior_node * HEDLEY_RESTRICT right, bool is_last = false) noexcept
{ {
HEDLEY_PRAGMA(GCC diagnostic push) tree::traverse01_x4(parents, cw0, cw1, left, right, is_last);
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
alignas(interior_node) interior_node seeds[4];
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
seeds[i] = unset_lo_2bits(parents[i]);
interior_prg::eval01_x4(seeds, left, right);
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
left[i] = dpf::xor_if_lo_bit(left[i], cw0, parents[i]);
right[i] = dpf::xor_if_lo_bit(right[i], cw1, parents[i]);
}
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <std::size_t I = 0> template <std::size_t I = 0>
@ -932,33 +1134,49 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
constexpr auto pos = constexpr auto pos =
meta[I].pos_base + meta[I].index_in_group * meta[I].block_len; meta[I].pos_base + meta[I].index_in_group * meta[I].block_len;
constexpr auto count = meta[I].block_len; constexpr auto count = meta[I].block_len;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using leaf_type = dpf::leaf_node_t<exterior_node, Out>; using leaf_type = dpf::leaf_node_t<exterior_node, Out>;
HEDLEY_PRAGMA(GCC diagnostic pop)
leaf_type mask{}; leaf_type mask{};
auto seed_ = auto seed_ =
utils::to_exterior_node<exterior_node>(unset_lo_2bits(node)); utils::to_exterior_node<exterior_node>(unset_lo_2bits(node));
if constexpr (IsExtractable)
{
detail::vdpf::extractable_leaf_prg<exterior_prg>::eval(seed_,
leaf_blocks<exterior_node>(mask),
static_cast<psnip_uint32_t>(count),
static_cast<psnip_uint32_t>(pos));
}
else
{
exterior_prg::eval(seed_, leaf_blocks<exterior_node>(mask), exterior_prg::eval(seed_, leaf_blocks<exterior_node>(mask),
static_cast<psnip_uint32_t>(count), static_cast<psnip_uint32_t>(count),
static_cast<psnip_uint32_t>(pos)); static_cast<psnip_uint32_t>(pos));
// Subtractive share: CW_if_t − mask so reconstruct(y0, y1) = y0 − y1 = β. }
return dpf::subtract_leaf<Out>( return dpf::subtract_leaf<Out>(
dpf::get_if_lo_bit(std::get<I>(leaf_nodes).get(), node), mask); dpf::get_if_lo_bit(std::get<I>(leaf_nodes).get(), node), mask);
} }
leaf_wrapper_tuple leaf_nodes; leaf_wrapper_tuple leaf_nodes;
offset_type offset_x; offset_type offset_x;
/// Public `if_false` addends for `eq` / `eq_at` slots. /// @brief Public `if_false` addends for `eq` / `eq_at` slots.
addend_tuple public_addends{}; addend_tuple public_addends{};
HEDLEY_NO_THROW HEDLEY_NO_THROW
bool has_cmp() const noexcept { return cmp_store_.has_cmp(); } bool has_cmp() const noexcept { return cmp_store_.has_cmp(); }
/// True once a wildcard comparison payload has been assigned (always true /// @brief True once a wildcard comparison payload has been assigned (always true
/// for concrete cmp keys and for keys without a comparison channel). /// for concrete cmp keys and for keys without a comparison channel).
/// @return True once a wildcard comparison payload has been assigned (always true for concrete
/// cmp keys and for keys without a comparison channel)
HEDLEY_NO_THROW HEDLEY_NO_THROW
bool cmp_assigned() const noexcept { return cmp_store_.cmp_assigned(); } bool cmp_assigned() const noexcept { return cmp_store_.cmp_assigned(); }
/// Patch the value CWs / `cw_last` for a resolved payload δ and install /// @brief Patch the value CWs / `cw_last` for a resolved payload δ and install
/// this party's `cmp_addend` share. Only valid for wildcard cmp keys; see /// this party's `cmp_addend` share. Only valid for wildcard cmp keys; see
/// the free `dpf::assign_cmp`. No tree re-walk / re-PRG. /// the free `dpf::assign_cmp`. No tree re-walk / re-PRG.
/// @param delta the payload difference `if_true - if_false`
/// @param addend_share the `addend_share`
void assign_cmp_delta(uint64_t delta, uint64_t addend_share) void assign_cmp_delta(uint64_t delta, uint64_t addend_share)
{ {
cmp_store_.assign_cmp_delta(delta, addend_share); cmp_store_.assign_cmp_delta(delta, addend_share);
@ -971,6 +1189,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
interior_node root_; interior_node root_;
correction_words_array correction_words_; correction_words_array correction_words_;
correction_advice_array correction_advice_; correction_advice_array correction_advice_;
correction_seeds_array correction_seeds_{};
digest_type common_part_hash_; digest_type common_part_hash_;
}; // struct incr_key_base }; // struct incr_key_base
@ -1013,7 +1232,13 @@ using dpf_key_base_t = std::conditional_t<
OutputT, OutputTs...>::cmp_block, OutputT, OutputTs...>::cmp_block,
dpf::detail::incr::normalize_pack< dpf::detail::incr::normalize_pack<
utils::bitlength_of_v<dpf::concrete_type_t<InputT>>, utils::bitlength_of_v<dpf::concrete_type_t<InputT>>,
OutputT, OutputTs...>::cmp_idcf>>; OutputT, OutputTs...>::cmp_idcf,
dpf::detail::incr::normalize_pack<
utils::bitlength_of_v<dpf::concrete_type_t<InputT>>,
OutputT, OutputTs...>::is_verifiable,
dpf::detail::incr::normalize_pack<
utils::bitlength_of_v<dpf::concrete_type_t<InputT>>,
OutputT, OutputTs...>::is_extractable>>;
} // namespace detail } // namespace detail
@ -1038,44 +1263,122 @@ namespace detail
namespace incr namespace incr
{ {
// Assemble the public dpf_key type for a (PlacedTuple, CmpDepth) pair by // Assemble the public dpf_key type for a (PlacedTuple, CmpDepth, flags) pack by
// expanding the placed slots into the output pack and appending the phantom // expanding the placed slots into the output pack and appending phantom tags.
// cmp tag when a comparison channel is present. //
template <std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild, // `dpf_key` always takes an output type. A comparison-only key (empty placed
std::size_t CmpBlock, bool CmpIdcf, typename InteriorPRG, // pack, CmpDepth > 0) is `dpf_key<..., cmp_channel_tag<...>>`. The no-comparison
typename ExteriorPRG, typename InputT, typename ...Ps> // form is a separate specialization so an empty pack is not named as
struct assemble_key // `dpf_key<Interior, Exterior, Input>` — `std::conditional_t` would require
// that type to be valid even when CmpDepth > 0.
template <bool WithCmp, typename InteriorPRG, typename ExteriorPRG,
typename InputT, std::size_t CmpDepth, std::size_t CmpOutBits,
bool CmpWild, std::size_t CmpBlock, bool CmpIdcf, typename ...Ps>
struct plain_assembled_key;
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild,
std::size_t CmpBlock, bool CmpIdcf, typename ...Ps>
struct plain_assembled_key<true, InteriorPRG, ExteriorPRG, InputT,
CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, Ps...>
{ {
using type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps..., using type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps...,
dpf::cmp_channel_tag<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf>>; dpf::cmp_channel_tag<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf>>;
}; };
template <std::size_t CmpOutBits, bool CmpWild, std::size_t CmpBlock,
bool CmpIdcf, typename InteriorPRG, template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename ExteriorPRG, typename InputT, typename ...Ps> std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild,
struct assemble_key<0, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, InteriorPRG, std::size_t CmpBlock, bool CmpIdcf, typename P0, typename ...Ps>
ExteriorPRG, InputT, Ps...> struct plain_assembled_key<false, InteriorPRG, ExteriorPRG, InputT,
CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, P0, Ps...>
{ {
using type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps...>; using type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, P0, Ps...>;
};
template <std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild,
std::size_t CmpBlock, bool CmpIdcf, bool IsVerifiable,
bool IsExtractable, typename InteriorPRG, typename ExteriorPRG,
typename InputT, typename ...Ps>
struct assemble_key;
template <std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild,
std::size_t CmpBlock, bool CmpIdcf, typename InteriorPRG,
typename ExteriorPRG, typename InputT, typename ...Ps>
struct assemble_key<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, false,
false, InteriorPRG, ExteriorPRG, InputT, Ps...>
{
using type = typename plain_assembled_key<(CmpDepth > 0),
InteriorPRG, ExteriorPRG, InputT, CmpDepth, CmpOutBits, CmpWild,
CmpBlock, CmpIdcf, Ps...>::type;
};
template <std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild,
std::size_t CmpBlock, bool CmpIdcf, typename InteriorPRG,
typename ExteriorPRG, typename InputT, typename ...Ps>
struct assemble_key<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, true,
false, InteriorPRG, ExteriorPRG, InputT, Ps...>
{
using with_cmp = std::conditional_t<
(CmpDepth > 0),
dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps...,
dpf::cmp_channel_tag<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf>,
dpf::verifiable>,
dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps..., dpf::verifiable>>;
using type = with_cmp;
};
template <std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild,
std::size_t CmpBlock, bool CmpIdcf, typename InteriorPRG,
typename ExteriorPRG, typename InputT, typename ...Ps>
struct assemble_key<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, false,
true, InteriorPRG, ExteriorPRG, InputT, Ps...>
{
using type = std::conditional_t<
(CmpDepth > 0),
dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps...,
dpf::cmp_channel_tag<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf>,
dpf::extractable>,
dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps..., dpf::extractable>>;
};
template <std::size_t CmpDepth, std::size_t CmpOutBits, bool CmpWild,
std::size_t CmpBlock, bool CmpIdcf, typename InteriorPRG,
typename ExteriorPRG, typename InputT, typename ...Ps>
struct assemble_key<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, true,
true, InteriorPRG, ExteriorPRG, InputT, Ps...>
{
using type = std::conditional_t<
(CmpDepth > 0),
dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps...,
dpf::cmp_channel_tag<CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf>,
dpf::verifiable, dpf::extractable>,
dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, Ps...,
dpf::verifiable, dpf::extractable>>;
}; };
template <typename InteriorPRG, typename ExteriorPRG, typename InputT, template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, std::size_t CmpDepth, std::size_t CmpOutBits = 0, typename PlacedTuple, std::size_t CmpDepth, std::size_t CmpOutBits = 0,
bool CmpWild = false, std::size_t CmpBlock = 0, bool CmpIdcf = false> bool CmpWild = false, std::size_t CmpBlock = 0, bool CmpIdcf = false,
bool IsVerifiable = false, bool IsExtractable = false>
struct incr_dpf_key_of; struct incr_dpf_key_of;
template <typename InteriorPRG, typename ExteriorPRG, typename InputT, template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename ...Ps, std::size_t CmpDepth, std::size_t CmpOutBits, typename ...Ps, std::size_t CmpDepth, std::size_t CmpOutBits,
bool CmpWild, std::size_t CmpBlock, bool CmpIdcf> bool CmpWild, std::size_t CmpBlock, bool CmpIdcf,
bool IsVerifiable, bool IsExtractable>
struct incr_dpf_key_of<InteriorPRG, ExteriorPRG, InputT, std::tuple<Ps...>, struct incr_dpf_key_of<InteriorPRG, ExteriorPRG, InputT, std::tuple<Ps...>,
CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf> CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf, IsVerifiable, IsExtractable>
{ {
using type = typename assemble_key<CmpDepth, CmpOutBits, CmpWild, CmpBlock, using type = typename assemble_key<CmpDepth, CmpOutBits, CmpWild, CmpBlock,
CmpIdcf, InteriorPRG, ExteriorPRG, InputT, Ps...>::type; CmpIdcf, IsVerifiable, IsExtractable, InteriorPRG, ExteriorPRG, InputT,
Ps...>::type;
}; };
template <typename InteriorPRG, typename ExteriorPRG, typename InputT, template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, std::size_t CmpDepth, std::size_t CmpOutBits = 0, typename PlacedTuple, std::size_t CmpDepth, std::size_t CmpOutBits = 0,
bool CmpWild = false, std::size_t CmpBlock = 0, bool CmpIdcf = false> bool CmpWild = false, std::size_t CmpBlock = 0, bool CmpIdcf = false,
bool IsVerifiable = false, bool IsExtractable = false>
using incr_dpf_key_of_t = typename incr_dpf_key_of<InteriorPRG, ExteriorPRG, using incr_dpf_key_of_t = typename incr_dpf_key_of<InteriorPRG, ExteriorPRG,
InputT, PlacedTuple, CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf>::type; InputT, PlacedTuple, CmpDepth, CmpOutBits, CmpWild, CmpBlock, CmpIdcf,
IsVerifiable, IsExtractable>::type;
} // namespace incr } // namespace incr
} // namespace detail } // namespace detail
@ -1167,10 +1470,12 @@ auto make_dpf_impl(dpfargs<InputT, OutputT, OutputTs...> args, root_sampler_t<In
utils::flip_msb_if_signed_integral(x); utils::flip_msb_if_signed_integral(x);
const interior_node root[2] = { using tree = dpf::tree_traits<InteriorPRG>;
dpf::unset_lo_bit(root_sampler()), HEDLEY_PRAGMA(GCC diagnostic push)
dpf::set_lo_bit(root_sampler()) HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
}; interior_node root[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(root, root_sampler);
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
@ -1179,32 +1484,29 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
correction_advice_array correction_advice; correction_advice_array correction_advice;
interior_node parent[2] = { root[0], root[1] }; interior_node parent[2] = { root[0], root[1] };
bool advice[2];
for (std::size_t level = 0; level < depth; ++level, mask >>= 1) for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{ {
bool bit = !!(mask & x); const bool bit = !!(mask & x);
const bool is_last = tree::is_last_level(level, depth);
const bool ctrl0 = static_cast<bool>(dpf::get_lo_bit(parent[0]));
const bool ctrl1 = static_cast<bool>(dpf::get_lo_bit(parent[1]));
advice[0] = dpf::get_lo_bit_and_clear_lo_2bits(parent[0]); const auto child0 = tree::expand(parent[0], is_last);
advice[1] = dpf::get_lo_bit_and_clear_lo_2bits(parent[1]); const auto child1 = tree::expand(parent[1], is_last);
auto child0 = InteriorPRG::eval01(parent[0]); interior_node cw{};
auto child1 = InteriorPRG::eval01(parent[1]); psnip_uint8_t advice = 0;
interior_node child[2] = { tree::make_cw(cw, advice, child0, child1, parent[0], parent[1], bit,
child0[0] ^ child1[0], is_last);
child0[1] ^ child1[1]
};
bool t[2] = { parent[0] = tree::advance(parent[0], child0, cw, advice, bit, ctrl0,
static_cast<bool>(dpf::get_lo_bit(child[0]) ^ !bit), is_last);
static_cast<bool>(dpf::get_lo_bit(child[1]) ^ bit) parent[1] = tree::advance(parent[1], child1, cw, advice, bit, ctrl1,
}; is_last);
auto cw = dpf::set_lo_bit(child[!bit], t[bit]);
parent[0] = dpf::xor_if(child0[bit], cw, advice[0]);
parent[1] = dpf::xor_if(child1[bit], cw, advice[1]);
correction_words[level] = child[!bit]; correction_words[level] = cw;
correction_advice[level] = static_cast<psnip_uint8_t>(t[1] << 1) | t[0]; correction_advice[level] = advice;
} }
bool sign0 = dpf::get_lo_bit(parent[0]); bool sign0 = dpf::get_lo_bit(parent[0]);

View file

@ -34,6 +34,7 @@ namespace utils
/// @brief Emplaces a `dpf::dpf_key` object into the specified, pre-allocated memory. /// @brief Emplaces a `dpf::dpf_key` object into the specified, pre-allocated memory.
/// @tparam DpfKey The concrete specialization of `dpf::dpf_key` to construct. /// @tparam DpfKey The concrete specialization of `dpf::dpf_key` to construct.
/// @tparam T value type
/// @param storage Reference to the container where the `dpf::dpf_key` object will be emplaced. /// @param storage Reference to the container where the `dpf::dpf_key` object will be emplaced.
/// @param root The root node used by the `dpf::dpf_key`. /// @param root The root node used by the `dpf::dpf_key`.
/// @param correction_words Correction words array for the `dpf::dpf_key`. /// @param correction_words Correction words array for the `dpf::dpf_key`.
@ -59,6 +60,14 @@ struct dpf_emplacer
using input_type = typename DpfKey::input_type; using input_type = typename DpfKey::input_type;
/// @brief Generic version is intentionally left undefined. /// @brief Generic version is intentionally left undefined.
/// @param storage the `storage`
/// @param root the root seed
/// @param correction_words the `correction_words`
/// @param correction_advice the advice bit on the correction word
/// @param leaves the leaf values
/// @param beavers the `beavers`
/// @param offset_share the `offset_share`
/// @return Generic version is intentionally left undefined
static auto emplace(T & storage, static auto emplace(T & storage,
const interior_node & root, const interior_node & root,
const correction_words_array & correction_words, const correction_words_array & correction_words,
@ -69,6 +78,7 @@ struct dpf_emplacer
}; };
/// @brief Specialization for `std::unique_ptr`. /// @brief Specialization for `std::unique_ptr`.
/// @tparam DpfKey DPF key type
template <typename DpfKey> template <typename DpfKey>
struct dpf_emplacer<DpfKey, std::unique_ptr<DpfKey>> struct dpf_emplacer<DpfKey, std::unique_ptr<DpfKey>>
{ {

View file

@ -1,6 +1,5 @@
/// @file dpf/eval_common.hpp /// @file dpf/eval_common.hpp
/// @brief /// @brief Shared evaluation types, party tags, and output cursors.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -23,11 +22,13 @@
namespace dpf namespace dpf
{ {
/// Sentinel: `dpf_output` converts to a bare `OutputT` (no party tag). /// @brief Sentinel: `dpf_output` converts to a bare `OutputT` (no party tag).
inline constexpr std::size_t no_party = std::numeric_limits<std::size_t>::max(); inline constexpr std::size_t no_party = std::numeric_limits<std::size_t>::max();
/// Eval result type for a leaf output of `KeyT`: subtractive share when the /// @brief Eval result type for a leaf output of `KeyT`: subtractive share when the
/// key is party-tagged, otherwise the concrete output type. /// key is party-tagged, otherwise the concrete output type.
/// @tparam KeyT key type
/// @tparam OutputT output type
template <typename KeyT, typename OutputT, bool = is_party_key_v<KeyT>> template <typename KeyT, typename OutputT, bool = is_party_key_v<KeyT>>
struct eval_leaf_result struct eval_leaf_result
{ {
@ -41,8 +42,10 @@ struct eval_leaf_result<KeyT, OutputT, true>
template <typename KeyT, typename OutputT> template <typename KeyT, typename OutputT>
using eval_leaf_result_t = typename eval_leaf_result<KeyT, OutputT>::type; using eval_leaf_result_t = typename eval_leaf_result<KeyT, OutputT>::type;
/// Eval result type for a comparison output of `KeyT`: additive share when /// @brief Eval result type for a comparison output of `KeyT`: additive share when
/// the key is party-tagged, otherwise `Beta`. /// the key is party-tagged, otherwise `Beta`.
/// @tparam KeyT key type
/// @tparam Beta payload type
template <typename KeyT, typename Beta, bool = is_party_key_v<KeyT>> template <typename KeyT, typename Beta, bool = is_party_key_v<KeyT>>
struct eval_cmp_result struct eval_cmp_result
{ {
@ -112,9 +115,14 @@ struct alignas(utils::max_align_v) dpf_output
friend auto make_dpf_output(const Node & node, Input x); friend auto make_dpf_output(const Node & node, Input x);
}; };
/// Copy one packed leaf into a buffer whose element type may differ /// @brief Copy one packed leaf into a buffer whose element type may differ
/// from `LeafT` (bit arrays store `word_type`, not the exterior node). /// from `LeafT` (bit arrays store `word_type`, not the exterior node).
/// Byte destination keeps the store free of strict-aliasing UB. /// @details Byte destination keeps the store free of strict-aliasing UB.
/// @tparam LeafT leaf type
/// @tparam Buffer output buffer type
/// @param buf the output buffer
/// @param index the index
/// @param leaf the leaf value
template <typename LeafT, typename Buffer> template <typename LeafT, typename Buffer>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -144,8 +152,16 @@ auto make_dpf_output(const Node & node, Input x)
offset_within_block<concrete_type_t<Output>, Node>(x)}; offset_within_block<concrete_type_t<Output>, Node>(x)};
} }
/// Wrap a raw leaf node into a party-tagged `dpf_output` when `KeyT` is a /// @brief Wrap a raw leaf node into a party-tagged `dpf_output` when `KeyT` is a
/// `party_key`, otherwise a bare `dpf_output`. /// `party_key`, otherwise a bare `dpf_output`.
/// @tparam KeyT key type
/// @tparam Output output
/// @tparam Input input domain type
/// @tparam Node node
/// @param node the GGM node
/// @param x the `x`
/// @return Wrap a raw leaf node into a party-tagged `dpf_output` when `KeyT` is a `party_key`,
/// otherwise a bare `dpf_output`
template <typename KeyT, typename Output, typename Input, typename Node> template <typename KeyT, typename Output, typename Input, typename Node>
auto make_eval_dpf_output(const Node & node, Input x) auto make_eval_dpf_output(const Node & node, Input x)
{ {
@ -155,8 +171,12 @@ auto make_eval_dpf_output(const Node & node, Input x)
return make_dpf_output<Output>(node, x); return make_dpf_output<Output>(node, x);
} }
/// Wrap a raw comparison `Beta` value as an additive share when `KeyT` is a /// @brief Wrap a raw comparison `Beta` value as an additive share when `KeyT` is a
/// `party_key`. /// `party_key`.
/// @tparam KeyT key type
/// @tparam Beta payload type
/// @param raw the underlying integer
/// @return Wrap a raw comparison `Beta` value as an additive share when `KeyT` is a `party_key`
template <typename KeyT, typename Beta> template <typename KeyT, typename Beta>
HEDLEY_NO_THROW HEDLEY_NO_THROW
auto make_eval_cmp_result(Beta raw) noexcept auto make_eval_cmp_result(Beta raw) noexcept

View file

@ -137,7 +137,13 @@ auto eval_full(const DpfKey & dpf,
return std::make_pair(std::move(outbufs), std::move(iterable)); return std::make_pair(std::move(outbufs), std::move(iterable));
} }
/// Evaluate the whole domain, allocating a basic full memoizer and a buffer. /// @brief Evaluate the whole domain, allocating a basic full memoizer and a buffer.
/// @tparam I output index
/// @tparam Is is
/// @tparam DpfKey DPF key type
/// @tparam DpfKey DPF key type
/// @param dpf the DPF key
/// @return the evaluation result
template <std::size_t I = 0, template <std::size_t I = 0,
std::size_t ...Is, std::size_t ...Is,
typename DpfKey, typename DpfKey,

View file

@ -107,7 +107,6 @@ struct ip_accum
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
if constexpr (simd64 && std::is_same_v<LeafT, simde__m128i>) if constexpr (simd64 && std::is_same_v<LeafT, simde__m128i>)
HEDLEY_PRAGMA(GCC diagnostic pop)
{ {
if (HEDLEY_LIKELY(opl == 2)) if (HEDLEY_LIKELY(opl == 2))
{ {
@ -126,6 +125,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
return; return;
} }
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
for (std::size_t p = 0; p < opl; ++p) for (std::size_t p = 0; p < opl; ++p)
{ {
@ -208,7 +208,10 @@ void eval_inner_product_exterior(const DpfKey & dpf, IntegralT from_node,
using node_type = typename DpfKey::exterior_node; using node_type = typename DpfKey::exterior_node;
using outputs_tuple = typename DpfKey::concrete_outputs_tuple; using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using range = ip_prg_range<node_type, outputs_tuple, Is...>; using range = ip_prg_range<node_type, outputs_tuple, Is...>;
HEDLEY_PRAGMA(GCC diagnostic pop)
constexpr std::size_t opl = DpfKey::outputs_per_leaf; constexpr std::size_t opl = DpfKey::outputs_per_leaf;
std::size_t nodes_in_interval = static_cast<std::size_t>(to_node - from_node); std::size_t nodes_in_interval = static_cast<std::size_t>(to_node - from_node);
@ -397,11 +400,18 @@ auto eval_inner_product_impl(const DpfKey & dpf, InputT from, InputT to,
} // namespace internal } // namespace internal
/// Expand the interior tree for `[from, to]`. A wrapping interval is left /// @brief Expand the interior tree for `[from, to]`. A wrapping interval is left
/// cold: the memoizer holds one half, and walking the first half of the later /// cold: the memoizer holds one half, and walking the first half of the later
/// inner product would clobber a cached second half. Safe to call before the /// inner product would clobber a cached second half. Safe to call before the
/// weight vector exists; a subsequent inner-product on the same memoizer /// weight vector exists; a subsequent inner-product on the same memoizer
/// skips the interior AES when the interval did not wrap. /// skips the interior AES when the interval did not wrap.
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @tparam IntervalMemoizer interval memoizer type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the `to`
/// @param memoizer the memoizer built for this key
template <typename DpfKey, template <typename DpfKey,
typename InputT, typename InputT,
typename IntervalMemoizer> typename IntervalMemoizer>
@ -425,11 +435,24 @@ void eval_prepare_full(const DpfKey & dpf, IntervalMemoizer && memoizer)
memoizer); memoizer);
} }
/// `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR). /// @brief `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR).
/// `w[j]` is the weight for the `j`-th output in the interval, matching /// @details `w[j]` is the weight for the `j`-th output in the interval, matching
/// `eval_interval`'s destination layout. Multiple `Is` take a tuple of /// `eval_interval`'s destination layout. Multiple `Is` take a tuple of
/// weight ranges and return a tuple of accumulators; a single `I` takes /// weight ranges and return a tuple of accumulators; a single `I` takes
/// one range and returns one accumulator. /// one range and returns one accumulator.
/// @tparam I output index
/// @tparam Is is
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @tparam Weights weights
/// @tparam IntervalMemoizer interval memoizer type
/// @tparam DpfKey DPF key type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the `to`
/// @param weights the weights
/// @param memoizer the memoizer built for this key
/// @return `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR)
template <std::size_t I = 0, template <std::size_t I = 0,
std::size_t ...Is, std::size_t ...Is,
typename DpfKey, typename DpfKey,

View file

@ -67,6 +67,8 @@ inline auto eval_interval_interior(const DpfKey & dpf, IntegralT from_node,
dpf.correction_word(level_index-1, 0), dpf.correction_word(level_index-1, 0),
dpf.correction_word(level_index-1, 1) dpf.correction_word(level_index-1, 1)
}; };
const bool is_last = dpf_type::tree::is_last_level(level_index - 1,
dpf.depth);
auto *prev = memoizer[level_index-1]; auto *prev = memoizer[level_index-1];
auto *curr = memoizer[level_index]; auto *curr = memoizer[level_index];
@ -74,7 +76,7 @@ inline auto eval_interval_interior(const DpfKey & dpf, IntegralT from_node,
// process node which only requires a right traversal // process node which only requires a right traversal
if (from_offset == true) if (from_offset == true)
{ {
curr[i++] = dpf_type::traverse_interior(prev[j++], cw[1], 1); curr[i++] = dpf_type::traverse_interior(prev[j++], cw[1], 1, is_last);
} }
// process all nodes which require both a left traversal and a right traversal // process all nodes which require both a left traversal and a right traversal
const std::size_t both_end = nodes_at_level - to_offset; const std::size_t both_end = nodes_at_level - to_offset;
@ -88,7 +90,8 @@ inline auto eval_interval_interior(const DpfKey & dpf, IntegralT from_node,
{ {
parents[t] = prev[j + t]; parents[t] = prev[j + t];
} }
dpf_type::traverse_interior01_x4(parents, cw[0], cw[1], left, right); dpf_type::traverse_interior01_x4(parents, cw[0], cw[1], left, right,
is_last);
DPF_UNROLL_LOOP DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t) for (std::size_t t = 0; t < 4; ++t)
{ {
@ -102,14 +105,15 @@ inline auto eval_interval_interior(const DpfKey & dpf, IntegralT from_node,
for (; i < both_end;) for (; i < both_end;)
{ {
auto cur_node = prev[j++]; auto cur_node = prev[j++];
auto kids = dpf_type::traverse_interior01(cur_node, cw[0], cw[1]); auto kids = dpf_type::traverse_interior01(cur_node, cw[0], cw[1],
is_last);
curr[i++] = kids[0]; curr[i++] = kids[0];
curr[i++] = kids[1]; curr[i++] = kids[1];
} }
// process node which only requires a left traversal // process node which only requires a left traversal
if (to_offset == true) if (to_offset == true)
{ {
curr[i] = dpf_type::traverse_interior(prev[j], cw[0], 0); curr[i] = dpf_type::traverse_interior(prev[j], cw[0], 0, is_last);
} }
} }
} }
@ -135,6 +139,7 @@ inline auto eval_interval_exterior(const DpfKey & dpf, IntegralT from_node,
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto cw = std::get<I>(dpf.leaf_nodes).get(); auto cw = std::get<I>(dpf.leaf_nodes).get();
HEDLEY_PRAGMA(GCC diagnostic pop)
auto *nodes = memoizer[dpf_type::depth]; auto *nodes = memoizer[dpf_type::depth];
DPF_UNROLL_LOOP DPF_UNROLL_LOOP
for (std::size_t j = 0, k = start; j < nodes_in_interval; ++j, ++k) for (std::size_t j = 0, k = start; j < nodes_in_interval; ++j, ++k)
@ -151,7 +156,6 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
sizeof(output_type) * dpf_type::outputs_per_leaf); sizeof(output_type) * dpf_type::outputs_per_leaf);
} }
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <std::size_t I, template <std::size_t I,
@ -175,9 +179,23 @@ void store_interval_leaf(OutputBuffer && outbuf, std::size_t k, const LeafT & le
} }
} }
/// One pass over the leaf-level interior nodes. When the selected output /// @brief One pass over the leaf-level interior nodes. When the selected output
/// indices occupy a contiguous PRG-position range, a single batched /// indices occupy a contiguous PRG-position range, a single batched
/// `ExteriorPRG::eval` produces every output's leaf mask. /// `ExteriorPRG::eval` produces every output's leaf mask.
/// @tparam Is is
/// @tparam DpfKey DPF key type
/// @tparam OutputBuffers tuple of output buffers
/// @tparam IntervalMemoizer interval memoizer type
/// @tparam IntegralT integral type
/// @tparam IIs iis
/// @param dpf the DPF key
/// @param from_node the `from_node`
/// @param to_node the `to_node`
/// @param outbufs the named output buffers
/// @param memoizer the memoizer built for this key
/// @param IIs the `IIs`
/// @param start the start of the range
/// @throws std::runtime_error if `to_node<from_node`
template <std::size_t ...Is, template <std::size_t ...Is,
typename DpfKey, typename DpfKey,
typename OutputBuffers, typename OutputBuffers,
@ -194,7 +212,10 @@ inline void eval_interval_exterior_fused(const DpfKey & dpf, IntegralT from_node
using node_type = typename DpfKey::exterior_node; using node_type = typename DpfKey::exterior_node;
using outputs_tuple = typename DpfKey::concrete_outputs_tuple; using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using range = leaf_prg_range<node_type, outputs_tuple, Is...>; using range = leaf_prg_range<node_type, outputs_tuple, Is...>;
HEDLEY_PRAGMA(GCC diagnostic pop)
std::size_t nodes_in_interval = static_cast<std::size_t>(to_node - from_node); std::size_t nodes_in_interval = static_cast<std::size_t>(to_node - from_node);
auto *nodes = memoizer[DpfKey::depth]; auto *nodes = memoizer[DpfKey::depth];
@ -312,7 +333,10 @@ void eval_interval_exterior_all(const DpfKey & dpf, IntegralT from_node,
{ {
using node_type = typename DpfKey::exterior_node; using node_type = typename DpfKey::exterior_node;
using outputs_tuple = typename DpfKey::concrete_outputs_tuple; using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using range = leaf_prg_range<node_type, outputs_tuple, Is...>; using range = leaf_prg_range<node_type, outputs_tuple, Is...>;
HEDLEY_PRAGMA(GCC diagnostic pop)
if constexpr (range::is_contiguous) if constexpr (range::is_contiguous)
{ {
eval_interval_exterior_fused<Is...>(dpf, from_node, to_node, outbufs, eval_interval_exterior_fused<Is...>(dpf, from_node, to_node, outbufs,
@ -397,10 +421,26 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
} // namespace internal } // namespace internal
/// Write outputs `I, Is...` for `[from, to]` into `outbufs`. /// @name Closed-interval evaluation
/// @param outbufs Named buffer, or a tuple of buffers when several outputs /// @tparam I output index
/// @tparam Is the remaining output indices
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the inclusive end of the range
/// @{
/// @brief Write outputs `I, Is...` for `[from, to]` into `outbufs`.
/// @tparam OutputBuffers tuple of output buffers
/// @tparam IntervalMemoizer interval memoizer type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the inclusive end of the range
/// @param outbufs named buffer, or a tuple of buffers when several outputs
/// are selected. Must outlive the returned iterable. /// are selected. Must outlive the returned iterable.
/// @param memoizer Workspace sized for at least this interval. /// @param memoizer workspace sized for at least this interval
/// @return an iterable over the written outputs
template <std::size_t I = 0, template <std::size_t I = 0,
std::size_t ...Is, std::size_t ...Is,
typename DpfKey, typename DpfKey,
@ -417,7 +457,13 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
return internal::eval_interval<I, Is...>(dpf, dpf.offset_x(from), dpf.offset_x(to), outbufs, memoizer, std::make_index_sequence<1+sizeof...(Is)>()); return internal::eval_interval<I, Is...>(dpf, dpf.offset_x(from), dpf.offset_x(to), outbufs, memoizer, std::make_index_sequence<1+sizeof...(Is)>());
} }
/// Evaluate `[from, to]` into `outbufs`, allocating a basic interval memoizer. /// @brief Evaluate `[from, to]` into `outbufs`, allocating a basic interval memoizer.
/// @tparam OutputBuffers tuple of output buffers
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the inclusive end of the range
/// @param outbufs the named output buffers
/// @return an iterable over the written outputs
template <std::size_t I = 0, template <std::size_t I = 0,
std::size_t ...Is, std::size_t ...Is,
typename DpfKey, typename DpfKey,
@ -435,7 +481,12 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
dpf::make_basic_interval_memoizer<DpfKey>(from, to)); dpf::make_basic_interval_memoizer<DpfKey>(from, to));
} }
/// Evaluate `[from, to]` with a caller-supplied memoizer. /// @brief Evaluate `[from, to]` with a caller-supplied memoizer.
/// @tparam IntervalMemoizer interval memoizer type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the inclusive end of the range
/// @param memoizer the memoizer built for this key
/// @return `std::pair` of a new buffer (or tuple of buffers) and an iterable /// @return `std::pair` of a new buffer (or tuple of buffers) and an iterable
/// into that buffer. /// into that buffer.
template <std::size_t I = 0, template <std::size_t I = 0,
@ -462,7 +513,9 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
return std::make_pair(std::move(outbufs), std::move(iterable)); return std::make_pair(std::move(outbufs), std::move(iterable));
} }
/// Evaluate `[from, to]`, allocating a basic interval memoizer and a buffer. /// @brief Evaluate `[from, to]`, allocating a basic interval memoizer and a buffer.
/// @return `std::pair` of a new buffer (or tuple of buffers) and an iterable
/// into that buffer.
template <std::size_t I = 0, template <std::size_t I = 0,
std::size_t ...Is, std::size_t ...Is,
typename DpfKey, typename DpfKey,
@ -475,6 +528,8 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to)
dpf::make_basic_interval_memoizer<DpfKey>(from, to)); dpf::make_basic_interval_memoizer<DpfKey>(from, to));
} }
/// @}
} // namespace dpf } // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_EVAL_INTERVAL_HPP__ #endif // LIBDPF_INCLUDE_DPF_EVAL_INTERVAL_HPP__

View file

@ -5,6 +5,7 @@
/// `eval_point<I0, I1, ...>` returns a tuple of shares. /// `eval_point<I0, I1, ...>` returns a tuple of shares.
/// Pass a `basic_path_memoizer` lvalue to resume a previous path. /// Pass a `basic_path_memoizer` lvalue to resume a previous path.
/// An unassigned wildcard output throws `std::runtime_error`. /// An unassigned wildcard output throws `std::runtime_error`.
/// `eval_point(key, x, dpf::prove(π))` folds a VDPF proof token.
/// @snippet evaluation/eval_point.cpp eval-point /// @snippet evaluation/eval_point.cpp eval-point
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @author Christopher Jiang <christopher.jiang@ucalgary.ca> /// @author Christopher Jiang <christopher.jiang@ucalgary.ca>
@ -25,6 +26,7 @@
#include "dpf/eval_common.hpp" #include "dpf/eval_common.hpp"
#include "dpf/eval_target.hpp" #include "dpf/eval_target.hpp"
#include "dpf/path_memoizer.hpp" #include "dpf/path_memoizer.hpp"
#include "dpf/verifiable.hpp"
namespace dpf namespace dpf
{ {
@ -35,7 +37,8 @@ namespace internal
template <typename DpfKey, template <typename DpfKey,
typename InputT, typename InputT,
typename PathMemoizer> typename PathMemoizer>
inline auto eval_point_interior(const DpfKey & dpf, InputT && x, PathMemoizer && path) inline auto eval_point_interior(const DpfKey & dpf, InputT && x, PathMemoizer && path,
proof_token * pi = nullptr)
{ {
using dpf_type = DpfKey; using dpf_type = DpfKey;
@ -47,7 +50,21 @@ inline auto eval_point_interior(const DpfKey & dpf, InputT && x, PathMemoizer &&
{ {
bool bit = !!(mask & x); bool bit = !!(mask & x);
auto cw = dpf.correction_word(level_index-1, bit); auto cw = dpf.correction_word(level_index-1, bit);
path[level_index] = dpf_type::traverse_interior(path[level_index-1], cw, bit); const bool is_last = dpf_type::tree::is_last_level(level_index - 1,
dpf.depth);
path[level_index] = dpf_type::traverse_interior(path[level_index-1],
cw, bit, is_last);
if constexpr (dpf_type::is_verifiable)
{
if (pi != nullptr)
{
const auto x_bits = static_cast<psnip_uint64_t>(
utils::to_integral_type<std::decay_t<InputT>>{}(x)
>> (utils::bitlength_of_v<std::decay_t<InputT>> - level_index));
detail::vdpf::fold_node(*pi, level_index - 1, x_bits,
path[level_index], dpf.correction_seeds()[level_index - 1]);
}
}
} }
detail::path_note_filled_to(path, dpf.depth); detail::path_note_filled_to(path, dpf.depth);
} }
@ -68,19 +85,17 @@ template <std::size_t I,
typename InputT, typename InputT,
typename PathMemoizer> typename PathMemoizer>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path) auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path,
proof_token * pi = nullptr)
{ {
utils::flip_msb_if_signed_integral(x); utils::flip_msb_if_signed_integral(x);
internal::eval_point_interior(dpf, x, path); internal::eval_point_interior(dpf, x, path, pi);
return internal::eval_point_exterior<I>(dpf, path); return internal::eval_point_exterior<I>(dpf, path);
} }
} // namespace internal } // namespace internal
/// Evaluate output `I` at `x`. /// Evaluate output `I` at `x`.
/// @param path Mutable path memoizer. The default is a fresh
/// nonmemoizing workspace for this call.
/// @return Handle whose `operator*` is the party's share.
template <std::size_t I = 0, template <std::size_t I = 0,
typename DpfKey, typename DpfKey,
typename InputT, typename InputT,
@ -97,8 +112,28 @@ auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path = PathMemo
internal::eval_point<I>(dpf, tx, path), tx); internal::eval_point<I>(dpf, tx, path), tx);
} }
/// Evaluate and fold a VDPF proof token for the walked path.
template <std::size_t I = 0,
typename DpfKey,
typename InputT,
typename PathMemoizer = dpf::nonmemoizing_path_memoizer<DpfKey>,
std::enable_if_t<looks_like_dpf_key_v<DpfKey>, bool> = true>
HEDLEY_ALWAYS_INLINE
auto eval_point(const DpfKey & dpf, InputT && x, prove_ref pr,
PathMemoizer && path = PathMemoizer{})
{
static_assert(DpfKey::is_verifiable,
"eval_point(..., prove(π)): key must carry dpf::verifiable");
assert_not_wildcard_output<I>(dpf);
using output_type = typename DpfKey::concrete_output_type<I>;
detail::vdpf::init_proof(pr.token, dpf);
auto tx = dpf.offset_x(x);
return make_eval_dpf_output<DpfKey, output_type>(
internal::eval_point<I>(dpf, tx, path, &pr.token), tx);
}
/// Evaluate several outputs at `x`. /// Evaluate several outputs at `x`.
/// @return Tuple of shares, already dereferenced.
template <std::size_t I0, template <std::size_t I0,
std::size_t I1, std::size_t I1,
std::size_t ...Is, std::size_t ...Is,
@ -115,6 +150,37 @@ auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path = PathMemo
*eval_point<Is>(dpf, x, path)...); *eval_point<Is>(dpf, x, path)...);
} }
/// Fold every point in `[from, to]` into `pi` (caller must `init_proof` first,
/// or pass a fresh token via `prove_interval` below).
template <typename KeyT, typename InputT>
void prove_fold_interval(const KeyT & key, InputT from, InputT to,
proof_token & pi)
{
static_assert(KeyT::is_verifiable,
"prove_fold_interval: key must carry dpf::verifiable");
using input_type = typename KeyT::input_type;
auto cur = static_cast<input_type>(from);
const auto last = static_cast<input_type>(to);
for (;;)
{
nonmemoizing_path_memoizer<KeyT> path{};
auto tx = key.offset_x(cur);
utils::flip_msb_if_signed_integral(tx);
internal::eval_point_interior(key, tx, path, &pi);
if (cur == last)
break;
++cur;
}
}
/// Initialise `pr.token` and fold `[from, to]`.
template <typename KeyT, typename InputT>
void prove_interval(const KeyT & key, InputT from, InputT to, prove_ref pr)
{
detail::vdpf::init_proof(pr.token, key);
prove_fold_interval(key, from, to, pr.token);
}
} // namespace dpf } // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_EVAL_POINT_HPP__ #endif // LIBDPF_INCLUDE_DPF_EVAL_POINT_HPP__

View file

@ -148,8 +148,18 @@ inline auto eval_sequence(const DpfKey & dpf, ForwardIterator begin, ForwardIter
} }
} }
/// Evaluate the sorted range `[begin, end)`, allocating a buffer. /// @brief Evaluate the sorted range `[begin, end)`, allocating a buffer.
/// @tparam I output index
/// @tparam Is is
/// @tparam DpfKey DPF key type
/// @tparam ForwardIterator forward iterator type
/// @tparam ReturnType return type
/// @tparam DpfKey DPF key type
/// @tparam ReturnType return type
/// @param return_type `return_entire_node_tag_{}` or `return_output_only_tag_{}`. /// @param return_type `return_entire_node_tag_{}` or `return_output_only_tag_{}`.
/// @param dpf the DPF key
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @return Pair of buffer (or tuple of buffers) and an iterable in list order. /// @return Pair of buffer (or tuple of buffers) and an iterable in list order.
template <std::size_t I = 0, template <std::size_t I = 0,
std::size_t ...Is, std::size_t ...Is,
@ -188,8 +198,8 @@ inline auto eval_sequence_breadth_first(const DpfKey & dpf, ForwardIterator begi
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using allocator = aligned_allocator<typename DpfKey::interior_node>; using allocator = aligned_allocator<typename DpfKey::interior_node>;
using unique_ptr = typename allocator::unique_ptr;
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
using unique_ptr = typename allocator::unique_ptr;
allocator alloc = allocator{}; allocator alloc = allocator{};
if (HEDLEY_UNLIKELY(!std::is_sorted(begin, end))) if (HEDLEY_UNLIKELY(!std::is_sorted(begin, end)))
@ -219,6 +229,8 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
dpf.correction_word(level_index-1, 0), dpf.correction_word(level_index-1, 0),
dpf.correction_word(level_index-1, 1) dpf.correction_word(level_index-1, 1)
}; };
const bool is_last = dpf_type::tree::is_last_level(level_index - 1,
dpf_type::depth);
// `lower` and `upper` are always adjacent elements of `splits` with `lower` < `upper` // `lower` and `upper` are always adjacent elements of `splits` with `lower` < `upper`
// [lower, upper) = "block" // [lower, upper) = "block"
for (auto upper = std::begin(splits), lower = upper++; upper != std::end(splits); lower = upper++) for (auto upper = std::begin(splits), lower = upper++; upper != std::end(splits); lower = upper++)
@ -228,16 +240,16 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
[&flip](auto a, auto b){ return static_cast<bool>(a&b) ^ flip; }); [&flip](auto a, auto b){ return static_cast<bool>(a&b) ^ flip; });
if (it == *lower) // right only since first element in "block" requires right traversal if (it == *lower) // right only since first element in "block" requires right traversal
{ {
memo[curhalf*nodes_in_sequence + i++] = dpf_type::traverse_interior(memo[!curhalf*nodes_in_sequence + j++], cw[1], 1); memo[curhalf*nodes_in_sequence + i++] = dpf_type::traverse_interior(memo[!curhalf*nodes_in_sequence + j++], cw[1], 1, is_last);
} }
else if (it == *upper) // left only since no element in "block" requires right traversal else if (it == *upper) // left only since no element in "block" requires right traversal
{ {
memo[curhalf*nodes_in_sequence + i++] = dpf_type::traverse_interior(memo[!curhalf*nodes_in_sequence + j++], cw[0], 0); memo[curhalf*nodes_in_sequence + i++] = dpf_type::traverse_interior(memo[!curhalf*nodes_in_sequence + j++], cw[0], 0, is_last);
} }
else // both ways since some (non-lower) element within "block" requires right traversal else // both ways since some (non-lower) element within "block" requires right traversal
{ {
auto cur_node = memo[!curhalf*nodes_in_sequence + j++]; auto cur_node = memo[!curhalf*nodes_in_sequence + j++];
auto kids = dpf_type::traverse_interior01(cur_node, cw[0], cw[1]); auto kids = dpf_type::traverse_interior01(cur_node, cw[0], cw[1], is_last);
memo[curhalf*nodes_in_sequence + i++] = kids[0]; memo[curhalf*nodes_in_sequence + i++] = kids[0];
memo[curhalf*nodes_in_sequence + i++] = kids[1]; memo[curhalf*nodes_in_sequence + i++] = kids[1];
splits.insert(upper, it); splits.insert(upper, it);
@ -259,6 +271,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto cw = dpf.template leaf<I>(); auto cw = dpf.template leaf<I>();
HEDLEY_PRAGMA(GCC diagnostic pop)
auto buf = memo.get(); auto buf = memo.get();
constexpr auto clz = utils::countl_zero_symmetric_difference<input_type>{}; constexpr auto clz = utils::countl_zero_symmetric_difference<input_type>{};
@ -278,7 +291,6 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
} }
prev = curr++; prev = curr++;
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
return subsequence_iterable<DpfKey, decltype(std::begin(outbuf)), ForwardIterator>(std::begin(outbuf), begin, end); return subsequence_iterable<DpfKey, decltype(std::begin(outbuf)), ForwardIterator>(std::begin(outbuf), begin, end);
} }
@ -324,6 +336,8 @@ inline auto eval_sequence_interior(const DpfKey & dpf, const sequence_recipe & r
dpf.correction_word(level_index-1, 0), dpf.correction_word(level_index-1, 0),
dpf.correction_word(level_index-1, 1) dpf.correction_word(level_index-1, 1)
}; };
const bool is_last = dpf_type::tree::is_last_level(level_index - 1,
dpf_type::depth);
auto prevbuf = memoizer[level_index-1]; auto prevbuf = memoizer[level_index-1];
auto currbuf = memoizer[level_index]; auto currbuf = memoizer[level_index];
@ -334,12 +348,12 @@ inline auto eval_sequence_interior(const DpfKey & dpf, const sequence_recipe & r
if (memoizer.traverse_first(recipe_index) == true) if (memoizer.traverse_first(recipe_index) == true)
{ {
bool dir = memoizer.get_direction(0); bool dir = memoizer.get_direction(0);
currbuf[output_index++] = dpf_type::traverse_interior(prevbuf[input_index], cw[dir], dir); currbuf[output_index++] = dpf_type::traverse_interior(prevbuf[input_index], cw[dir], dir, is_last);
} }
if (memoizer.traverse_second(recipe_index) == true) if (memoizer.traverse_second(recipe_index) == true)
{ {
bool dir = memoizer.get_direction(1); bool dir = memoizer.get_direction(1);
currbuf[output_index++] = dpf_type::traverse_interior(prevbuf[input_index], cw[dir], dir); currbuf[output_index++] = dpf_type::traverse_interior(prevbuf[input_index], cw[dir], dir, is_last);
} }
} }
} }
@ -362,6 +376,7 @@ inline auto eval_sequence_exterior_entire_node(const DpfKey & dpf, const sequenc
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto buf = memoizer[dpf.depth]; auto buf = memoizer[dpf.depth];
HEDLEY_PRAGMA(GCC diagnostic pop)
DPF_UNROLL_LOOP DPF_UNROLL_LOOP
for (std::size_t j = 0; j < nodes_in_interval; ++j) for (std::size_t j = 0; j < nodes_in_interval; ++j)
{ {
@ -375,7 +390,6 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::memcpy(&outbuf[j*dpf_type::outputs_per_leaf], &leaf, sizeof(output_type)*dpf_type::outputs_per_leaf); std::memcpy(&outbuf[j*dpf_type::outputs_per_leaf], &leaf, sizeof(output_type)*dpf_type::outputs_per_leaf);
} }
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <std::size_t I, template <std::size_t I,
@ -393,6 +407,7 @@ inline auto eval_sequence_exterior_output_only(const DpfKey & dpf, const sequenc
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto cw = dpf.template leaf<I>(); auto cw = dpf.template leaf<I>();
HEDLEY_PRAGMA(GCC diagnostic pop)
using node_type = typename DpfKey::exterior_node; using node_type = typename DpfKey::exterior_node;
using leaf_node_type = std::tuple_element_t<I, typename DpfKey::leaf_tuple>; using leaf_node_type = std::tuple_element_t<I, typename DpfKey::leaf_tuple>;
auto buf = memoizer[dpf.depth]; auto buf = memoizer[dpf.depth];
@ -417,7 +432,6 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
else else
outbuf[i] = v; outbuf[i] = v;
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <std::size_t ...Is, template <std::size_t ...Is,
@ -454,9 +468,21 @@ auto eval_sequence(const DpfKey & dpf, const sequence_recipe & recipe,
} // namespace internal } // namespace internal
/// Evaluate `recipe` into a named buffer, reusing `memoizer`. /// @brief Evaluate `recipe` into a named buffer, reusing `memoizer`.
/// @tparam I output index
/// @tparam Is is
/// @tparam DpfKey DPF key type
/// @tparam OutputBuffers tuple of output buffers
/// @tparam SequenceMemoizer sequence memoizer
/// @tparam ReturnType return type
/// @tparam SequenceMemoizer sequence memoizer
/// @tparam ReturnType return type
/// @param recipe The same object `memoizer` was constructed from. /// @param recipe The same object `memoizer` was constructed from.
/// @param outbufs Named buffer. The returned iterable refers into it. /// @param outbufs Named buffer. The returned iterable refers into it.
/// @param dpf the DPF key
/// @param memoizer the memoizer built for this key
/// @param return_type `return_entire_node_tag_{}` or `return_output_only_tag_{}`
/// @return the evaluation result
template <std::size_t I = 0, template <std::size_t I = 0,
std::size_t ...Is, std::size_t ...Is,
typename DpfKey, typename DpfKey,

View file

@ -10,14 +10,18 @@
#include <limits> #include <limits>
#include <type_traits> #include <type_traits>
#include "hedley/hedley.h"
namespace dpf namespace dpf
{ {
/// Sentinel: deduce point-slot prefix from the key (`meta[I].prefix`). /// @brief Sentinel: deduce point-slot prefix from the key (`meta[I].prefix`).
inline constexpr std::size_t prefix_deduce = inline constexpr std::size_t prefix_deduce =
std::numeric_limits<std::size_t>::max(); std::numeric_limits<std::size_t>::max();
/// Point-output channel: slot `I`, optional prefix check `N`. /// @brief Point-output channel: slot `I`, optional prefix check `N`.
/// @tparam I output index
/// @tparam N prefix length, or `prefix_deduce`
template <std::size_t I = 0, std::size_t N = prefix_deduce> template <std::size_t I = 0, std::size_t N = prefix_deduce>
struct out_t struct out_t
{ {
@ -29,13 +33,14 @@ struct out_t
template <std::size_t I = 0, std::size_t N = prefix_deduce> template <std::size_t I = 0, std::size_t N = prefix_deduce>
inline constexpr out_t<I, N> out{}; inline constexpr out_t<I, N> out{};
/// Comparison (DCF) channel. /// @brief Comparison (DCF) channel.
struct cmp_t struct cmp_t
{ {
}; };
inline constexpr cmp_t cmp{}; inline constexpr cmp_t cmp{};
/// Prefix of an `idcf` comparison. `L` is the number of leading bits. /// @brief Prefix of an `idcf` comparison. `L` is the number of leading bits.
/// @tparam L number of leading bits
template <std::size_t L> template <std::size_t L>
struct cmp_prefix_t struct cmp_prefix_t
{ {
@ -74,7 +79,7 @@ template <typename T>
inline constexpr bool is_cmp_prefix_target_v = inline constexpr bool is_cmp_prefix_target_v =
is_cmp_prefix_target<std::decay_t<T>>::value; is_cmp_prefix_target<std::decay_t<T>>::value;
/// True for channel tags that must not bind as the key in classic eval_*. /// @brief True for channel tags that must not bind as the key in classic eval_*.
template <typename T> template <typename T>
inline constexpr bool is_eval_channel_tag_v = inline constexpr bool is_eval_channel_tag_v =
is_out_v<T> || is_cmp_target_v<T> || is_cmp_prefix_target_v<T>; is_out_v<T> || is_cmp_target_v<T> || is_cmp_prefix_target_v<T>;
@ -83,12 +88,15 @@ template <typename T, typename = void>
struct looks_like_dpf_key : std::false_type struct looks_like_dpf_key : std::false_type
{ {
}; };
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
template <typename T> template <typename T>
struct looks_like_dpf_key<T, struct looks_like_dpf_key<T,
std::void_t<typename T::input_type, typename T::interior_node>> std::void_t<typename T::input_type, typename T::interior_node>>
: std::true_type : std::true_type
{ {
}; };
HEDLEY_PRAGMA(GCC diagnostic pop)
template <typename T> template <typename T>
inline constexpr bool looks_like_dpf_key_v = inline constexpr bool looks_like_dpf_key_v =
looks_like_dpf_key<std::decay_t<T>>::value; looks_like_dpf_key<std::decay_t<T>>::value;
@ -107,12 +115,13 @@ template <typename T>
inline constexpr bool is_incremental_dpf_key_v = inline constexpr bool is_incremental_dpf_key_v =
is_incremental_dpf_key<std::decay_t<T>>::value; is_incremental_dpf_key<std::decay_t<T>>::value;
/// True only for keys that must use the slot-aware (multi-level / comparison) /// @brief True only for keys that must use the slot-aware (multi-level / comparison)
/// eval path. Every key now carries a `slot_meta` table (so /// eval path. Every key now carries a `slot_meta` table (so
/// `is_incremental_dpf_key_v` is true for all keys), but classic single-level /// `is_incremental_dpf_key_v` is true for all keys), but classic single-level
/// equal-width keys keep using the classic `eval_*` fast paths; they set /// equal-width keys keep using the classic `eval_*` fast paths; they set
/// `is_multilevel == false`. Multi-level (`at<N>`) and comparison keys set it /// `is_multilevel == false`. Multi-level (`at<N>`) and comparison keys set it
/// to true. /// to true.
/// @tparam T value type
template <typename T, typename = void> template <typename T, typename = void>
struct is_multilevel_key : std::false_type struct is_multilevel_key : std::false_type
{ {

View file

@ -410,7 +410,10 @@ auto eval_out_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to,
const bool wraps = utils::interval_wraps(from_i, to_i, N); const bool wraps = utils::interval_wraps(from_i, to_i, N);
const auto segs = utils::split_leaf_nodes(from_node, to_node, to_level, wraps); const auto segs = utils::split_leaf_nodes(from_node, to_node, to_level, wraps);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ml_ip_accum<output_type, exterior_node> acc{}; ml_ip_accum<output_type, exterior_node> acc{};
HEDLEY_PRAGMA(GCC diagnostic pop)
std::size_t start = 0; std::size_t start = 0;
for (std::size_t s = 0; s < segs.n; ++s) for (std::size_t s = 0; s < segs.n; ++s)
{ {
@ -552,7 +555,10 @@ void eval_out_sequence_breadth_first_impl(const KeyT & dpf,
if (begin == end) if (begin == end)
return; return;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using allocator = aligned_allocator<node_type>; using allocator = aligned_allocator<node_type>;
HEDLEY_PRAGMA(GCC diagnostic pop)
allocator alloc{}; allocator alloc{};
const std::size_t nseq = static_cast<std::size_t>(std::distance(begin, end)); const std::size_t nseq = static_cast<std::size_t>(std::distance(begin, end));
auto memo = alloc.allocate_unique_ptr(nseq * 2); auto memo = alloc.allocate_unique_ptr(nseq * 2);
@ -570,6 +576,8 @@ void eval_out_sequence_breadth_first_impl(const KeyT & dpf,
const node_type cw[2] = { const node_type cw[2] = {
dpf.correction_word(level_index - 1, 0), dpf.correction_word(level_index - 1, 0),
dpf.correction_word(level_index - 1, 1)}; dpf.correction_word(level_index - 1, 1)};
const bool is_last = key_type::tree::is_last_level(level_index - 1,
key_type::depth);
const std::size_t cur = static_cast<std::size_t>(curhalf) * nseq; const std::size_t cur = static_cast<std::size_t>(curhalf) * nseq;
const std::size_t prv = static_cast<std::size_t>(!curhalf) * nseq; const std::size_t prv = static_cast<std::size_t>(!curhalf) * nseq;
for (auto upper = std::begin(splits), lower = upper++; for (auto upper = std::begin(splits), lower = upper++;
@ -580,17 +588,17 @@ void eval_out_sequence_breadth_first_impl(const KeyT & dpf,
if (it == *lower) if (it == *lower)
{ {
memo[cur + i++] = key_type::traverse_interior( memo[cur + i++] = key_type::traverse_interior(
memo[prv + j++], cw[1], 1); memo[prv + j++], cw[1], 1, is_last);
} }
else if (it == *upper) else if (it == *upper)
{ {
memo[cur + i++] = key_type::traverse_interior( memo[cur + i++] = key_type::traverse_interior(
memo[prv + j++], cw[0], 0); memo[prv + j++], cw[0], 0, is_last);
} }
else else
{ {
auto kids = key_type::traverse_interior01(memo[prv + j++], auto kids = key_type::traverse_interior01(memo[prv + j++],
cw[0], cw[1]); cw[0], cw[1], is_last);
memo[cur + i++] = kids[0]; memo[cur + i++] = kids[0];
memo[cur + i++] = kids[1]; memo[cur + i++] = kids[1];
splits.insert(upper, it); splits.insert(upper, it);
@ -649,7 +657,17 @@ auto eval_sequence_breadth_first(out_t<I, N>, const KeyT & key,
return buf; return buf;
} }
/// Build a sequence recipe stopped at slot `I`'s tree level (prefix domain). /// @brief Build a sequence recipe stopped at slot `I`'s tree level (prefix domain).
/// @tparam I output index
/// @tparam N width in bits
/// @tparam KeyT key type
/// @tparam ForwardIterator forward iterator type
/// @tparam KeyT key type
/// @param N the `N`
/// @param key the `key`
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @return the constructed object
template <std::size_t I, std::size_t N, typename KeyT, typename ForwardIterator, template <std::size_t I, std::size_t N, typename KeyT, typename ForwardIterator,
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true> std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
auto make_sequence_recipe(out_t<I, N>, const KeyT & key, ForwardIterator begin, auto make_sequence_recipe(out_t<I, N>, const KeyT & key, ForwardIterator begin,

241
include/dpf/fp61.hpp Normal file
View file

@ -0,0 +1,241 @@
/// @file dpf/fp61.hpp
/// @brief Prime field \(\mathbb{F}_{2^{61}-1}\) additive output type.
/// @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_FP61_HPP__
#define LIBDPF_INCLUDE_DPF_FP61_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <ostream>
#include <type_traits>
#include "hedley/hedley.h"
#include "dpf/utils.hpp"
#include "dpf/leaf_arithmetic.hpp"
namespace dpf
{
/// @brief Modulus \(p = 2^{61}-1\). `p` itself reduces to 0.
inline constexpr std::uint64_t fp61_mod = (std::uint64_t{1} << 61) - 1;
/// @brief Additive element of \(\mathbb{F}_{2^{61}-1}\).
class fp61
{
public:
/// @brief Underlying unsigned word. Values are stored already reduced.
using integral_type = std::uint64_t;
static constexpr std::size_t num_bits = 61;
static constexpr bool dpf_modint = true;
static constexpr bool dpf_fp61 = true;
/// @brief Reduce `v` into the field.
/// @param v the integer to reduce. Defaults to 0
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr fp61(integral_type v = 0) noexcept
: val{reduce(v)}
{ }
/// @brief Copy constructor.
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr fp61(const fp61 &) noexcept = default;
/// @brief Move constructor.
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr fp61(fp61 &&) noexcept = default;
/// @brief Copy assignment.
/// @return `*this`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr fp61 & operator=(const fp61 &) noexcept = default;
/// @brief Move assignment.
/// @return `*this`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr fp61 & operator=(fp61 &&) noexcept = default;
/// @brief The reduced representative in `[0, p)`.
/// @return the stored field element
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
constexpr integral_type raw() const noexcept { return val; }
/// @brief Same value as `raw()`.
/// @return the stored field element
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
explicit constexpr operator integral_type() const noexcept { return val; }
/// @brief Mersenne reduction of a 64-bit word.
/// @param x the integer to reduce
/// @return `x` modulo `2^61-1`, with `p` itself represented as 0
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static constexpr integral_type reduce(integral_type x) noexcept
{
x = (x & fp61_mod) + (x >> 61);
if (x >= fp61_mod)
x -= fp61_mod;
return x;
}
/// @brief Field addition.
/// @param a left addend
/// @param b right addend
/// @return `a + b` in the field
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
friend constexpr fp61 operator+(fp61 a, fp61 b) noexcept
{
return fp61{a.val + b.val};
}
/// @brief Field subtraction.
/// @param a minuend
/// @param b subtrahend
/// @return `a - b` in the field
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
friend constexpr fp61 operator-(fp61 a, fp61 b) noexcept
{
return fp61{a.val + fp61_mod - b.val};
}
/// @brief Field negation.
/// @param a the element to negate
/// @return `-a`, with `-0 = 0`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
friend constexpr fp61 operator-(fp61 a) noexcept
{
return fp61{a.val == 0 ? 0 : fp61_mod - a.val};
}
/// @brief Field multiplication.
/// @param a left factor
/// @param b right factor
/// @return `a * b` in the field
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
friend constexpr fp61 operator*(fp61 a, fp61 b) noexcept
{
using u128 = unsigned __int128;
const u128 p = static_cast<u128>(a.val) * static_cast<u128>(b.val);
const auto lo = static_cast<integral_type>(p) & fp61_mod;
const auto mid = static_cast<integral_type>(p >> 61) & fp61_mod;
const auto hi = static_cast<integral_type>(p >> 122);
return fp61{lo + mid + hi};
}
/// @brief Field equality.
/// @param a left element
/// @param b right element
/// @return `true` when the reduced values match
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
friend constexpr bool operator==(fp61 a, fp61 b) noexcept
{
return a.val == b.val;
}
/// @brief Field inequality.
/// @param a left element
/// @param b right element
/// @return `true` when the reduced values differ
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
friend constexpr bool operator!=(fp61 a, fp61 b) noexcept
{
return a.val != b.val;
}
/// @brief Write the reduced representative in decimal.
/// @param os the output stream
/// @param a the element to write
/// @return `os`
friend std::ostream & operator<<(std::ostream & os, fp61 a)
{
return os << a.val;
}
private:
integral_type val;
};
namespace utils
{
template <>
struct bitlength_of<fp61>
: std::integral_constant<std::size_t, 61>
{ };
template <>
struct has_characteristic_two<fp61> : std::false_type
{ };
} // namespace utils
namespace leaf_arithmetic
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
template <>
struct add_t<fp61, simde__m128i>
{
auto operator()(const simde__m128i & a, const simde__m128i & b) const
{
return add_t<fp61::integral_type, simde__m128i>{}(a, b);
}
};
template <>
struct subtract_t<fp61, simde__m128i>
{
auto operator()(const simde__m128i & a, const simde__m128i & b) const
{
return subtract_t<fp61::integral_type, simde__m128i>{}(a, b);
}
};
template <>
struct add_t<fp61, simde__m256i>
{
auto operator()(const simde__m256i & a, const simde__m256i & b) const
{
return add_t<fp61::integral_type, simde__m256i>{}(a, b);
}
};
template <>
struct subtract_t<fp61, simde__m256i>
{
auto operator()(const simde__m256i & a, const simde__m256i & b) const
{
return subtract_t<fp61::integral_type, simde__m256i>{}(a, b);
}
};
HEDLEY_PRAGMA(GCC diagnostic pop)
} // namespace leaf_arithmetic
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_FP61_HPP__

View file

@ -50,16 +50,21 @@
namespace dpf namespace dpf
{ {
/// Shares and the correction words opened along the query trie. /// @brief Shares and the correction words opened along the query trie.
/// `correction_words[i]` / `correction_advice[i]` match a reusable key at /// @details `correction_words[i]` / `correction_advice[i]` match a reusable key at
/// the same target for every `i < live_levels`. `leaf_live` means the /// the same target for every `i < live_levels`. `leaf_live` means the
/// target's leaf was in the trie, so `leaf` is that key's leaf word. /// target's leaf was in the trie, so `leaf` is that key's leaf word.
/// @tparam Output output
/// @tparam Leaf leaf
template <typename Output, typename Leaf> template <typename Output, typename Leaf>
struct geneval_result struct geneval_result
{ {
std::vector<Output> party0; std::vector<Output> party0;
std::vector<Output> party1; std::vector<Output> party1;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::vector<simde__m128i, aligned_allocator<simde__m128i>> correction_words; std::vector<simde__m128i, aligned_allocator<simde__m128i>> correction_words;
HEDLEY_PRAGMA(GCC diagnostic pop)
std::vector<uint8_t> correction_advice; std::vector<uint8_t> correction_advice;
std::size_t live_levels = 0; std::size_t live_levels = 0;
bool leaf_live = false; bool leaf_live = false;
@ -88,8 +93,11 @@ T geneval_flipped(T x)
return x; return x;
} }
/// Leaf-node id of an already MSB-flipped input. The id is the high /// @brief Leaf-node id of an already MSB-flipped input. The id is the high
/// `depth` bits; the low `lg(outputs_per_leaf)` bits select the lane. /// `depth` bits; the low `lg(outputs_per_leaf)` bits select the lane.
/// @tparam Dpf dpf
/// @param x the `x`
/// @return Leaf-node id of an already MSB-flipped input
template <typename Dpf> template <typename Dpf>
uint64_t geneval_leaf_id(typename Dpf::input_type x) uint64_t geneval_leaf_id(typename Dpf::input_type x)
{ {
@ -134,9 +142,9 @@ template <typename InteriorPRG,
typename OutputT, typename OutputT,
typename RootSampler, typename RootSampler,
typename PadRng> typename PadRng>
auto geneval_run(bool arith, InputT x0, InputT x1, auto geneval_run(bool arith, bool arith_out, InputT x0, InputT x1,
const std::vector<InputT> & queries, RootSampler & root_sampler, const std::vector<InputT> & queries, RootSampler & root_sampler,
PadRng & pads, OutputT y) PadRng & pads, OutputT y0, OutputT y1 = OutputT{})
{ {
static_assert(std::is_integral_v<InputT>, static_assert(std::is_integral_v<InputT>,
"geneval input shares are an integral domain"); "geneval input shares are an integral domain");
@ -147,7 +155,11 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
using dpf_type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, InputT, OutputT>; using dpf_type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, InputT, OutputT>;
using node = typename dpf_type::interior_node; using node = typename dpf_type::interior_node;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using leaf_node = leaf_node_t<node, OutputT>; using leaf_node = leaf_node_t<node, OutputT>;
HEDLEY_PRAGMA(GCC diagnostic pop)
using outputs_tuple = std::tuple<OutputT>;
constexpr std::size_t depth = dpf_type::depth; constexpr std::size_t depth = dpf_type::depth;
if (queries.empty()) if (queries.empty())
@ -183,8 +195,16 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
constexpr auto to_int = utils::to_integral_type<InputT>{}; constexpr auto to_int = utils::to_integral_type<InputT>{};
const node root0 = dpf::unset_lo_bit(static_cast<node>(root_sampler())); using tree = dpf::tree_traits<InteriorPRG>;
const node root1 = dpf::set_lo_bit(static_cast<node>(root_sampler())); HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
node roots[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(roots, [&]() -> node {
return static_cast<node>(root_sampler());
});
const node root0 = roots[0];
const node root1 = roots[1];
struct slot struct slot
{ {
@ -195,7 +215,10 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
std::vector<slot> frontier; std::vector<slot> frontier;
frontier.push_back(slot{0, root0, root1}); frontier.push_back(slot{0, root0, root1});
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
geneval_result<OutputT, leaf_node> result; geneval_result<OutputT, leaf_node> result;
HEDLEY_PRAGMA(GCC diagnostic pop)
std::memset(&result.leaf, 0, sizeof(result.leaf)); std::memset(&result.leaf, 0, sizeof(result.leaf));
result.correction_words.reserve(depth); result.correction_words.reserve(depth);
result.correction_advice.reserve(depth); result.correction_advice.reserve(depth);
@ -207,6 +230,7 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
const uint8_t bit0 = static_cast<uint8_t>(!!(to_int(mask) & to_int(x0c))); const uint8_t bit0 = static_cast<uint8_t>(!!(to_int(mask) & to_int(x0c)));
const uint8_t bit1 = static_cast<uint8_t>(!!(to_int(mask) & to_int(x1c))); const uint8_t bit1 = static_cast<uint8_t>(!!(to_int(mask) & to_int(x1c)));
const uint64_t parent_id = geneval_prefix(secret_leaf, depth, level); const uint64_t parent_id = geneval_prefix(secret_leaf, depth, level);
const bool is_last = tree::is_last_level(level, depth);
node L0 = simde_mm_setzero_si128(); node L0 = simde_mm_setzero_si128();
node R0 = simde_mm_setzero_si128(); node R0 = simde_mm_setzero_si128();
@ -225,8 +249,8 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
{ {
if (n.id == parent_id) if (n.id == parent_id)
level_live = true; level_live = true;
const auto c0 = InteriorPRG::eval01(dpf::unset_lo_2bits(n.s0)); const auto c0 = tree::expand(n.s0, is_last);
const auto c1 = InteriorPRG::eval01(dpf::unset_lo_2bits(n.s1)); const auto c1 = tree::expand(n.s1, is_last);
L0 = ds_xor(L0, c0[0]); L0 = ds_xor(L0, c0[0]);
R0 = ds_xor(R0, c0[1]); R0 = ds_xor(R0, c0[1]);
L1 = ds_xor(L1, c1[0]); L1 = ds_xor(L1, c1[0]);
@ -242,21 +266,42 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
auto opened = proto.open_cw(blinds); auto opened = proto.open_cw(blinds);
cw = opened.first; cw = opened.first;
advice = opened.second; advice = opened.second;
if constexpr (tree::is_half_tree)
{
if (!is_last)
advice = 0;
}
++result.live_levels; ++result.live_levels;
} }
else else
{ {
still_live = false; still_live = false;
cw = pads.block(); cw = pads.block();
if constexpr (tree::is_half_tree)
{
if (!is_last)
{
advice = 0;
}
else
{
const uint8_t t0 = static_cast<uint8_t>(pads.bit() & 1u); const uint8_t t0 = static_cast<uint8_t>(pads.bit() & 1u);
const uint8_t t1 = static_cast<uint8_t>(pads.bit() & 1u); const uint8_t t1 = static_cast<uint8_t>(pads.bit() & 1u);
advice = static_cast<uint8_t>((t1 << 1) | t0); advice = static_cast<uint8_t>((t1 << 1) | t0);
} }
}
else
{
const uint8_t t0 = static_cast<uint8_t>(pads.bit() & 1u);
const uint8_t t1 = static_cast<uint8_t>(pads.bit() & 1u);
advice = static_cast<uint8_t>((t1 << 1) | t0);
}
}
result.correction_words.push_back(cw); result.correction_words.push_back(cw);
result.correction_advice.push_back(advice); result.correction_advice.push_back(advice);
const node cw0 = dpf::set_lo_bit(cw, advice & 1u); const node cw0 = tree::pack_cw(cw, advice, false, is_last);
const node cw1 = dpf::set_lo_bit(cw, (advice >> 1) & 1u); const node cw1 = tree::pack_cw(cw, advice, true, is_last);
const std::size_t child_bits = level + 1; const std::size_t child_bits = level + 1;
std::vector<slot> next; std::vector<slot> next;
next.reserve(exps.size() * 2); next.reserve(exps.size() * 2);
@ -294,12 +339,24 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
} }
if (on == nullptr) if (on == nullptr)
throw std::logic_error("geneval: secret leaf missing from trie"); throw std::logic_error("geneval: secret leaf missing from trie");
if (arith_out)
{
const uint8_t t0 = static_cast<uint8_t>(dpf::get_lo_bit(on->s0));
const uint8_t t1 = static_cast<uint8_t>(dpf::get_lo_bit(on->s1));
const std::size_t lane = static_cast<std::size_t>(to_int(alpha));
result.leaf = proto.template open_arith_leaf<ExteriorPRG, 0, outputs_tuple>(
dpf::unset_lo_2bits(on->s0), dpf::unset_lo_2bits(on->s1), t0, t1,
y0, y1, std::size_t{0}, lane);
}
else
{
const bool sign0 = dpf::get_lo_bit(on->s0); const bool sign0 = dpf::get_lo_bit(on->s0);
auto built = dpf::make_leaves<ExteriorPRG>(alpha, auto built = dpf::make_leaves<ExteriorPRG>(alpha,
dpf::unset_lo_2bits(on->s0), dpf::unset_lo_2bits(on->s1), sign0, dpf::unset_lo_2bits(on->s0), dpf::unset_lo_2bits(on->s1), sign0,
std::size_t{0}, y); std::size_t{0}, y0);
result.leaf = std::get<0>(built.first.first); result.leaf = std::get<0>(built.first.first);
} }
}
result.party0.reserve(flipped.size()); result.party0.reserve(flipped.size());
result.party1.reserve(flipped.size()); result.party1.reserve(flipped.size());
@ -326,6 +383,20 @@ auto geneval_run(bool arith, InputT x0, InputT x1,
return result; return result;
} }
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
auto geneval_run(bool arith, InputT x0, InputT x1,
const std::vector<InputT> & queries, RootSampler & root_sampler,
PadRng & pads, OutputT y)
{
return geneval_run<InteriorPRG, ExteriorPRG>(arith, false, x0, x1, queries,
root_sampler, pads, y, OutputT{});
}
template <typename InteriorPRG, template <typename InteriorPRG,
typename ExteriorPRG, typename ExteriorPRG,
typename InputT, typename InputT,
@ -335,8 +406,8 @@ template <typename InteriorPRG,
auto geneval_run(InputT x0, InputT x1, const std::vector<InputT> & queries, auto geneval_run(InputT x0, InputT x1, const std::vector<InputT> & queries,
RootSampler & root_sampler, PadRng & pads, OutputT y) RootSampler & root_sampler, PadRng & pads, OutputT y)
{ {
return geneval_run<InteriorPRG, ExteriorPRG>(false, x0, x1, queries, return geneval_run<InteriorPRG, ExteriorPRG>(false, false, x0, x1, queries,
root_sampler, pads, y); root_sampler, pads, y, OutputT{});
} }
template <typename InputT> template <typename InputT>
@ -400,7 +471,26 @@ std::vector<InputT> geneval_inclusive(InputT from, InputT to)
} // namespace detail } // namespace detail
/// Geneval at one public point. The secret point is `x0 XOR x1`. /// @name Point geneval
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param query the query point
/// @param rng the Doerner–Shelat randomness tapes
/// @{
/// @brief The secret point is `x0 XOR x1`.
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param query the query point
/// @param rng the Doerner–Shelat randomness tapes
/// @param y the payload
/// @return the opened shares and correction words
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -411,11 +501,17 @@ HEDLEY_WARN_UNUSED_RESULT
auto geneval_point(InputT x0, InputT x1, InputT query, auto geneval_point(InputT x0, InputT x1, InputT query,
ds_randomness<RootSampler, PadRng> rng, OutputT y) ds_randomness<RootSampler, PadRng> rng, OutputT y)
{ {
return detail::geneval_run<InteriorPRG, ExteriorPRG>(false, x0, x1, return detail::geneval_run<InteriorPRG, ExteriorPRG>(false, false, x0, x1,
std::vector<InputT>{query}, rng.root, rng.pad, y); std::vector<InputT>{query}, rng.root, rng.pad, y, OutputT{});
} }
/// Geneval at one public point. The secret point is `x0 + x1`. /// @brief The secret point is `x0 + x1`.
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param query the query point
/// @param rng the Doerner–Shelat randomness tapes
/// @param y the payload
/// @return the opened shares and correction words
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -426,11 +522,70 @@ HEDLEY_WARN_UNUSED_RESULT
auto geneval_point(arith_input_t, InputT x0, InputT x1, InputT query, auto geneval_point(arith_input_t, InputT x0, InputT x1, InputT query,
ds_randomness<RootSampler, PadRng> rng, OutputT y) ds_randomness<RootSampler, PadRng> rng, OutputT y)
{ {
return detail::geneval_run<InteriorPRG, ExteriorPRG>(true, x0, x1, return detail::geneval_run<InteriorPRG, ExteriorPRG>(true, false, x0, x1,
std::vector<InputT>{query}, rng.root, rng.pad, y); std::vector<InputT>{query}, rng.root, rng.pad, y, OutputT{});
} }
/// Geneval on the inclusive interval `[from, to]`. /// @brief XOR-index shares, additively shared payload `y0 + y1 = β`.
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param query the query point
/// @param rng the Doerner–Shelat randomness tapes
/// @param y0 party 0's share of the payload
/// @param y1 party 1's share of the payload
/// @return the opened shares and correction words
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_point(arith_output_t, InputT x0, InputT x1, InputT query,
ds_randomness<RootSampler, PadRng> rng, OutputT y0, OutputT y1)
{
return detail::geneval_run<InteriorPRG, ExteriorPRG>(false, true, x0, x1,
std::vector<InputT>{query}, rng.root, rng.pad, y0, y1);
}
/// @brief Additive index and additive payload shares.
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param query the query point
/// @param rng the Doerner–Shelat randomness tapes
/// @param y0 party 0's share of the payload
/// @param y1 party 1's share of the payload
/// @return the opened shares and correction words
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_point(arith_input_t, arith_output_t, InputT x0, InputT x1,
InputT query, ds_randomness<RootSampler, PadRng> rng, OutputT y0, OutputT y1)
{
return detail::geneval_run<InteriorPRG, ExteriorPRG>(true, true, x0, x1,
std::vector<InputT>{query}, rng.root, rng.pad, y0, y1);
}
/// @}
/// @brief Geneval on the inclusive interval `[from, to]`.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param from the inclusive start of the range
/// @param to the `to`
/// @param rng the Doerner–Shelat randomness tapes
/// @param y the `y`
/// @return Geneval on the inclusive interval `[from, to]`
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -459,7 +614,18 @@ auto geneval_interval(arith_input_t, InputT x0, InputT x1, InputT from,
detail::geneval_inclusive(from, to), rng.root, rng.pad, y); detail::geneval_inclusive(from, to), rng.root, rng.pad, y);
} }
/// Geneval on the whole domain. Refuses a domain above 2^20 inputs. /// @brief Geneval on the whole domain. Refuses a domain above 2^20 inputs.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param rng the Doerner–Shelat randomness tapes
/// @param y the `y`
/// @return Geneval on the whole domain
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -488,7 +654,21 @@ auto geneval_full(arith_input_t, InputT x0, InputT x1,
detail::geneval_full_domain<InputT>(), rng.root, rng.pad, y); detail::geneval_full_domain<InputT>(), rng.root, rng.pad, y);
} }
/// Geneval on a public sequence, in the order given. /// @brief Geneval on a public sequence, in the order given.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam ForwardIterator forward iterator type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param rng the Doerner–Shelat randomness tapes
/// @param y the `y`
/// @return Geneval on a public sequence, in the order given
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -522,14 +702,17 @@ auto geneval_sequence(arith_input_t, InputT x0, InputT x1,
std::move(qs), rng.root, rng.pad, y); std::move(qs), rng.root, rng.pad, y);
} }
/// Opened comparison key material and one prefix share per endpoint. /// @brief Opened comparison key material and one prefix share per endpoint.
/// `live_levels` is the full depth: a comparison value word depends on the /// @details `live_levels` is the full depth: a comparison value word depends on the
/// secret path at every level, so there is no early dummy-word tail. /// secret path at every level, so there is no early dummy-word tail.
struct geneval_cmp_result struct geneval_cmp_result
{ {
std::vector<uint64_t> party0; std::vector<uint64_t> party0;
std::vector<uint64_t> party1; std::vector<uint64_t> party1;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::vector<simde__m128i, aligned_allocator<simde__m128i>> correction_words; std::vector<simde__m128i, aligned_allocator<simde__m128i>> correction_words;
HEDLEY_PRAGMA(GCC diagnostic pop)
std::vector<uint8_t> correction_advice; std::vector<uint8_t> correction_advice;
std::vector<uint64_t> value_cw; std::vector<uint64_t> value_cw;
std::vector<uint64_t> tail_cw; std::vector<uint64_t> tail_cw;
@ -540,10 +723,30 @@ struct geneval_cmp_result
std::size_t live_levels = 0; std::size_t live_levels = 0;
}; };
/// Doerner–Shelat comparison geneval. `x0 XOR x1` is the secret point, in the /// @name Comparison geneval
/// same share convention as `geneval_point`. `spec` is an `lt` / `leq` / `gt` /// @tparam InputT input domain type
/// / `geq` pack. Each endpoint is returned in order as the two parties' /// @tparam ForwardIterator forward iterator type
/// `eval_point(cmp, ...)` shares. An empty range opens nothing. /// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param rng the Doerner–Shelat randomness tapes
/// @return the opened comparison shares
/// @{
/// @brief `x0 XOR x1` is the secret point, in the same share convention as
/// `geneval_point`. `spec` is an `lt` / `leq` / `gt` / `geq` pack. Each
/// endpoint is returned in order as the two parties' `eval_point(cmp, ...)`
/// shares. An empty range opens nothing.
/// @tparam Spec comparison or interval specification
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param rng the Doerner–Shelat randomness tapes
/// @param spec the comparison specification
template <typename InputT, template <typename InputT,
typename ForwardIterator, typename ForwardIterator,
typename RootSampler, typename RootSampler,
@ -597,7 +800,14 @@ geneval_cmp_result geneval_cmp(InputT x0, InputT x1,
return out; return out;
} }
/// Comparison geneval with additive shares of the point (`x0 + x1`). /// @brief Additive shares of the point (`x0 + x1`).
/// @tparam Spec comparison or interval specification
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param rng the Doerner–Shelat randomness tapes
/// @param spec the comparison specification
template <typename InputT, template <typename InputT,
typename ForwardIterator, typename ForwardIterator,
typename RootSampler, typename RootSampler,
@ -651,7 +861,13 @@ geneval_cmp_result geneval_cmp(arith_input_t, InputT x0, InputT x1,
return out; return out;
} }
/// `gt(beta)` comparison geneval. `if_false` is 0. /// @brief `gt(beta)` on XOR shares of the point. `if_false` is 0.
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param rng the Doerner–Shelat randomness tapes
/// @param beta the true payload
template <typename InputT, template <typename InputT,
typename ForwardIterator, typename ForwardIterator,
typename RootSampler, typename RootSampler,
@ -665,6 +881,13 @@ geneval_cmp_result geneval_cmp(InputT x0, InputT x1,
std::move(rng), dpf::gt(beta)); std::move(rng), dpf::gt(beta));
} }
/// @brief `gt(beta)` on additive shares of the point. `if_false` is 0.
/// @param x0 party 0's share of the secret point
/// @param x1 party 1's share of the secret point
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param rng the Doerner–Shelat randomness tapes
/// @param beta the true payload
template <typename InputT, template <typename InputT,
typename ForwardIterator, typename ForwardIterator,
typename RootSampler, typename RootSampler,
@ -678,6 +901,8 @@ geneval_cmp_result geneval_cmp(arith_input_t, InputT x0, InputT x1,
std::move(rng), dpf::gt(beta)); std::move(rng), dpf::gt(beta));
} }
/// @}
} // namespace dpf } // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_GENEVAL_HPP__ #endif // LIBDPF_INCLUDE_DPF_GENEVAL_HPP__

File diff suppressed because it is too large Load diff

View file

@ -45,7 +45,7 @@ struct ic_pack
Beta if_false{}; Beta if_false{};
}; };
/// Spec tag and factory. `dpf::ic(p, q, beta)` builds a pack; /// @brief Spec tag and factory. `dpf::ic(p, q, beta)` builds a pack;
/// `eval_point(dpf::ic, key, x)` evaluates it. /// `eval_point(dpf::ic, key, x)` evaluates it.
struct ic_fn struct ic_fn
{ {
@ -68,6 +68,12 @@ inline constexpr ic_fn ic{};
template <typename T> template <typename T>
struct is_ic_key : std::false_type {}; struct is_ic_key : std::false_type {};
/// @brief One party's interval key: the inner comparison, the public bounds, and the
/// secret correction shares.
/// @tparam Party party index, `0` or `1`
/// @tparam Key key type
/// @tparam Input input domain type
/// @tparam Beta payload type
template <std::size_t Party, typename Key, typename Input, typename Beta> template <std::size_t Party, typename Key, typename Input, typename Beta>
struct ic_key struct ic_key
{ {
@ -76,24 +82,27 @@ struct ic_key
using input_type = Input; using input_type = Input;
using key_type = party_key<Party, Key>; using key_type = party_key<Party, Key>;
using beta_type = Beta; using beta_type = Beta;
using share_type = std::conditional_t<
detail::cmp_group_info<dpf::concrete_type_t<Beta>>::custom,
detail::group_elem, uint64_t>;
key_type key; key_type key;
uint64_t lo = 0; uint64_t lo = 0;
uint64_t hi = 0; uint64_t hi = 0;
uint64_t input_mask = 0; uint64_t input_mask = 0;
uint64_t group_mask = 0; uint64_t group_mask = 0;
/// Share of `δ`. Public `c_x ∈ {-1,0,1}` scales it locally. /// @brief Share of `δ`. Public `c_x ∈ {-1,0,1}` scales it locally.
uint64_t delta_share = 0; share_type delta_share{};
/// Share of `δ · c_r + if_false`. /// @brief Share of `δ · c_r + if_false`.
uint64_t cr_share = 0; share_type cr_share{};
/// Wildcard only: shares of `1` and of `c_r`, scaled by `δ` in `assign_cmp`. /// @brief Wildcard only: shares of `1` and of `c_r`, scaled by `δ` in `assign_cmp`.
uint64_t delta_coeff = 0; share_type delta_coeff{};
uint64_t cr_coeff = 0; share_type cr_coeff{};
bool assigned = !wildcard; bool assigned = !wildcard;
ic_key(key_type k, uint64_t lo_in, uint64_t hi_in, uint64_t nmask, ic_key(key_type k, uint64_t lo_in, uint64_t hi_in, uint64_t nmask,
uint64_t gmask, uint64_t dshare, uint64_t cshare, uint64_t dcoeff, uint64_t gmask, share_type dshare, share_type cshare, share_type dcoeff,
uint64_t ccoeff) share_type ccoeff) noexcept(std::is_nothrow_move_constructible_v<key_type>)
: key(std::move(k)) : key(std::move(k))
, lo(lo_in) , lo(lo_in)
, hi(hi_in) , hi(hi_in)
@ -167,20 +176,40 @@ constexpr uint64_t mul_mask(uint64_t a, uint64_t b, uint64_t mask) noexcept
return static_cast<uint64_t>(static_cast<unsigned __int128>(a) * b) & mask; return static_cast<uint64_t>(static_cast<unsigned __int128>(a) * b) & mask;
} }
/// Dealer correction in Fig. 3, as an element of the payload group. /// @brief Public integer in Fig. 3, before it is embedded in the payload group.
/// @param r the `r`
/// @param p the `p`
/// @param q the `q`
/// @param nmask the mask of the live input bits
/// @return Public integer in Fig. 3, before it is embedded in the payload group
HEDLEY_CONST
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr int correction_s(uint64_t r, uint64_t p, uint64_t q,
uint64_t nmask) noexcept
{
const uint64_t aq = (q + r) & nmask;
const uint64_t ap = (p + r) & nmask;
const uint64_t q0 = (q + 1ULL) & nmask;
const uint64_t aq0 = (q0 + r) & nmask;
return (ap > aq ? 1 : 0) - (ap > p ? 1 : 0)
+ (aq0 > q0 ? 1 : 0) + (aq == nmask ? 1 : 0);
}
/// @brief Dealer correction in Fig. 3, as an element of the payload group.
/// @param r the `r`
/// @param p the `p`
/// @param q the `q`
/// @param nmask the mask of the live input bits
/// @param gmask the mask of the live payload bits
/// @return Dealer correction in Fig. 3, as an element of the payload group
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr uint64_t correction(uint64_t r, uint64_t p, uint64_t q, constexpr uint64_t correction(uint64_t r, uint64_t p, uint64_t q,
uint64_t nmask, uint64_t gmask) noexcept uint64_t nmask, uint64_t gmask) noexcept
{ {
const uint64_t aq = (q + r) & nmask; return embed_small(correction_s(r, p, q, nmask), gmask);
const uint64_t ap = (p + r) & nmask;
const uint64_t q0 = (q + 1ULL) & nmask;
const uint64_t aq0 = (q0 + r) & nmask;
const int s = (ap > aq ? 1 : 0) - (ap > p ? 1 : 0)
+ (aq0 > q0 ? 1 : 0) + (aq == nmask ? 1 : 0);
return embed_small(s, gmask);
} }
HEDLEY_CONST HEDLEY_CONST
@ -251,8 +280,10 @@ void check_bounds(const ic_pack<Beta> & spec)
} }
template <typename Beta> template <typename Beta>
HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
uint64_t group_mask_of() noexcept HEDLEY_ALWAYS_INLINE
constexpr uint64_t group_mask_of() noexcept
{ {
using B = concrete_type_t<Beta>; using B = concrete_type_t<Beta>;
if constexpr (std::is_same_v<B, dpf::bit>) if constexpr (std::is_same_v<B, dpf::bit>)
@ -261,11 +292,11 @@ uint64_t group_mask_of() noexcept
return dcf_impl::default_mask_for_bits(utils::bitlength_of_v<B>); return dcf_impl::default_mask_for_bits(utils::bitlength_of_v<B>);
} }
template <std::size_t Party, typename Key, typename Input, typename Beta> template <std::size_t Party, typename Key, typename Input, typename Beta,
typename Share>
ic_key<Party, Key, Input, Beta> make_side(party_key<Party, Key> key, ic_key<Party, Key, Input, Beta> make_side(party_key<Party, Key> key,
uint64_t lo, uint64_t hi, uint64_t nmask, uint64_t gmask, uint64_t lo, uint64_t hi, uint64_t nmask, uint64_t gmask,
uint64_t delta_share, uint64_t cr_share, Share delta_share, Share cr_share, Share delta_coeff, Share cr_coeff)
uint64_t delta_coeff, uint64_t cr_coeff)
{ {
return ic_key<Party, Key, Input, Beta>(std::move(key), lo, hi, nmask, gmask, return ic_key<Party, Key, Input, Beta>(std::move(key), lo, hi, nmask, gmask,
delta_share, cr_share, delta_coeff, cr_coeff); delta_share, cr_share, delta_coeff, cr_coeff);
@ -281,6 +312,46 @@ auto finish(uint64_t r_bits, const ic_pack<Beta> & spec, Pair && inner)
const uint64_t nmask = input_mask_of<in_type>(); const uint64_t nmask = input_mask_of<in_type>();
const uint64_t gmask = group_mask_of<Beta>(); const uint64_t gmask = group_mask_of<Beta>();
constexpr bool wild = is_wildcard_v<Beta>; constexpr bool wild = is_wildcard_v<Beta>;
if constexpr (detail::cmp_group_info<Beta>::custom)
{
using prg = typename raw_key::interior_prg;
const auto layout = detail::group_layout<out_beta>();
auto delta = detail::group_zero(layout);
auto fval = detail::group_zero(layout);
if constexpr (!wild)
{
delta = detail::group_sub(detail::group_from_beta(spec.if_true),
detail::group_from_beta(spec.if_false));
fval = detail::group_from_beta(spec.if_false);
}
const auto cr = detail::group_scalar(
correction_s(r_bits, spec.lo, spec.hi, nmask), layout);
auto splitg = [&](const detail::group_elem & target,
detail::group_elem & a, detail::group_elem & b) {
const auto blind = detail::group_from_node<prg>(
dpf::uniform_sample<typename raw_key::interior_node>(), layout);
a = blind;
b = detail::group_sub(target, blind);
};
detail::group_elem d0{}, d1{}, c0{}, c1{}, dc0{}, dc1{}, cc0{}, cc1{};
if constexpr (wild)
{
splitg(detail::group_one(layout), dc0, dc1);
splitg(cr, cc0, cc1);
}
else
{
splitg(delta, d0, d1);
splitg(detail::group_add(detail::group_mul(delta, cr), fval), c0, c1);
}
auto k0 = make_side<0, raw_key, in_type, out_beta>(std::move(inner.first),
spec.lo, spec.hi, nmask, gmask, d0, c0, dc0, cc0);
auto k1 = make_side<1, raw_key, in_type, out_beta>(std::move(inner.second),
spec.lo, spec.hi, nmask, gmask, d1, c1, dc1, cc1);
return std::make_pair(std::move(k0), std::move(k1));
}
else
{
uint64_t delta = 0; uint64_t delta = 0;
uint64_t fval = 0; uint64_t fval = 0;
if constexpr (!wild) if constexpr (!wild)
@ -310,6 +381,7 @@ auto finish(uint64_t r_bits, const ic_pack<Beta> & spec, Pair && inner)
auto k1 = make_side<1, raw_key, in_type, out_beta>(std::move(inner.second), auto k1 = make_side<1, raw_key, in_type, out_beta>(std::move(inner.second),
spec.lo, spec.hi, nmask, gmask, d1, c1, dc1, cc1); spec.lo, spec.hi, nmask, gmask, d1, c1, dc1, cc1);
return std::make_pair(std::move(k0), std::move(k1)); return std::make_pair(std::move(k0), std::move(k1));
}
} }
template <typename Beta> template <typename Beta>
@ -318,6 +390,15 @@ auto inner_lt(const ic_pack<Beta> & spec)
using B = std::decay_t<Beta>; using B = std::decay_t<Beta>;
if constexpr (is_wildcard_v<B>) if constexpr (is_wildcard_v<B>)
return lt(spec.if_true, spec.if_false); return lt(spec.if_true, spec.if_false);
else if constexpr (detail::cmp_group_info<B>::custom)
{
const auto layout = detail::group_layout<concrete_type_t<B>>();
const auto delta = detail::group_sub(
detail::group_from_beta(spec.if_true),
detail::group_from_beta(spec.if_false));
return lt(detail::group_to_beta<B>(delta),
detail::group_to_beta<B>(detail::group_zero(layout)));
}
else else
{ {
const uint64_t gmask = group_mask_of<B>(); const uint64_t gmask = group_mask_of<B>();
@ -342,9 +423,35 @@ auto eval_one(const IcKey & k, Query && x, Memo & memo)
throw std::invalid_argument( throw std::invalid_argument(
"ic eval: wildcard payload not assigned (call assign_cmp)"); "ic eval: wildcard payload not assigned (call assign_cmp)");
using in_type = typename IcKey::input_type; using in_type = typename IcKey::input_type;
using beta = typename IcKey::beta_type;
const uint64_t xu = bits_of(in_type(std::forward<Query>(x))); const uint64_t xu = bits_of(in_type(std::forward<Query>(x)));
const uint64_t xp = shift_p(xu, k.lo, k.input_mask); const uint64_t xp = shift_p(xu, k.lo, k.input_mask);
const uint64_t xq = shift_q0(xu, k.hi, k.input_mask); const uint64_t xq = shift_q0(xu, k.hi, k.input_mask);
if constexpr (detail::cmp_group_info<beta>::custom)
{
auto opened = [](const auto & v) {
if constexpr (is_secret_share_v<std::decay_t<decltype(v)>>)
return detail::group_from_beta(v.raw());
else
return detail::group_from_beta(v);
};
const auto a = opened(eval_point<beta>(dpf::cmp, k.key,
input_from_bits<in_type>(xp), memo));
const auto b = opened(eval_point<beta>(dpf::cmp, k.key,
input_from_bits<in_type>(xq), memo));
const int cx = public_cx(xu, k.lo, k.hi, k.input_mask);
auto scaled = detail::group_zero(a);
if (cx == 1)
scaled = k.delta_share;
else if (cx == -1)
scaled = detail::group_neg(k.delta_share);
const auto y = detail::group_add(detail::group_add(
detail::group_add(detail::group_neg(a), b), k.cr_share), scaled);
return make_eval_cmp_result<typename IcKey::key_type>(
detail::group_to_beta<beta>(y));
}
else
{
const uint64_t a = opened_u64( const uint64_t a = opened_u64(
eval_point(dpf::cmp, k.key, input_from_bits<in_type>(xp), memo), eval_point(dpf::cmp, k.key, input_from_bits<in_type>(xp), memo),
k.group_mask); k.group_mask);
@ -361,12 +468,20 @@ auto eval_one(const IcKey & k, Query && x, Memo & memo)
+ scaled) & k.group_mask; + scaled) & k.group_mask;
return make_eval_cmp_result<typename IcKey::key_type>( return make_eval_cmp_result<typename IcKey::key_type>(
dcf_impl::u64_to_beta<typename IcKey::beta_type>(y)); dcf_impl::u64_to_beta<typename IcKey::beta_type>(y));
}
} }
} // namespace ic_impl } // namespace ic_impl
} // namespace detail } // namespace detail
/// Dealer key for public bounds `spec` and secret mask `r`. /// @brief Dealer key for public bounds `spec` and secret mask `r`.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam Beta payload type
/// @param r the secret input mask
/// @param spec the public bounds and payloads
/// @return Dealer key for public bounds `spec` and secret mask `r`
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -384,7 +499,24 @@ auto make_dpf(InputT && r, const ic_pack<Beta> & spec)
return detail::ic_impl::finish<input_type>(r_bits, spec, std::move(inner)); return detail::ic_impl::finish<input_type>(r_bits, spec, std::move(inner));
} }
/// Doerner–Shelat key. `r0 XOR r1` is the secret mask. /// @name Doerner–Shelat interval keys
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam Beta payload type
/// @param r0 party 0's share of the mask
/// @param r1 party 1's share of the mask
/// @param spec the public bounds and payloads
/// @{
/// @brief XOR shares. `r0 XOR r1` is the secret mask.
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @param r0 party 0's share of the mask
/// @param r1 party 1's share of the mask
/// @param rng the Doerner–Shelat randomness tapes
/// @param spec the public bounds and payloads
/// @return the two party keys
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename RootSampler, typename RootSampler,
@ -408,7 +540,14 @@ auto make_dpf_doerner_shelat(InputT r0, InputT r1,
return detail::ic_impl::finish<input_type>(r_bits, spec, std::move(inner)); return detail::ic_impl::finish<input_type>(r_bits, spec, std::move(inner));
} }
/// Doerner–Shelat IC key. `r0 + r1` is the secret mask; γ = (r0 + r1) − 1. /// @brief Additive shares. `r0 + r1` is the secret mask; γ = (r0 + r1) − 1.
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @param r0 party 0's share of the mask
/// @param r1 party 1's share of the mask
/// @param rng the Doerner–Shelat randomness tapes
/// @param spec the public bounds and payloads
/// @return the two party keys
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename RootSampler, typename RootSampler,
@ -436,7 +575,8 @@ auto make_dpf_doerner_shelat(arith_input_t, InputT r0, InputT r1,
detail::ic_impl::bits_of(r), spec, std::move(inner)); detail::ic_impl::bits_of(r), spec, std::move(inner));
} }
/// Doerner–Shelat key sampled from the library entropy source. /// @brief XOR shares, sampled from the library entropy source.
/// @return the two party keys
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -451,6 +591,8 @@ auto make_dpf_doerner_shelat(InputT r0, InputT r1, const ic_pack<Beta> & spec)
std::move(r0), std::move(r1), rng, spec); std::move(r0), std::move(r1), rng, spec);
} }
/// @brief Additive shares, sampled from the library entropy source.
/// @return the two party keys
template <typename InteriorPRG = dpf::prg::aes128, template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG, typename ExteriorPRG = InteriorPRG,
typename InputT, typename InputT,
@ -466,7 +608,17 @@ auto make_dpf_doerner_shelat(arith_input_t, InputT r0, InputT r1,
arith_input, std::move(r0), std::move(r1), rng, spec); arith_input, std::move(r0), std::move(r1), rng, spec);
} }
/// Open a wildcard interval payload onto an existing key pair. /// @}
/// @brief Open a wildcard interval payload onto an existing key pair.
/// @tparam Key key type
/// @tparam Input input domain type
/// @tparam Beta payload type
/// @tparam Payload concrete payload type
/// @param k0 the `k0`
/// @param k1 the `k1`
/// @param if_true the payload on a true comparison
/// @param if_false the payload on a false comparison
template <typename Key, typename Input, typename Beta, typename Payload> template <typename Key, typename Input, typename Beta, typename Payload>
void assign_cmp(ic_key<0, Key, Input, Beta> & k0, void assign_cmp(ic_key<0, Key, Input, Beta> & k0,
ic_key<1, Key, Input, Beta> & k1, const Payload & if_true, ic_key<1, Key, Input, Beta> & k1, const Payload & if_true,
@ -474,6 +626,31 @@ void assign_cmp(ic_key<0, Key, Input, Beta> & k0,
{ {
static_assert(Key::cmp_is_wildcard, static_assert(Key::cmp_is_wildcard,
"assign_cmp: interval payload is not a wildcard"); "assign_cmp: interval payload is not a wildcard");
if constexpr (detail::cmp_group_info<dpf::concrete_type_t<Beta>>::custom)
{
using prg = typename Key::interior_prg;
const auto layout = detail::group_layout<dpf::concrete_type_t<Beta>>();
const auto delta = detail::group_sub(
detail::group_from_beta(if_true), detail::group_from_beta(if_false));
const auto fval = detail::group_from_beta(if_false);
assign_cmp(k0.key, k1.key,
detail::group_to_beta<Beta>(delta),
detail::group_to_beta<Beta>(detail::group_zero(layout)));
k0.delta_share = detail::group_mul(k0.delta_coeff, delta);
k1.delta_share = detail::group_mul(k1.delta_coeff, delta);
k0.cr_share = detail::group_mul(k0.cr_coeff, delta);
k1.cr_share = detail::group_mul(k1.cr_coeff, delta);
const auto blind = detail::group_from_node<prg>(
dpf::uniform_sample<typename Key::interior_node>(), layout);
const auto f1 = detail::group_sub(fval, blind);
k0.cr_share = detail::group_add(k0.cr_share, blind);
k1.cr_share = detail::group_add(k1.cr_share, f1);
k0.assigned = true;
k1.assigned = true;
return;
}
else
{
const uint64_t mask = k0.group_mask; const uint64_t mask = k0.group_mask;
const uint64_t delta = const uint64_t delta =
detail::dcf_impl::beta_delta_u64(if_true, if_false, mask); detail::dcf_impl::beta_delta_u64(if_true, if_false, mask);
@ -492,9 +669,17 @@ void assign_cmp(ic_key<0, Key, Input, Beta> & k0,
k1.cr_share = (k1.cr_share + f1) & mask; k1.cr_share = (k1.cr_share + f1) & mask;
k0.assigned = true; k0.assigned = true;
k1.assigned = true; k1.assigned = true;
}
} }
/// Point evaluation. `memo` is a path memoizer for the inner comparison key. /// @brief Point evaluation. `memo` is a path memoizer for the inner comparison key.
/// @tparam IcKey interval-containment key type
/// @tparam Query query point type
/// @tparam Memo path memoizer type
/// @param key the key to evaluate
/// @param x the `x`
/// @param memo the memoizer reused across queries
/// @return Point evaluation
template <typename IcKey, typename Query, template <typename IcKey, typename Query,
typename Memo = basic_path_memoizer<typename IcKey::key_type>, typename Memo = basic_path_memoizer<typename IcKey::key_type>,
typename = std::enable_if_t<is_ic_key_v<IcKey>>> typename = std::enable_if_t<is_ic_key_v<IcKey>>>
@ -504,7 +689,24 @@ auto eval_point(ic_fn, const IcKey & key, Query && x, Memo && memo = Memo{})
return detail::ic_impl::eval_one(key, std::forward<Query>(x), memo); return detail::ic_impl::eval_one(key, std::forward<Query>(x), memo);
} }
/// Inclusive interval `[from, to]` on the input domain. /// @name Interval evaluation
/// @tparam IcKey interval-containment key type
/// @tparam Lane input-domain lane type
/// @tparam Buffer output buffer type
/// @param key the interval key
/// @param from the inclusive start of the range
/// @param to the inclusive end of the range
/// @param buf the output buffer
/// @throws std::invalid_argument if `to < from`
/// @{
/// @brief Inclusive interval `[from, to]` on the input domain.
/// @tparam Memo path memoizer type
/// @param key the interval key
/// @param from the inclusive start of the range
/// @param to the inclusive end of the range
/// @param buf the output buffer
/// @param memo the memoizer reused across queries
template <typename IcKey, typename Lane, typename Buffer, typename Memo, template <typename IcKey, typename Lane, typename Buffer, typename Memo,
typename = std::enable_if_t<is_ic_key_v<IcKey>>> typename = std::enable_if_t<is_ic_key_v<IcKey>>>
void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to, void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to,
@ -528,6 +730,7 @@ void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to,
} }
} }
/// @brief Inclusive interval `[from, to]`, with a fresh path memoizer.
template <typename IcKey, typename Lane, typename Buffer, template <typename IcKey, typename Lane, typename Buffer,
typename = std::enable_if_t<is_ic_key_v<IcKey>>> typename = std::enable_if_t<is_ic_key_v<IcKey>>>
void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to, Buffer && buf) void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to, Buffer && buf)
@ -536,7 +739,25 @@ void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to, Buffer && buf)
eval_interval(ic, key, from, to, std::forward<Buffer>(buf), memo); eval_interval(ic, key, from, to, std::forward<Buffer>(buf), memo);
} }
/// Evaluate the points in `[begin, end)`. /// @}
/// @name Sequence evaluation
/// @tparam IcKey interval-containment key type
/// @tparam Iter iterator type
/// @tparam Buffer output buffer type
/// @param key the interval key
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param buf the output buffer
/// @{
/// @brief Evaluate the points in `[begin, end)`.
/// @tparam Memo path memoizer type
/// @param key the interval key
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param buf the output buffer
/// @param memo the memoizer reused across queries
template <typename IcKey, typename Iter, typename Buffer, typename Memo, template <typename IcKey, typename Iter, typename Buffer, typename Memo,
typename = std::enable_if_t<is_ic_key_v<IcKey>>> typename = std::enable_if_t<is_ic_key_v<IcKey>>>
void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end, void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end,
@ -547,6 +768,7 @@ void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end,
buf[i] = detail::ic_impl::eval_one(key, *it, memo); buf[i] = detail::ic_impl::eval_one(key, *it, memo);
} }
/// @brief Evaluate `[begin, end)`, with a fresh path memoizer.
template <typename IcKey, typename Iter, typename Buffer, template <typename IcKey, typename Iter, typename Buffer,
typename = std::enable_if_t<is_ic_key_v<IcKey>>> typename = std::enable_if_t<is_ic_key_v<IcKey>>>
void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end, Buffer && buf) void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end, Buffer && buf)
@ -555,7 +777,12 @@ void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end, Buffer && buf
eval_sequence(ic, key, begin, end, std::forward<Buffer>(buf), memo); eval_sequence(ic, key, begin, end, std::forward<Buffer>(buf), memo);
} }
/// Buffer of `n` interval shares. /// @}
/// @brief Buffer of `n` interval shares.
/// @tparam IcKey interval-containment key type
/// @param n the `n`
/// @return Buffer of `n` interval shares
template <typename IcKey, typename = std::enable_if_t<is_ic_key_v<IcKey>>> template <typename IcKey, typename = std::enable_if_t<is_ic_key_v<IcKey>>>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
auto make_output_buffer(ic_fn, const IcKey &, std::size_t n) auto make_output_buffer(ic_fn, const IcKey &, std::size_t n)
@ -565,7 +792,14 @@ auto make_output_buffer(ic_fn, const IcKey &, std::size_t n)
return output_buffer<elem>(n); return output_buffer<elem>(n);
} }
/// Buffer large enough for the inclusive interval `[from, to]`. /// @brief Buffer large enough for the inclusive interval `[from, to]`.
/// @tparam IcKey interval-containment key type
/// @tparam Lane input-domain lane type
/// @param key the `key`
/// @param from the inclusive start of the range
/// @param to the `to`
/// @return Buffer large enough for the inclusive interval `[from, to]`
/// @throws std::invalid_argument if `to < from`
template <typename IcKey, typename Lane, template <typename IcKey, typename Lane,
typename = std::enable_if_t<is_ic_key_v<IcKey>>> typename = std::enable_if_t<is_ic_key_v<IcKey>>>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
@ -580,7 +814,22 @@ auto make_output_buffer(ic_fn, const IcKey & key, Lane from, Lane to)
return make_output_buffer(ic, key, static_cast<std::size_t>(n)); return make_output_buffer(ic, key, static_cast<std::size_t>(n));
} }
/// Doerner–Shelat geneval. `r0 XOR r1` is the secret mask. Each query is /// @name Interval geneval
/// @tparam InputT input domain type
/// @tparam Iter iterator type
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam Beta payload type
/// @param r0 party 0's share of the mask
/// @param r1 party 1's share of the mask
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param rng the Doerner–Shelat randomness tapes
/// @param spec the public bounds and payloads
/// @return the opened party shares
/// @{
/// @brief XOR mask. `r0 XOR r1` is the secret mask. Each query is
/// returned already combined into the interval share. /// returned already combined into the interval share.
template <typename InputT, typename Iter, typename RootSampler, typename PadRng, template <typename InputT, typename Iter, typename RootSampler, typename PadRng,
typename Beta> typename Beta>
@ -625,7 +874,7 @@ geneval_cmp_result geneval_ic(InputT r0, InputT r1, Iter begin, Iter end,
return out; return out;
} }
/// Additive-share geneval_ic. `r0 + r1` is the secret mask. /// @brief Additive mask. `r0 + r1` is the secret mask.
template <typename InputT, typename Iter, typename RootSampler, typename PadRng, template <typename InputT, typename Iter, typename RootSampler, typename PadRng,
typename Beta> typename Beta>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
@ -669,6 +918,8 @@ geneval_cmp_result geneval_ic(arith_input_t, InputT r0, InputT r1, Iter begin,
return out; return out;
} }
/// @}
} // namespace dpf } // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_INTERVAL_HPP__ #endif // LIBDPF_INCLUDE_DPF_INTERVAL_HPP__

View file

@ -36,14 +36,16 @@
namespace dpf namespace dpf
{ {
/// Ping-pong pivot math underflows at 0 leaves. Keep a one-node slab so the /// @brief Ping-pong pivot math underflows at 0 leaves. Keep a one-node slab so the
/// root still has a place to land; callers never walk a 0-leaf interval. /// root still has a place to land; callers never walk a 0-leaf interval.
/// @param output_len the `output_len`
/// @return Ping-pong pivot math underflows at 0 leaves
inline std::size_t interval_memoizer_slots(std::size_t output_len) inline std::size_t interval_memoizer_slots(std::size_t output_len)
{ {
return output_len == 0 ? std::size_t{1} : output_len; return output_len == 0 ? std::size_t{1} : output_len;
} }
/// Interval memoizers key on the underlying DPF key type (same rule as path /// @brief Interval memoizers key on the underlying DPF key type (same rule as path
/// memoizers): a memoizer built from `party_key<0, Key>` also accepts /// memoizers): a memoizer built from `party_key<0, Key>` also accepts
/// `party_key<1, Key>` and bare `Key`. /// `party_key<1, Key>` and bare `Key`.
template <typename DpfKey> template <typename DpfKey>
@ -161,18 +163,20 @@ struct interval_memoizer_base
std::optional<integral_type> to_; std::optional<integral_type> to_;
}; };
/// Two-level workspace for one interval. This is what /// @brief Two-level workspace for one interval. This is what
/// `eval_interval(key, from, to)` allocates when you omit the memoizer. /// `eval_interval(key, from, to)` allocates when you omit the memoizer.
/// @tparam DpfKey DPF key type
/// @tparam interior_node interior node
template <typename DpfKey, template <typename DpfKey,
typename Allocator = aligned_allocator< typename Allocator = aligned_allocator<
typename interval_memoizer_key_t<DpfKey>::interior_node>> typename interval_memoizer_key_t<DpfKey>::interior_node>>
struct basic_interval_memoizer final : public interval_memoizer_base<DpfKey> struct basic_interval_memoizer final : public interval_memoizer_base<DpfKey>
{ {
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
private: private:
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using parent = interval_memoizer_base<DpfKey>; using parent = interval_memoizer_base<DpfKey>;
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
public: public:
using unique_ptr = typename Allocator::unique_ptr; using unique_ptr = typename Allocator::unique_ptr;
using return_type = typename interval_memoizer_key_t<DpfKey>::interior_node *; using return_type = typename interval_memoizer_key_t<DpfKey>::interior_node *;
@ -241,17 +245,19 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
unique_ptr buf; unique_ptr buf;
}; };
/// Every level of the interval. `retains_all_levels` is true. /// @brief Every level of the interval. `retains_all_levels` is true.
/// @tparam DpfKey DPF key type
/// @tparam interior_node interior node
template <typename DpfKey, template <typename DpfKey,
typename Allocator = aligned_allocator< typename Allocator = aligned_allocator<
typename interval_memoizer_key_t<DpfKey>::interior_node>> typename interval_memoizer_key_t<DpfKey>::interior_node>>
struct full_tree_interval_memoizer final : public interval_memoizer_base<DpfKey> struct full_tree_interval_memoizer final : public interval_memoizer_base<DpfKey>
{ {
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
private: private:
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using parent = interval_memoizer_base<DpfKey>; using parent = interval_memoizer_base<DpfKey>;
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
public: public:
using node_type = typename interval_memoizer_key_t<DpfKey>::interior_node; using node_type = typename interval_memoizer_key_t<DpfKey>::interior_node;
using unique_ptr = typename Allocator::unique_ptr; using unique_ptr = typename Allocator::unique_ptr;
@ -333,7 +339,10 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
} }
}; };
/// Interval memoizer whose leaf depth is `StopLevel` (incremental `eval_interval`). /// @brief Interval memoizer whose leaf depth is `StopLevel` (incremental `eval_interval`).
/// @tparam DpfKey DPF key type
/// @tparam StopLevel stop level
/// @tparam interior_node interior node
template <typename DpfKey, std::size_t StopLevel, template <typename DpfKey, std::size_t StopLevel,
typename Allocator = aligned_allocator< typename Allocator = aligned_allocator<
typename interval_memoizer_key_t<DpfKey>::interior_node>> typename interval_memoizer_key_t<DpfKey>::interior_node>>
@ -441,10 +450,13 @@ auto make_interval_memoizer(InputT from, InputT to)
} // namespace detail } // namespace detail
/// Two-level workspace sized for the closed interval `[from, to]`. /// @brief Two-level workspace sized for the closed interval `[from, to]`.
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @param from Inclusive start, in the key's input domain. /// @param from Inclusive start, in the key's input domain.
/// @param to Inclusive end. `to` is at least `from` in that domain. /// @param to Inclusive end. `to` is at least `from` in that domain.
/// @snippet evaluation/memoizers.cpp interval-memoizer /// @snippet evaluation/memoizers.cpp interval-memoizer
/// @return Two-level workspace sized for the closed interval `[from, to]`
template <typename DpfKey, template <typename DpfKey,
typename InputT> typename InputT>
inline auto make_basic_interval_memoizer(InputT from, InputT to) inline auto make_basic_interval_memoizer(InputT from, InputT to)
@ -463,7 +475,9 @@ inline auto make_basic_interval_memoizer(const DpfKey &, InputT from, InputT to)
return make_basic_interval_memoizer<DpfKey>(from, to); return make_basic_interval_memoizer<DpfKey>(from, to);
} }
/// `make_basic_interval_memoizer` sized for the whole input domain. /// @brief `make_basic_interval_memoizer` sized for the whole input domain.
/// @tparam DpfKey DPF key type
/// @return `make_basic_interval_memoizer` sized for the whole input domain
template <typename DpfKey> template <typename DpfKey>
inline auto make_basic_full_memoizer() inline auto make_basic_full_memoizer()
{ {
@ -480,7 +494,12 @@ inline auto make_basic_full_memoizer(const DpfKey &)
return make_basic_full_memoizer<DpfKey>(); return make_basic_full_memoizer<DpfKey>();
} }
/// Full-tree workspace sized for the closed interval `[from, to]`. /// @brief Full-tree workspace sized for the closed interval `[from, to]`.
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @param from the inclusive start of the range
/// @param to the `to`
/// @return Full-tree workspace sized for the closed interval `[from, to]`
template <typename DpfKey, template <typename DpfKey,
typename InputT> typename InputT>
inline auto make_full_tree_interval_memoizer(InputT from, InputT to) inline auto make_full_tree_interval_memoizer(InputT from, InputT to)
@ -499,7 +518,9 @@ inline auto make_full_tree_interval_memoizer(const DpfKey &, InputT from, InputT
return make_full_tree_interval_memoizer<DpfKey>(from, to); return make_full_tree_interval_memoizer<DpfKey>(from, to);
} }
/// `make_full_tree_interval_memoizer` sized for the whole input domain. /// @brief `make_full_tree_interval_memoizer` sized for the whole input domain.
/// @tparam DpfKey DPF key type
/// @return `make_full_tree_interval_memoizer` sized for the whole input domain
template <typename DpfKey> template <typename DpfKey>
inline auto make_full_tree_full_memoizer() inline auto make_full_tree_full_memoizer()
{ {
@ -522,10 +543,17 @@ inline auto make_basic_interval_memoizer_at(std::size_t leaf_nodes)
return basic_interval_memoizer_at<DpfKey, StopLevel>(leaf_nodes); return basic_interval_memoizer_at<DpfKey, StopLevel>(leaf_nodes);
} }
/// Stop-level interval memoizer for output slot `I` of a multi-level key. /// @brief Stop-level interval memoizer for output slot `I` of a multi-level key.
/// Sizes the ping-pong buffer for the lane-domain interval `[from, to]` /// @details Sizes the ping-pong buffer for the lane-domain interval `[from, to]`
/// expanded to `meta[I].tree_level` (the leaf level of slot `I`). This is the /// expanded to `meta[I].tree_level` (the leaf level of slot `I`). This is the
/// default memoizer for a multi-level `eval_interval(out<I>, ...)`. /// default memoizer for a multi-level `eval_interval(out<I>, ...)`.
/// @tparam DpfKey DPF key type
/// @tparam I output index
/// @tparam InputT input domain type
/// @tparam is_multilevel is multilevel
/// @param from the inclusive start of the range
/// @param to the `to`
/// @return Stop-level interval memoizer for output slot `I` of a multi-level key
template <typename DpfKey, std::size_t I, template <typename DpfKey, std::size_t I,
typename InputT, typename InputT,
std::enable_if_t<DpfKey::is_multilevel, bool> = true> std::enable_if_t<DpfKey::is_multilevel, bool> = true>
@ -546,7 +574,10 @@ inline auto make_basic_interval_memoizer(InputT from, InputT to)
const bool wraps = utils::interval_wraps(from_i, to_i, const bool wraps = utils::interval_wraps(from_i, to_i,
utils::bitlength_of_v<InputT>); utils::bitlength_of_v<InputT>);
const auto segs = utils::split_leaf_nodes(from_node, to_node, stop, wraps); const auto segs = utils::split_leaf_nodes(from_node, to_node, stop, wraps);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return basic_interval_memoizer_at<DpfKey, stop>(segs.total); return basic_interval_memoizer_at<DpfKey, stop>(segs.total);
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <typename DpfKey, std::size_t I, template <typename DpfKey, std::size_t I,

View file

@ -1,7 +1,11 @@
/// @file dpf/json.hpp /// @file dpf/json.hpp
/// @brief nlohmann::json serializers for DPF keys and beaver triples. /// @brief nlohmann::json serializers for DPF keys and beaver triples.
/// @details ADL `adl_serializer` specializations so `nlohmann::json` can /// @details ADL `adl_serializer` specializations so `nlohmann::json` can
/// convert the library's key and triple types. /// convert the library's key and triple types. Works with the
/// vendored nlohmann 3.12 headers and with a newer nlohmann already
/// included by the caller. Classic keys, multi-level `at<>` keys,
/// comparison channels (including payloads wider than 64 bits),
/// wildcard coefficients, and verifiable correction seeds round-trip.
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -12,79 +16,262 @@
#include <cstddef> #include <cstddef>
#include <cstdint> #include <cstdint>
#include <cstring>
#include <stdexcept>
#include <string>
#include <tuple> #include <tuple>
#include <array> #include <array>
#include <string>
#include <bitset>
#include <type_traits> #include <type_traits>
#include <utility> #include <utility>
#if !defined(NLOHMANN_JSON_VERSION_MAJOR)
#include "json/include/nlohmann/json.hpp" #include "json/include/nlohmann/json.hpp"
#endif
#include "portable-snippets/exact-int/exact-int.h" #include "portable-snippets/exact-int/exact-int.h"
#include "dpf/dpf_key.hpp" #include "dpf/dpf_key.hpp"
#include "dpf/secret_share.hpp"
namespace nlohmann namespace dpf
{
namespace json
{
namespace codec
{ {
template <typename NodeT, template <typename T>
typename OutputT> struct is_std_array : std::false_type {};
struct adl_serializer<dpf::beaver<true, NodeT, OutputT>> template <typename T, std::size_t N>
struct is_std_array<std::array<T, N>> : std::true_type {};
template <typename T>
struct is_std_tuple : std::false_type {};
template <typename ...Ts>
struct is_std_tuple<std::tuple<Ts...>> : std::true_type {};
inline std::uint64_t as_u64(const nlohmann::json & j)
{ {
static void from_json(const nlohmann::json & j, dpf::beaver<true, NodeT, OutputT> & beaver) // NOLINT(runtime/references) if (j.is_number_unsigned())
return j.get<std::uint64_t>();
if (j.is_number_integer())
return static_cast<std::uint64_t>(j.get<std::int64_t>());
throw std::invalid_argument("dpf::json: expected an integer");
}
inline std::string hex_encode(const void * data, std::size_t n)
{
static constexpr char digits[] = "0123456789abcdef";
const auto * bytes = static_cast<const unsigned char *>(data);
std::string out(n * 2, '\0');
for (std::size_t i = 0; i < n; ++i)
{ {
j.get_to(beaver.output_blind); out[2 * i] = digits[bytes[i] >> 4];
j.get_to(beaver.vector_blind); out[2 * i + 1] = digits[bytes[i] & 0x0f];
j.get_to(beaver.blinded_vector);
} }
return out;
}
static void to_json(nlohmann::json & j, const dpf::beaver<true, NodeT, OutputT> & beaver) // NOLINT(runtime/references) inline int hex_nybble(char c)
{
if (c >= '0' && c <= '9')
return c - '0';
if (c >= 'a' && c <= 'f')
return c - 'a' + 10;
if (c >= 'A' && c <= 'F')
return c - 'A' + 10;
return -1;
}
template <typename T>
nlohmann::json dump(const T & value);
template <typename T>
T load(const nlohmann::json & j);
template <typename Tuple, std::size_t ...Is>
nlohmann::json dump_tuple(const Tuple & value, std::index_sequence<Is...>)
{
return nlohmann::json::array({dump(std::get<Is>(value))...});
}
template <typename Tuple, std::size_t ...Is>
Tuple load_tuple(const nlohmann::json & j, std::index_sequence<Is...>)
{
return Tuple{load<std::tuple_element_t<Is, Tuple>>(j.at(Is))...};
}
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
template <typename T>
nlohmann::json dump(const T & value)
{
using U = std::remove_cv_t<T>;
if constexpr (std::is_same_v<U, simde__m128i>)
{ {
j = nlohmann::json{ HEDLEY_PRAGMA(GCC diagnostic push)
{"output_blind", beaver.output_blind}, HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
{"vector_blind", beaver.vector_blind}, std::uint64_t lane[2];
{"blinded_vector", beaver.blinded_vector} std::memcpy(lane, &value, sizeof(lane));
HEDLEY_PRAGMA(GCC diagnostic pop)
nlohmann::json out = nlohmann::json::array();
out.push_back(lane[0]);
out.push_back(lane[1]);
return out;
}
else if constexpr (std::is_same_v<U, simde__m256i>)
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::uint64_t lane[4];
std::memcpy(lane, &value, sizeof(lane));
HEDLEY_PRAGMA(GCC diagnostic pop)
nlohmann::json out = nlohmann::json::array();
for (std::uint64_t limb : lane)
out.push_back(limb);
return out;
}
else if constexpr (std::is_same_v<U, simde_uint128>)
{
nlohmann::json out = nlohmann::json::array();
out.push_back(static_cast<std::uint64_t>(value));
out.push_back(static_cast<std::uint64_t>(value >> 64));
return out;
}
else if constexpr (std::is_same_v<U, uint128_t>)
{
nlohmann::json out = nlohmann::json::array();
out.push_back(value.lower());
out.push_back(value.upper());
return out;
}
else if constexpr (std::is_same_v<U, uint256_t>)
{
nlohmann::json out = nlohmann::json::array();
out.push_back(value.lower().lower());
out.push_back(value.lower().upper());
out.push_back(value.upper().lower());
out.push_back(value.upper().upper());
return out;
}
else if constexpr (std::is_same_v<U, dpf::detail::cmp_meta>)
{
nlohmann::json out = nlohmann::json{
{"nbits", value.nbits},
{"mask", value.mask},
{"kind", static_cast<psnip_uint8_t>(value.kind)},
{"trivial", static_cast<psnip_uint8_t>(value.trivial)},
{"eval_as_ge", value.eval_as_ge},
{"include_eq", value.include_eq},
{"active", value.active}
}; };
if (value.incremental)
out["incremental"] = true;
if (value.block_width != 0)
{
out["block_width"] = value.block_width;
out["tail_bits"] = value.tail_bits;
} }
}; return out;
}
else if constexpr (is_std_array<U>::value)
{
nlohmann::json out = nlohmann::json::array();
for (const auto & elem : value)
out.push_back(dump(elem));
return out;
}
else if constexpr (is_std_tuple<U>::value)
{
return dump_tuple(value, std::make_index_sequence<std::tuple_size_v<U>>{});
}
else if constexpr (dpf::is_wildcard_v<U>)
{
return nlohmann::json(nullptr);
}
else if constexpr (std::is_enum_v<U>)
{
return dump(static_cast<std::underlying_type_t<U>>(value));
}
else if constexpr (std::is_same_v<U, bool>)
{
return nlohmann::json(value);
}
else if constexpr (std::is_integral_v<U> && sizeof(U) <= sizeof(std::uint64_t))
{
if constexpr (std::is_signed_v<U>)
return nlohmann::json(static_cast<std::int64_t>(value));
else
return nlohmann::json(static_cast<std::uint64_t>(value));
}
else if constexpr (std::is_same_v<U, float> || std::is_same_v<U, double>)
{
return nlohmann::json(value);
}
else if constexpr (std::is_trivially_copyable_v<U>)
{
nlohmann::json out = nlohmann::json::object();
out["$bytes"] = hex_encode(&value, sizeof(U));
return out;
}
else
{
static_assert(sizeof(U) == 0, "dpf::json: no conversion for this type");
return nlohmann::json(nullptr);
}
}
template <> template <typename T>
struct adl_serializer<simde__m128i> T load(const nlohmann::json & j)
{ {
static void from_json(const nlohmann::json & j, simde__m128i & a) // NOLINT(runtime/references) using U = std::remove_cv_t<T>;
if constexpr (std::is_same_v<U, simde__m128i>)
{ {
std::array<psnip_uint64_t, 2> A; std::uint64_t lane[2] = {as_u64(j.at(0)), as_u64(j.at(1))};
j.get_to(A); HEDLEY_PRAGMA(GCC diagnostic push)
a = simde_mm_set_epi64x(A[1], A[0]); HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
simde__m128i out;
std::memcpy(&out, lane, sizeof(out));
HEDLEY_PRAGMA(GCC diagnostic pop)
return out;
} }
else if constexpr (std::is_same_v<U, simde__m256i>)
static void to_json(nlohmann::json & j, const simde__m128i & a) // NOLINT(runtime/references)
{ {
j = nlohmann::json{a[0], a[1]}; std::uint64_t lane[4] = {
as_u64(j.at(0)), as_u64(j.at(1)), as_u64(j.at(2)), as_u64(j.at(3))
};
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
simde__m256i out;
std::memcpy(&out, lane, sizeof(out));
HEDLEY_PRAGMA(GCC diagnostic pop)
return out;
} }
}; else if constexpr (std::is_same_v<U, simde_uint128>)
template <>
struct adl_serializer<simde__m256i>
{
static void from_json(const nlohmann::json & j, simde__m256i & a) // NOLINT(runtime/references)
{ {
std::array<psnip_uint64_t, 4> A; if (j.is_number())
j.get_to(A); return static_cast<simde_uint128>(as_u64(j));
a = simde_mm256_set_epi64x(A[3], A[2], A[1], A[0]); const simde_uint128 lo = as_u64(j.at(0));
const simde_uint128 hi = as_u64(j.at(1));
return lo | (hi << 64);
} }
else if constexpr (std::is_same_v<U, uint128_t>)
static void to_json(nlohmann::json & j, const simde__m256i & a) // NOLINT(runtime/references)
{ {
j = nlohmann::json{a[0], a[1], a[2], a[3]}; if (j.is_number())
return uint128_t{as_u64(j)};
return uint128_t{as_u64(j.at(1)), as_u64(j.at(0))};
} }
}; else if constexpr (std::is_same_v<U, uint256_t>)
template <>
struct adl_serializer<dpf::detail::cmp_meta>
{
static void from_json(const nlohmann::json & j, dpf::detail::cmp_meta & c) // NOLINT(runtime/references)
{ {
if (j.is_number())
return uint256_t{as_u64(j)};
const uint128_t lo{as_u64(j.at(1)), as_u64(j.at(0))};
const uint128_t hi{as_u64(j.at(3)), as_u64(j.at(2))};
return uint256_t{hi, lo};
}
else if constexpr (std::is_same_v<U, dpf::detail::cmp_meta>)
{
dpf::detail::cmp_meta c;
j.at("nbits").get_to(c.nbits); j.at("nbits").get_to(c.nbits);
j.at("mask").get_to(c.mask); j.at("mask").get_to(c.mask);
c.kind = static_cast<dpf::cmp_kind>(j.at("kind").get<psnip_uint8_t>()); c.kind = static_cast<dpf::cmp_kind>(j.at("kind").get<psnip_uint8_t>());
@ -96,162 +283,89 @@ struct adl_serializer<dpf::detail::cmp_meta>
c.incremental = j.value("incremental", false); c.incremental = j.value("incremental", false);
c.block_width = j.value("block_width", 0); c.block_width = j.value("block_width", 0);
c.tail_bits = j.value("tail_bits", 0); c.tail_bits = j.value("tail_bits", 0);
return c;
} }
else if constexpr (is_std_array<U>::value)
static void to_json(nlohmann::json & j, const dpf::detail::cmp_meta & c) // NOLINT(runtime/references)
{ {
j = nlohmann::json{ U out{};
{"nbits", c.nbits}, if (j.size() != out.size())
{"mask", c.mask}, throw std::invalid_argument("dpf::json: array length mismatch");
{"kind", static_cast<psnip_uint8_t>(c.kind)}, for (std::size_t i = 0; i < out.size(); ++i)
{"trivial", static_cast<psnip_uint8_t>(c.trivial)}, out[i] = load<typename U::value_type>(j.at(i));
{"eval_as_ge", c.eval_as_ge}, return out;
{"include_eq", c.include_eq}, }
{"active", c.active} else if constexpr (is_std_tuple<U>::value)
};
if (c.incremental)
j["incremental"] = true;
if (c.block_width != 0)
{ {
j["block_width"] = c.block_width; if (j.size() != std::tuple_size_v<U>)
j["tail_bits"] = c.tail_bits; throw std::invalid_argument("dpf::json: tuple length mismatch");
return load_tuple<U>(j, std::make_index_sequence<std::tuple_size_v<U>>{});
} }
else if constexpr (dpf::is_wildcard_v<U>)
{
return U{};
} }
}; else if constexpr (std::is_enum_v<U>)
{
return static_cast<U>(load<std::underlying_type_t<U>>(j));
}
else if constexpr (std::is_same_v<U, bool>)
{
if (j.is_boolean())
return j.get<bool>();
return as_u64(j) != 0;
}
else if constexpr (std::is_integral_v<U> && sizeof(U) <= sizeof(std::uint64_t))
{
if constexpr (std::is_signed_v<U>)
return static_cast<U>(static_cast<std::int64_t>(as_u64(j)));
else
return static_cast<U>(as_u64(j));
}
else if constexpr (std::is_same_v<U, float> || std::is_same_v<U, double>)
{
return j.get<U>();
}
else if constexpr (std::is_trivially_copyable_v<U>)
{
if (!j.is_object() || !j.contains("$bytes"))
throw std::invalid_argument("dpf::json: expected a byte blob");
const auto hex = j.at("$bytes").get<std::string>();
if (hex.size() != sizeof(U) * 2)
throw std::invalid_argument("dpf::json: byte blob has the wrong size");
U out{};
auto * bytes = reinterpret_cast<unsigned char *>(&out);
for (std::size_t i = 0; i < sizeof(U); ++i)
{
const int hi = hex_nybble(hex[2 * i]);
const int lo = hex_nybble(hex[2 * i + 1]);
if (hi < 0 || lo < 0)
throw std::invalid_argument("dpf::json: bad hex");
bytes[i] = static_cast<unsigned char>((hi << 4) | lo);
}
return out;
}
else
{
static_assert(sizeof(U) == 0, "dpf::json: no conversion for this type");
return U{};
}
}
HEDLEY_PRAGMA(GCC diagnostic pop)
// Classic single-level key (no `at<>` / no comparison channel). template <typename Word>
template <typename InteriorPRG, std::uint64_t low64(const Word & word)
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs>
struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>,
std::enable_if_t<!dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT,
OutputTs...>::is_multilevel>>
{ {
using dpf_type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>; if constexpr (std::is_same_v<Word, uint256_t>)
using interior_node = typename dpf_type::interior_node; return static_cast<std::uint64_t>(word.lower().lower());
using leaf_tuple = typename dpf_type::leaf_tuple; else if constexpr (std::is_same_v<Word, uint128_t>)
using beaver_tuple = typename dpf_type::beaver_tuple; return word.lower();
else if constexpr (sizeof(Word) <= sizeof(std::uint64_t))
return static_cast<std::uint64_t>(word);
else
return static_cast<std::uint64_t>(word);
}
static dpf_type from_json(const nlohmann::json & j) } // namespace codec
{
interior_node root;
j.at("root").get_to(root);
std::array<interior_node, dpf_type::depth> correction_words;
j.at("correction_words").get_to(correction_words);
std::array<psnip_uint8_t, dpf_type::depth> correction_advice;
j.at("correction_advice").get_to(correction_advice);
leaf_tuple leaves;
j.at("leaves").get_to(leaves);
std::string wildcard_mask_str;
j.at("wildcards").get_to(wildcard_mask_str);
beaver_tuple beavers;
j.at("beavers").get_to(beavers);
return dpf_type{
root,
correction_words,
correction_advice,
leaves,
std::bitset<std::tuple_size_v<leaf_tuple>>(wildcard_mask_str),
beavers
};
}
static void to_json(nlohmann::json & j, const dpf_type & dpf) // NOLINT(runtime/references)
{
j = nlohmann::json{
{"root", dpf.root()},
{"correction_words", dpf.correction_words()},
{"correction_advice", dpf.correction_advice()},
{"leaves", dpf.mutable_leaf_tuple()},
{"wildcards", dpf.mutable_wildcard_mask()},
{"beavers", dpf.mutable_beaver_tuple()}
};
}
};
// Multi-level / comparison key (`at<>` and/or a `cmp` channel). Round-trips
// the public tree (root, CWs, advice) and the comparison channel (cmp meta,
// value CWs, `cw_last`, and this party's `cmp_addend` share). Leaf outputs are
// not yet serialized here, so this path currently supports comparison-only
// keys (`num_outputs == 0`, e.g. `make_dpf(x, dpf::lt(...))`).
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs>
struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>,
std::enable_if_t<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT,
OutputTs...>::is_multilevel>>
{
using dpf_type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>;
using interior_node = typename dpf_type::interior_node;
using input_type = typename dpf_type::input_type;
static dpf_type from_json(const nlohmann::json & j)
{
static_assert(dpf_type::num_outputs == 0,
"dpf::json round-trip currently supports comparison-only "
"multi-level keys (no leaf outputs)");
interior_node root;
j.at("root").get_to(root);
typename dpf_type::correction_words_array correction_words;
j.at("correction_words").get_to(correction_words);
typename dpf_type::correction_advice_array correction_advice;
j.at("correction_advice").get_to(correction_advice);
dpf::detail::cmp_meta cmp;
j.at("cmp").get_to(cmp);
typename dpf_type::value_cw_array value_cws;
j.at("value_cw").get_to(value_cws);
uint64_t cw_last = j.at("cw_last").template get<uint64_t>();
uint64_t cmp_addend = j.at("cmp_addend").template get<uint64_t>();
typename dpf_type::tail_array tail{};
if constexpr (dpf_type::cmp_block > 0)
j.at("tail_cw").get_to(tail);
typename dpf_type::prefix_cw_array prefix{};
if constexpr (dpf_type::cmp_idcf)
j.at("prefix_cw").get_to(prefix);
typename dpf_type::leaf_wrapper_tuple leaves{};
input_type offset_share{};
typename dpf_type::addend_tuple addends{};
return dpf_type{root, correction_words, correction_advice,
std::move(leaves), offset_share, cmp, value_cws, cw_last,
cmp_addend, addends, {}, 0, tail, {}, prefix};
}
static void to_json(nlohmann::json & j, const dpf_type & dpf) // NOLINT(runtime/references)
{
static_assert(dpf_type::num_outputs == 0,
"dpf::json round-trip currently supports comparison-only "
"multi-level keys (no leaf outputs)");
j = nlohmann::json{
{"root", dpf.root()},
{"correction_words", dpf.correction_words()},
{"correction_advice", dpf.correction_advice()},
{"cmp", dpf.cmp()},
{"value_cw", dpf.value_cw()},
{"cw_last", static_cast<uint64_t>(dpf.cw_last())},
{"cmp_addend", static_cast<uint64_t>(dpf.cmp_addend())}
};
if constexpr (dpf_type::cmp_block > 0)
j["tail_cw"] = dpf.tail_cw();
if constexpr (dpf_type::cmp_idcf)
j["prefix_cw"] = dpf.prefix_cws();
}
};
} // namespace nlohmann
namespace dpf
{
namespace json
{
template <typename DpfKey> template <typename DpfKey>
static std::string to_json(const DpfKey & dpf) static std::string to_json(const DpfKey & dpf)
@ -268,7 +382,429 @@ static auto from_json(const std::string & json_string)
} }
} // namespace json } // namespace json
} // namespace dpf } // namespace dpf
NLOHMANN_JSON_NAMESPACE_BEGIN
template <typename NodeT,
typename OutputT>
struct adl_serializer<dpf::beaver<true, NodeT, OutputT>, void>
{
static void from_json(const nlohmann::json & j, dpf::beaver<true, NodeT, OutputT> & beaver) // NOLINT(runtime/references)
{
using beaver_type = dpf::beaver<true, NodeT, OutputT>;
beaver.output_blind = dpf::json::codec::load<OutputT>(j.at("output_blind"));
beaver.vector_blind = dpf::json::codec::load<typename beaver_type::LeafT>(j.at("vector_blind"));
beaver.blinded_vector = dpf::json::codec::load<typename beaver_type::LeafT>(j.at("blinded_vector"));
}
static void to_json(nlohmann::json & j, const dpf::beaver<true, NodeT, OutputT> & beaver) // NOLINT(runtime/references)
{
j = nlohmann::json{
{"output_blind", dpf::json::codec::dump(beaver.output_blind)},
{"vector_blind", dpf::json::codec::dump(beaver.vector_blind)},
{"blinded_vector", dpf::json::codec::dump(beaver.blinded_vector)}
};
}
};
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
template <>
struct adl_serializer<simde__m128i, void>
{
static simde__m128i from_json(const nlohmann::json & j)
{
return dpf::json::codec::load<simde__m128i>(j);
}
static void to_json(nlohmann::json & j, const simde__m128i & a) // NOLINT(runtime/references)
{
j = dpf::json::codec::dump(a);
}
};
template <>
struct adl_serializer<simde__m256i, void>
{
static simde__m256i from_json(const nlohmann::json & j)
{
return dpf::json::codec::load<simde__m256i>(j);
}
static void to_json(nlohmann::json & j, const simde__m256i & a) // NOLINT(runtime/references)
{
j = dpf::json::codec::dump(a);
}
};
HEDLEY_PRAGMA(GCC diagnostic pop)
template <>
struct adl_serializer<simde_uint128, void>
{
static simde_uint128 from_json(const nlohmann::json & j)
{
return dpf::json::codec::load<simde_uint128>(j);
}
static void to_json(nlohmann::json & j, const simde_uint128 & a) // NOLINT(runtime/references)
{
j = dpf::json::codec::dump(a);
}
};
template <>
struct adl_serializer<uint128_t, void>
{
static uint128_t from_json(const nlohmann::json & j)
{
return dpf::json::codec::load<uint128_t>(j);
}
static void to_json(nlohmann::json & j, const uint128_t & a) // NOLINT(runtime/references)
{
j = dpf::json::codec::dump(a);
}
};
template <>
struct adl_serializer<uint256_t, void>
{
static uint256_t from_json(const nlohmann::json & j)
{
return dpf::json::codec::load<uint256_t>(j);
}
static void to_json(nlohmann::json & j, const uint256_t & a) // NOLINT(runtime/references)
{
j = dpf::json::codec::dump(a);
}
};
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs>
struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>, void>
{
using dpf_type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>;
using input_type = typename dpf_type::input_type;
using leaf_tuple = typename dpf_type::leaf_tuple;
using leaf_wrapper_tuple = typename dpf_type::leaf_wrapper_tuple;
static constexpr bool classic =
dpf::detail::incr::is_classic_pack_v<OutputT, OutputTs...>;
template <std::size_t I>
static nlohmann::json dump_leaf(const std::tuple_element_t<I, leaf_wrapper_tuple> & wrapper)
{
nlohmann::json entry = nlohmann::json::object();
entry["leaf"] = dpf::json::codec::dump(wrapper.raw_leaf());
if constexpr (dpf::is_wildcard_v<typename dpf_type::template output_type_t<I>>)
{
const auto & beaver = wrapper.beaver();
entry["beaver"] = nlohmann::json{
{"output_blind", dpf::json::codec::dump(beaver.output_blind)},
{"vector_blind", dpf::json::codec::dump(beaver.vector_blind)},
{"blinded_vector", dpf::json::codec::dump(beaver.blinded_vector)}
};
entry["output_share"] = dpf::json::codec::dump(wrapper.output_share());
entry["state"] = wrapper.state();
}
return entry;
}
template <std::size_t ...Is>
static nlohmann::json dump_leaves(const leaf_wrapper_tuple & leaves,
std::index_sequence<Is...>)
{
return nlohmann::json::array({dump_leaf<Is>(std::get<Is>(leaves))...});
}
template <std::size_t I>
static auto load_beaver(const nlohmann::json & entry)
{
using beaver_type = std::tuple_element_t<I, typename dpf_type::beaver_tuple>;
if constexpr (dpf::is_wildcard_v<typename dpf_type::template output_type_t<I>>)
{
beaver_type beaver{};
const auto & stored = entry.at("beaver");
beaver.output_blind = dpf::json::codec::load<decltype(beaver.output_blind)>(
stored.at("output_blind"));
beaver.vector_blind = dpf::json::codec::load<decltype(beaver.vector_blind)>(
stored.at("vector_blind"));
beaver.blinded_vector = dpf::json::codec::load<decltype(beaver.blinded_vector)>(
stored.at("blinded_vector"));
return beaver;
}
else
{
return beaver_type{};
}
}
template <std::size_t ...Is>
static leaf_tuple load_leaf_tuple(const nlohmann::json & leaves,
std::index_sequence<Is...>)
{
return leaf_tuple{
dpf::json::codec::load<std::tuple_element_t<Is, leaf_tuple>>(
leaves.at(Is).at("leaf"))...
};
}
template <std::size_t ...Is>
static auto load_beaver_tuple(const nlohmann::json & leaves,
std::index_sequence<Is...>)
{
return typename dpf_type::beaver_tuple{load_beaver<Is>(leaves.at(Is))...};
}
template <std::size_t I, typename Wrapper>
static void restore_leaf(Wrapper & wrapper, const nlohmann::json & entry)
{
if constexpr (dpf::is_wildcard_v<typename dpf_type::template output_type_t<I>>)
{
if (!entry.contains("state"))
return;
using output_type = typename Wrapper::output_type;
output_type share{};
if (entry.contains("output_share"))
share = dpf::json::codec::load<output_type>(entry.at("output_share"));
wrapper.restore_state(std::move(share),
entry.at("state").template get<psnip_uint8_t>());
}
else
{
(void)wrapper;
(void)entry;
}
}
template <std::size_t ...Is>
static void restore_leaves(dpf_type & key, const nlohmann::json & leaves,
std::index_sequence<Is...>)
{
(restore_leaf<Is>(std::get<Is>(key.leaf_nodes), leaves.at(Is)), ...);
}
template <std::size_t I>
static auto load_wrapper(const nlohmann::json & entry)
{
using wrapper = std::tuple_element_t<I, leaf_wrapper_tuple>;
auto leaf = dpf::json::codec::load<typename wrapper::leaf_type>(entry.at("leaf"));
if constexpr (dpf::is_wildcard_v<typename dpf_type::template output_type_t<I>>)
{
using beaver_type = typename wrapper::beaver_type;
beaver_type beaver{};
const auto & stored = entry.at("beaver");
beaver.output_blind = dpf::json::codec::load<decltype(beaver.output_blind)>(
stored.at("output_blind"));
beaver.vector_blind = dpf::json::codec::load<decltype(beaver.vector_blind)>(
stored.at("vector_blind"));
beaver.blinded_vector = dpf::json::codec::load<decltype(beaver.blinded_vector)>(
stored.at("blinded_vector"));
wrapper out{std::move(leaf), std::move(beaver)};
restore_leaf<I>(out, entry);
return out;
}
else
{
return wrapper{std::move(leaf)};
}
}
template <std::size_t ...Is>
static leaf_wrapper_tuple load_wrappers(const nlohmann::json & leaves,
std::index_sequence<Is...>)
{
return leaf_wrapper_tuple{load_wrapper<Is>(leaves.at(Is))...};
}
static void restore_offset(dpf_type & key, const nlohmann::json & j)
{
if constexpr (dpf::is_wildcard_v<InputT>)
{
if (j.contains("offset_state"))
{
key.offset_x.restore(
dpf::json::codec::load<input_type>(j.at("offset")),
j.at("offset_state").template get<std::uint8_t>());
}
}
else
{
(void)key;
(void)j;
}
}
static dpf_type from_json(const nlohmann::json & j)
{
const auto root = dpf::json::codec::load<typename dpf_type::interior_node>(j.at("root"));
const auto correction_words =
dpf::json::codec::load<typename dpf_type::correction_words_array>(j.at("correction_words"));
const auto correction_advice =
dpf::json::codec::load<typename dpf_type::correction_advice_array>(j.at("correction_advice"));
input_type offset{};
if (j.contains("offset"))
offset = dpf::json::codec::load<input_type>(j.at("offset"));
if constexpr (classic)
{
constexpr auto idx = std::make_index_sequence<dpf_type::num_outputs>{};
const auto & leaves_json = j.at("leaves");
dpf_type key{root, correction_words, correction_advice,
load_leaf_tuple(leaves_json, idx),
load_beaver_tuple(leaves_json, idx),
offset};
restore_leaves(key, leaves_json, idx);
restore_offset(key, j);
return key;
}
else
{
auto leaves = [&]() {
if constexpr (dpf_type::num_outputs == 0)
return leaf_wrapper_tuple{};
else
return load_wrappers(j.at("leaves"),
std::make_index_sequence<dpf_type::num_outputs>{});
}();
dpf::detail::cmp_meta cmp{};
if (j.contains("cmp"))
cmp = dpf::json::codec::load<dpf::detail::cmp_meta>(j.at("cmp"));
typename dpf_type::value_cw_array value_cws{};
if (j.contains("value_cw"))
value_cws = dpf::json::codec::load<typename dpf_type::value_cw_array>(j.at("value_cw"));
using word = typename dpf_type::value_cw_word;
word cw_last{};
if (j.contains("cw_last"))
cw_last = dpf::json::codec::load<word>(j.at("cw_last"));
word cmp_addend{};
if (j.contains("cmp_addend"))
cmp_addend = dpf::json::codec::load<word>(j.at("cmp_addend"));
typename dpf_type::addend_tuple addends{};
if constexpr (dpf_type::num_outputs > 0)
{
if (j.contains("addends"))
addends = dpf::json::codec::load<typename dpf_type::addend_tuple>(j.at("addends"));
}
typename dpf_type::value_cw_array value_cw_coeff{};
if (j.contains("value_cw_coeff"))
value_cw_coeff = dpf::json::codec::load<typename dpf_type::value_cw_array>(
j.at("value_cw_coeff"));
word cw_last_coeff{};
if (j.contains("cw_last_coeff"))
cw_last_coeff = dpf::json::codec::load<word>(j.at("cw_last_coeff"));
typename dpf_type::tail_array tail{};
typename dpf_type::tail_array tail_coeff{};
if constexpr (dpf_type::cmp_block > 0)
{
if (j.contains("tail_cw"))
tail = dpf::json::codec::load<typename dpf_type::tail_array>(j.at("tail_cw"));
if (j.contains("tail_coeff"))
tail_coeff = dpf::json::codec::load<typename dpf_type::tail_array>(j.at("tail_coeff"));
}
typename dpf_type::prefix_cw_array prefix{};
typename dpf_type::prefix_cw_array prefix_coeff{};
if constexpr (dpf_type::cmp_idcf)
{
if (j.contains("prefix_cw"))
prefix = dpf::json::codec::load<typename dpf_type::prefix_cw_array>(j.at("prefix_cw"));
if (j.contains("prefix_coeff"))
prefix_coeff = dpf::json::codec::load<typename dpf_type::prefix_cw_array>(
j.at("prefix_coeff"));
}
typename dpf_type::correction_seeds_array seeds{};
if constexpr (dpf_type::is_verifiable)
{
if (j.contains("correction_seeds"))
seeds = dpf::json::codec::load<typename dpf_type::correction_seeds_array>(
j.at("correction_seeds"));
}
dpf_type key{root, correction_words, correction_advice,
std::move(leaves), offset, cmp, value_cws,
dpf::json::codec::low64(cw_last), dpf::json::codec::low64(cmp_addend),
std::move(addends), value_cw_coeff, dpf::json::codec::low64(cw_last_coeff),
tail, tail_coeff, prefix, prefix_coeff, seeds};
key.set_cmp_scalars(cw_last, cmp_addend, cw_last_coeff);
if (j.contains("cmp_assigned"))
key.set_cmp_assigned(j.at("cmp_assigned").template get<bool>());
restore_offset(key, j);
return key;
}
}
static void to_json(nlohmann::json & j, const dpf_type & dpf) // NOLINT(runtime/references)
{
j = nlohmann::json::object();
j["root"] = dpf::json::codec::dump(dpf.root());
j["correction_words"] = dpf::json::codec::dump(dpf.correction_words());
j["correction_advice"] = dpf::json::codec::dump(dpf.correction_advice());
if constexpr (dpf_type::num_outputs > 0)
{
j["leaves"] = dump_leaves(dpf.leaf_nodes,
std::make_index_sequence<dpf_type::num_outputs>{});
}
j["offset"] = dpf::json::codec::dump(dpf.offset_x.raw());
if constexpr (dpf::is_wildcard_v<InputT>)
j["offset_state"] = dpf.offset_x.state();
if constexpr (!classic)
{
if constexpr (dpf_type::cmp_depth > 0)
{
j["cmp"] = dpf::json::codec::dump(dpf.cmp());
j["value_cw"] = dpf::json::codec::dump(dpf.value_cw());
j["cw_last"] = dpf::json::codec::dump(dpf.cw_last_word());
j["cmp_addend"] = dpf::json::codec::dump(dpf.cmp_addend_word());
if constexpr (dpf_type::cmp_block > 0)
j["tail_cw"] = dpf::json::codec::dump(dpf.tail_cw());
if constexpr (dpf_type::cmp_idcf)
j["prefix_cw"] = dpf::json::codec::dump(dpf.prefix_cws());
if constexpr (dpf_type::cmp_is_wildcard)
{
j["value_cw_coeff"] = dpf::json::codec::dump(dpf.value_cw_coeff());
j["cw_last_coeff"] = dpf::json::codec::dump(dpf.cw_last_coeff_word());
if constexpr (dpf_type::cmp_block > 0)
j["tail_coeff"] = dpf::json::codec::dump(dpf.tail_coeff());
if constexpr (dpf_type::cmp_idcf)
j["prefix_coeff"] = dpf::json::codec::dump(dpf.prefix_cw_coeff());
j["cmp_assigned"] = dpf.cmp_assigned();
}
}
if constexpr (dpf_type::num_outputs > 0)
j["addends"] = dpf::json::codec::dump(dpf.public_addends);
if constexpr (dpf_type::is_verifiable)
j["correction_seeds"] = dpf::json::codec::dump(dpf.correction_seeds());
}
}
};
template <std::size_t Party, typename Key>
struct adl_serializer<dpf::party_key<Party, Key>, void>
{
static dpf::party_key<Party, Key> from_json(const nlohmann::json & j)
{
return dpf::party_key<Party, Key>(j.template get<Key>());
}
static void to_json(nlohmann::json & j, const dpf::party_key<Party, Key> & key) // NOLINT(runtime/references)
{
j = key.key();
}
};
NLOHMANN_JSON_NAMESPACE_END
#endif // LIBDPF_INCLUDE_DPF_JSON_HPP__ #endif // LIBDPF_INCLUDE_DPF_JSON_HPP__

View file

@ -214,6 +214,7 @@ class basic_fixed_length_string : public dpf::modint<static_cast<std::size_t>(st
/// @} /// @}
/// @brief assign the `basic_fixed_length_string` /// @brief assign the `basic_fixed_length_string`
/// @return `*this`
/// @{ /// @{
/// @brief value assignment /// @brief value assignment
@ -283,8 +284,10 @@ class basic_fixed_length_string : public dpf::modint<static_cast<std::size_t>(st
/// @brief converts a string of length at-most `max_length` over /// @brief converts a string of length at-most `max_length` over
/// `alphabet` into an integer /// `alphabet` into an integer
/// @throws `std::length_error` if `str` exceeds `max_length` /// @param str the source string
/// @throws `std::domain_error` if `str` contains a char not in `alphabet` /// @return the returned `integral_type`
/// @throws std::length_error if `str` exceeds `max_length`
/// @throws std::domain_error if `str` contains a char not in `alphabet`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
static constexpr integral_type encode_(string_view str) static constexpr integral_type encode_(string_view str)
{ {
@ -309,6 +312,8 @@ class basic_fixed_length_string : public dpf::modint<static_cast<std::size_t>(st
/// @brief Index of `c` in `alphabet`, or `npos` when `c` is absent. /// @brief Index of `c` in `alphabet`, or `npos` when `c` is absent.
/// Byte alphabets use a 256-entry table; wider character types scan. /// Byte alphabets use a 256-entry table; wider character types scan.
/// @param c the `c`
/// @return Index of `c` in `alphabet`, or `npos` when `c` is absent
static constexpr std::size_t digit_of_(CharT c) static constexpr std::size_t digit_of_(CharT c)
{ {
constexpr auto missing = string_view::npos; constexpr auto missing = string_view::npos;
@ -340,6 +345,9 @@ class basic_fixed_length_string : public dpf::modint<static_cast<std::size_t>(st
/// @{ /// @{
/// @brief Writes the decoded string, not the packed integer. /// @brief Writes the decoded string, not the packed integer.
/// @param os the character output stream
/// @param k the `k`
/// @return the stream
friend std::basic_ostream<CharT, Traits> & friend std::basic_ostream<CharT, Traits> &
operator<<(std::basic_ostream<CharT, Traits> & os, operator<<(std::basic_ostream<CharT, Traits> & os,
const basic_fixed_length_string & k) const basic_fixed_length_string & k)
@ -348,6 +356,9 @@ class basic_fixed_length_string : public dpf::modint<static_cast<std::size_t>(st
} }
/// @brief Reads a whitespace-delimited token and encodes it. /// @brief Reads a whitespace-delimited token and encodes it.
/// @param is the character input stream
/// @param k the `k`
/// @return the stream
friend std::basic_istream<CharT, Traits> & friend std::basic_istream<CharT, Traits> &
operator>>(std::basic_istream<CharT, Traits> & is, operator>>(std::basic_istream<CharT, Traits> & is,
basic_fixed_length_string & k) basic_fixed_length_string & k)
@ -386,6 +397,13 @@ using keyword = basic_fixed_length_string<MaxLen, char, Alphabet>;
/// @details Uses a `static_cast` to convert `str` to recreate the string /// @details Uses a `static_cast` to convert `str` to recreate the string
/// representation of a `basic_fixed_length_string` /// representation of a `basic_fixed_length_string`
/// @complexity `O(MaxLen)` where `MaxLen` is the maximum string length /// @complexity `O(MaxLen)` where `MaxLen` is the maximum string length
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alphabet alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Allocator allocator type
/// @param str the source string
/// @return the returned `std::basic_string<CharT, Traits, Allocator>`
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alphabet, const CharT * Alphabet,
@ -402,6 +420,11 @@ namespace utils
{ {
/// @brief specializes `dpf::bitlength_of` for `dpf::basic_fixed_length_string` /// @brief specializes `dpf::bitlength_of` for `dpf::basic_fixed_length_string`
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alpha alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Alloc allocator type
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alpha, const CharT * Alpha,
@ -413,6 +436,11 @@ struct bitlength_of<
dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>::bits> { }; dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>::bits> { };
/// @brief specializes `dpf::msb_of` for `dpf::basic_fixed_length_string` /// @brief specializes `dpf::msb_of` for `dpf::basic_fixed_length_string`
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alpha alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Alloc allocator type
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alpha, const CharT * Alpha,
@ -427,6 +455,11 @@ struct msb_of<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc
/// @brief specializes `dpf::countl_zero_symmetric_difference` for /// @brief specializes `dpf::countl_zero_symmetric_difference` for
/// `dpf::basic_fixed_length_string` /// `dpf::basic_fixed_length_string`
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alpha alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Alloc allocator type
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alpha, const CharT * Alpha,
@ -472,6 +505,11 @@ namespace std
/// @{ /// @{
/// @details specializes `std::numeric_limits` for `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>` /// @details specializes `std::numeric_limits` for `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>`
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alpha alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Alloc allocator type
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alpha, const CharT * Alpha,
@ -527,6 +565,11 @@ class numeric_limits<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits
}; };
/// @details specializes `std::numeric_limits` for `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc> const` /// @details specializes `std::numeric_limits` for `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc> const`
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alpha alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Alloc allocator type
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alpha, const CharT * Alpha,
@ -536,7 +579,12 @@ class numeric_limits<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits
: public numeric_limits<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>> {}; : public numeric_limits<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc> volatile` /// @brief `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc> volatile`
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alpha alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Alloc allocator type
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alpha, const CharT * Alpha,
@ -546,7 +594,12 @@ class numeric_limits<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits
: public numeric_limits<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>> {}; : public numeric_limits<dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc> const volatile` /// @brief `dpf::basic_fixed_length_string<MaxLen, CharT, Alpha, Traits, Alloc> const volatile`
/// @tparam MaxLen maximum string length
/// @tparam CharT character type
/// @tparam Alpha alphabet the string is drawn from
/// @tparam Traits character traits
/// @tparam Alloc allocator type
template <std::size_t MaxLen, template <std::size_t MaxLen,
typename CharT, typename CharT,
const CharT * Alpha, const CharT * Alpha,

View file

@ -1740,6 +1740,8 @@ class keyword2
~keyword2() noexcept = default; ~keyword2() noexcept = default;
/// @brief Rank, including values outside the language that fill the bit width. /// @brief Rank, including values outside the language that fill the bit width.
/// @param rank the `rank`
/// @return Rank, including values outside the language that fill the bit width
static constexpr keyword2 from_rank(integral_type rank) noexcept static constexpr keyword2 from_rank(integral_type rank) noexcept
{ {
return keyword2{rank}; return keyword2{rank};
@ -1751,6 +1753,7 @@ class keyword2
} }
/// @brief Bitwise complement of the rank, in the `2^bits` domain. /// @brief Bitwise complement of the rank, in the `2^bits` domain.
/// @return Bitwise complement of the rank, in the `2^bits` domain
constexpr keyword2 operator~() const noexcept constexpr keyword2 operator~() const noexcept
{ {
return keyword2{parent::operator~()}; return keyword2{parent::operator~()};
@ -1801,6 +1804,8 @@ class keyword2
}; };
/// @brief Compile diagnostic for a pattern that should not become a type. /// @brief Compile diagnostic for a pattern that should not become a type.
/// @tparam Pattern null-terminated pattern with static storage
/// @return Compile diagnostic for a pattern that should not become a type
template <const char * Pattern> template <const char * Pattern>
constexpr keyword2_error keyword2_status() noexcept constexpr keyword2_error keyword2_status() noexcept
{ {

View file

@ -1,6 +1,5 @@
/// @file dpf/leaf_arithmetic.hpp /// @file dpf/leaf_arithmetic.hpp
/// @brief /// @brief Addition, subtraction, and multiplication of packed leaves.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -361,6 +360,7 @@ template <typename OutputT, typename NodeT, std::size_t N> struct add_t<OutputT,
template <std::size_t Nbits, typename WordT> struct add_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {}; template <std::size_t Nbits, typename WordT> struct add_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {};
/// @brief Bitwise XOR, not IEEE addition. Float addition does not form an /// @brief Bitwise XOR, not IEEE addition. Float addition does not form an
/// exact secret-sharing group; XOR of the representation does. /// exact secret-sharing group; XOR of the representation does.
/// @tparam NodeT GGM node type
template <typename NodeT> struct add_t<float, NodeT> final : public std::bit_xor<> {}; template <typename NodeT> struct add_t<float, NodeT> final : public std::bit_xor<> {};
template <typename NodeT> struct add_t<double, NodeT> final : public std::bit_xor<> {}; template <typename NodeT> struct add_t<double, NodeT> final : public std::bit_xor<> {};
template <> struct add_t<dpf::bit, void> final : public std::bit_xor<> {}; template <> struct add_t<dpf::bit, void> final : public std::bit_xor<> {};
@ -372,6 +372,7 @@ template <typename T, typename NodeT> struct add_t<xor_wrapper<T>, NodeT> final
/// @brief Integer outputs whose width matches a SIMD lane but whose type is /// @brief Integer outputs whose width matches a SIMD lane but whose type is
/// not one of the explicitly specialized aliases (`char`, `long long`, /// not one of the explicitly specialized aliases (`char`, `long long`,
/// `char16_t`, and so on). /// `char16_t`, and so on).
/// @tparam OutputT output type
template <typename OutputT> template <typename OutputT>
struct add_t<OutputT, simde__m128i, struct add_t<OutputT, simde__m128i,
std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>> std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>>
@ -405,7 +406,6 @@ struct add_t<dpf::modint<N>, simde__m256i>
return add_t<typename dpf::modint<N>::integral_type, simde__m256i>{}(a, b); return add_t<typename dpf::modint<N>::integral_type, simde__m256i>{}(a, b);
} }
}; };
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
namespace detail namespace detail
@ -646,6 +646,7 @@ template <> struct subtract_t<simde_uint128, simde__m256i> final
template <typename OutputT, typename NodeT, std::size_t N> struct subtract_t<OutputT, std::array<NodeT, N>> final : public detail::sub_array_t<OutputT> {}; template <typename OutputT, typename NodeT, std::size_t N> struct subtract_t<OutputT, std::array<NodeT, N>> final : public detail::sub_array_t<OutputT> {};
template <std::size_t Nbits, typename WordT> struct subtract_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {}; template <std::size_t Nbits, typename WordT> struct subtract_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {};
/// @brief Bitwise XOR, not IEEE subtraction. /// @brief Bitwise XOR, not IEEE subtraction.
/// @tparam NodeT GGM node type
template <typename NodeT> struct subtract_t<float, NodeT> final : public std::bit_xor<> {}; template <typename NodeT> struct subtract_t<float, NodeT> final : public std::bit_xor<> {};
template <typename NodeT> struct subtract_t<double, NodeT> final : public std::bit_xor<> {}; template <typename NodeT> struct subtract_t<double, NodeT> final : public std::bit_xor<> {};
template <typename NodeT> struct subtract_t<dpf::bit, NodeT> final : public std::bit_xor<> {}; template <typename NodeT> struct subtract_t<dpf::bit, NodeT> final : public std::bit_xor<> {};
@ -687,7 +688,6 @@ struct subtract_t<dpf::modint<N>, simde__m256i>
return subtract_t<typename dpf::modint<N>::integral_type, simde__m256i>{}(a, b); return subtract_t<typename dpf::modint<N>::integral_type, simde__m256i>{}(a, b);
} }
}; };
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
namespace detail namespace detail
@ -806,7 +806,6 @@ struct mul4x64_t
static_cast<int64_t>(a[3]*b)}; static_cast<int64_t>(a[3]*b)};
} }
}; };
HEDLEY_PRAGMA(GCC diagnostic pop)
} // namespace detail } // namespace detail
@ -1279,8 +1278,8 @@ struct multiply_t<dpf::nyble, simde__m256i> final
return dpf::lane_arith::mul_epi4(a, b); return dpf::lane_arith::mul_epi4(a, b);
} }
}; };
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
} // namespace leaf_arithmetic } // namespace leaf_arithmetic
} // namespace dpf } // namespace dpf

View file

@ -1,6 +1,5 @@
/// @file dpf/leaf_node.hpp /// @file dpf/leaf_node.hpp
/// @brief /// @brief The packed leaf image of one output group.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @author Christopher Jiang <christopher.jiang@ucalgary.ca> /// @author Christopher Jiang <christopher.jiang@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
@ -136,10 +135,13 @@ struct const_max_size<Only>
static constexpr std::size_t value = Only; static constexpr std::size_t value = Only;
}; };
/// PRG position span covering output indices `Is...` of `OutputsTuple`. /// @brief PRG position span covering output indices `Is...` of `OutputsTuple`.
/// `is_contiguous` is true when the selected outputs occupy a hole-free /// @details `is_contiguous` is true when the selected outputs occupy a hole-free
/// range, so one `ExteriorPRG::eval(..., count, pos_min)` produces every /// range, so one `ExteriorPRG::eval(..., count, pos_min)` produces every
/// leaf mask. /// leaf mask.
/// @tparam NodeT GGM node type
/// @tparam OutputsTuple outputs tuple
/// @tparam Is is
template <typename NodeT, template <typename NodeT,
typename OutputsTuple, typename OutputsTuple,
std::size_t ...Is> std::size_t ...Is>
@ -268,8 +270,12 @@ auto make_naked_leaf(InputT x, OutputT y) noexcept
return Y; return Y;
} }
/// Address of the first `NodeT` block inside a leaf. /// @brief Address of the first `NodeT` block inside a leaf.
/// A one-block leaf *is* a `NodeT`; a longer leaf is `std::array<NodeT, N>`. /// @details A one-block leaf *is* a `NodeT`; a longer leaf is `std::array<NodeT, N>`.
/// @tparam NodeT GGM node type
/// @tparam LeafT leaf type
/// @param leaf the leaf value
/// @return Address of the first `NodeT` block inside a leaf
template <typename NodeT, typename LeafT> template <typename NodeT, typename LeafT>
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -314,6 +320,7 @@ auto make_leaf_mask(const InteriorBlock & seed0, const InteriorBlock & seed1,
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using output_type = concrete_type_t<std::tuple_element_t<I, OutputsTuple>>; using output_type = concrete_type_t<std::tuple_element_t<I, OutputsTuple>>;
HEDLEY_PRAGMA(GCC diagnostic pop)
auto mask0 = make_leaf_mask_inner<ExteriorPRG, I, OutputsTuple, InteriorBlock>( auto mask0 = make_leaf_mask_inner<ExteriorPRG, I, OutputsTuple, InteriorBlock>(
seed0, pos_base); seed0, pos_base);
@ -321,7 +328,6 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
seed1, pos_base); seed1, pos_base);
return dpf::subtract_leaf<output_type>(mask1, mask0); return dpf::subtract_leaf<output_type>(mask1, mask0);
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <typename ExteriorPRG, template <typename ExteriorPRG,
@ -340,6 +346,7 @@ auto make_leaf(InputT x, const ExteriorBlock & seed0, const ExteriorBlock & seed
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using node_type = typename ExteriorPRG::block_type; using node_type = typename ExteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic pop)
return sign ? dpf::subtract_leaf<output_type>( return sign ? dpf::subtract_leaf<output_type>(
make_naked_leaf<node_type>(x, Y), make_naked_leaf<node_type>(x, Y),
@ -349,7 +356,6 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
make_leaf_mask<ExteriorPRG, I, output_tuple_type, ExteriorBlock>( make_leaf_mask<ExteriorPRG, I, output_tuple_type, ExteriorBlock>(
seed0, seed1, pos_base), seed0, seed1, pos_base),
make_naked_leaf<node_type>(x, Y)); make_naked_leaf<node_type>(x, Y));
HEDLEY_PRAGMA(GCC diagnostic pop)
} }
template <typename ExteriorPRG, template <typename ExteriorPRG,
@ -475,9 +481,9 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
}, return_tuple.first.first); }, return_tuple.first.first);
}, leaves); }, leaves);
}, std::make_tuple(y, ys...)); }, std::make_tuple(y, ys...));
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
return return_tuple; return return_tuple;
} }

View file

@ -1,6 +1,5 @@
/// @file dpf/leaf_wrapper.hpp /// @file dpf/leaf_wrapper.hpp
/// @brief /// @brief A concrete or wildcard leaf and the slot for its Beaver triple.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -167,7 +166,11 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
return blinded_output_share; return blinded_output_share;
} }
/// Accept a party-tagged share; convert to additive before Beaver math. /// @brief Accept a party-tagged share; convert to additive before Beaver math.
/// @tparam Party party index, `0` or `1`
/// @tparam Scheme scheme
/// @param output_share the `output_share`
/// @return the returned `const output_type`
template <std::size_t Party, sharing Scheme> template <std::size_t Party, sharing Scheme>
const output_type compute_and_get_blinded_output_share( const output_type compute_and_get_blinded_output_share(
const secret_share<output_type, Party, Scheme> & output_share) const secret_share<output_type, Party, Scheme> & output_share)
@ -212,6 +215,32 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
HEDLEY_NO_THROW HEDLEY_NO_THROW
const beaver_type & beaver() const noexcept { return beaver_; } const beaver_type & beaver() const noexcept { return beaver_; }
/// @brief Output share captured during Beaver blinding, if any.
/// @return the output share
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
const output_type & output_share() const noexcept { return output_share_; }
/// @brief `leaf_status` as a byte, for serialization.
/// @return the status byte
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
constexpr std::uint8_t state() const noexcept
{
return static_cast<std::uint8_t>(leaf_state_);
}
/// @brief Restore a serialized blinding share and status byte.
/// @param share the output share
/// @param state the status byte
HEDLEY_ALWAYS_INLINE
void restore_state(output_type share, std::uint8_t state) noexcept
{
output_share_ = std::move(share);
leaf_state_ = static_cast<leaf_status>(state);
}
private: private:
enum class leaf_status : psnip_uint8_t { ready = 0, waiting = 1, computing = 2, blinded = 3, notset = 4 }; enum class leaf_status : psnip_uint8_t { ready = 0, waiting = 1, computing = 2, blinded = 3, notset = 4 };

View file

@ -13,6 +13,7 @@
#define LIBDPF_INCLUDE_DPF_MODINT_HPP__ #define LIBDPF_INCLUDE_DPF_MODINT_HPP__
#include <cstddef> #include <cstddef>
#include <cstring>
#include <cmath> #include <cmath>
#include <type_traits> #include <type_traits>
#include <functional> #include <functional>
@ -32,6 +33,7 @@ namespace dpf
{ {
/// @brief represents an unsigned integer modulo `2^Nbits` for small values of `Nbits` /// @brief represents an unsigned integer modulo `2^Nbits` for small values of `Nbits`
/// @tparam Nbits width in bits
template <std::size_t Nbits> template <std::size_t Nbits>
class modint class modint
{ {
@ -40,6 +42,7 @@ class modint
using integral_type = dpf::utils::nonvoid_integral_type_from_bitlength_t<Nbits>; using integral_type = dpf::utils::nonvoid_integral_type_from_bitlength_t<Nbits>;
static constexpr std::size_t num_bits = Nbits; static constexpr std::size_t num_bits = Nbits;
static constexpr bool dpf_modint = true;
/// @brief construct the `modint` /// @brief construct the `modint`
/// @{ /// @{
@ -78,6 +81,7 @@ class modint
/// @} /// @}
/// @brief assign the `modint` /// @brief assign the `modint`
/// @return `*this`
/// @{ /// @{
/// @brief value assignment /// @brief value assignment
@ -110,10 +114,11 @@ class modint
~modint() = default; ~modint() = default;
/// @brief addition operator /// @brief addition operator
/// @param rhs the other addend
/// @return the sum
/// @{ /// @{
/// @details Performs addition with an `integral_type`. /// @details Performs addition with an `integral_type`.
/// @param rhs the other addend
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -123,7 +128,6 @@ class modint
} }
/// @details Performs addition with another `modint`. /// @details Performs addition with another `modint`.
/// @param rhs the other addend
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -135,10 +139,11 @@ class modint
/// @} /// @}
/// @brief addition-assignment operator /// @brief addition-assignment operator
/// @param rhs the other addend
/// @return `*this`
/// @{ /// @{
/// @details Adds an `integral_type` to this `modint`. /// @details Adds an `integral_type` to this `modint`.
/// @param rhs the other addend
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator+=(integral_type rhs) noexcept constexpr modint & operator+=(integral_type rhs) noexcept
@ -148,7 +153,6 @@ class modint
} }
/// @details Adds another `modint` to this one. /// @details Adds another `modint` to this one.
/// @param rhs the other addend
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator+=(modint rhs) noexcept constexpr modint & operator+=(modint rhs) noexcept
@ -164,6 +168,7 @@ class modint
/// @brief pre-increment operator /// @brief pre-increment operator
/// @details Increments this `modint` and returns a reference to the /// @details Increments this `modint` and returns a reference to the
/// result. /// result.
/// @return `*this`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator++() noexcept constexpr modint & operator++() noexcept
@ -174,6 +179,7 @@ class modint
/// @brief post-increment operator /// @brief post-increment operator
/// @details Creates a copy of this `modint`, and then increments this /// @details Creates a copy of this `modint`, and then increments this
/// `modint` and returns the copy from before the increment. /// `modint` and returns the copy from before the increment.
/// @return `*this`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint operator++(int) noexcept constexpr modint operator++(int) noexcept
@ -189,6 +195,7 @@ class modint
/// @details Returns the additive inverse modulo `2^Nbits` (two's /// @details Returns the additive inverse modulo `2^Nbits` (two's
/// complement on the underlying word). Required by /// complement on the underlying word). Required by
/// `grotto::for_each_offset`, which computes `-offset`. /// `grotto::for_each_offset`, which computes `-offset`.
/// @return unary negation
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -198,10 +205,11 @@ class modint
} }
/// @brief subtraction operator /// @brief subtraction operator
/// @param rhs the subtrahend
/// @return the difference
/// @{ /// @{
/// @details Performs subtraction by an `integral_type`. /// @details Performs subtraction by an `integral_type`.
/// @param rhs the subtrahend
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -211,7 +219,6 @@ class modint
} }
/// @details Performs subtraction by another `modint`. /// @details Performs subtraction by another `modint`.
/// @param rhs the subtrahend
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -223,10 +230,11 @@ class modint
/// @} /// @}
/// @brief subtraction-assignment operator /// @brief subtraction-assignment operator
/// @param rhs the subtrahend
/// @return `*this`
/// @{ /// @{
/// @details Subtracts an `integral_type` from this `modint`. /// @details Subtracts an `integral_type` from this `modint`.
/// @param rhs the subtrahend
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator-=(integral_type rhs) noexcept constexpr modint & operator-=(integral_type rhs) noexcept
@ -236,7 +244,6 @@ class modint
} }
/// @details Subtracts another `modint` from this one. /// @details Subtracts another `modint` from this one.
/// @param rhs the subtrahend
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator-=(modint rhs) noexcept constexpr modint & operator-=(modint rhs) noexcept
@ -252,6 +259,7 @@ class modint
/// @brief pre-decrement operator /// @brief pre-decrement operator
/// @details Decrements this `modint` and returns a reference to the /// @details Decrements this `modint` and returns a reference to the
/// result. /// result.
/// @return `*this`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator--() noexcept constexpr modint & operator--() noexcept
@ -262,6 +270,7 @@ class modint
/// @brief post-decrement operator /// @brief post-decrement operator
/// @details Creates a copy of this `modint`, and then decrements this /// @details Creates a copy of this `modint`, and then decrements this
/// `modint` and returns the copy from before the decrement. /// `modint` and returns the copy from before the decrement.
/// @return `*this`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint operator--(int) noexcept constexpr modint operator--(int) noexcept
@ -278,6 +287,7 @@ class modint
/// one by `shift_amount` bits to the left. The value of `a<<b` /// one by `shift_amount` bits to the left. The value of `a<<b`
/// is therefore a `modint` congruent to `a * 2^b` modulo `2^Nbits`. /// is therefore a `modint` congruent to `a * 2^b` modulo `2^Nbits`.
/// @param shift_amount the number of bits to shift by /// @param shift_amount the number of bits to shift by
/// @return bitwise-left-shift operator
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -295,6 +305,7 @@ class modint
/// and returns a reference to the result. Upon invoking `a<<=b`, /// and returns a reference to the result. Upon invoking `a<<=b`,
/// `a` is congruent to `a * 2^b` modulo `2^Nbits`. /// `a` is congruent to `a * 2^b` modulo `2^Nbits`.
/// @param shift_amount the number of bits to shift by /// @param shift_amount the number of bits to shift by
/// @return `*this`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator<<=(std::size_t shift_amount) noexcept constexpr modint & operator<<=(std::size_t shift_amount) noexcept
@ -311,6 +322,7 @@ class modint
/// one by `shift_amount` bits to the right. The value of `a>>b` /// one by `shift_amount` bits to the right. The value of `a>>b`
/// is therefore a `modint` equal to the integer part of `a/2^b`. /// is therefore a `modint` equal to the integer part of `a/2^b`.
/// @param shift_amount the number of bits to shift by /// @param shift_amount the number of bits to shift by
/// @return bitwise-right-shift operator
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -338,6 +350,8 @@ class modint
} }
/// @brief Integer division of the reduced values. /// @brief Integer division of the reduced values.
/// @param rhs the right-hand operand
/// @return Integer division of the reduced values
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -364,10 +378,11 @@ class modint
} }
/// @brief multiplication operator /// @brief multiplication operator
/// @param rhs the other multiplicand
/// @return the product
/// @{ /// @{
/// @brief Multiplies this `modint` with an `integral_type`. /// @brief Multiplies this `modint` with an `integral_type`.
/// @param rhs the other multiplicand
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -377,7 +392,6 @@ class modint
} }
/// @brief Multiplies another `modint` with this one. /// @brief Multiplies another `modint` with this one.
/// @param rhs the other multiplicand
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -389,10 +403,11 @@ class modint
/// @} /// @}
/// @brief multiplication-assignment operator /// @brief multiplication-assignment operator
/// @param rhs the other multiplicand
/// @return `*this`
/// @{ /// @{
/// @details Multiplies an `integral_type` into this `modint`. /// @details Multiplies an `integral_type` into this `modint`.
/// @param rhs the other multiplicand
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator*=(integral_type rhs) noexcept constexpr modint & operator*=(integral_type rhs) noexcept
@ -402,7 +417,6 @@ class modint
} }
/// @details Multiplies another `modint` into this one. /// @details Multiplies another `modint` into this one.
/// @param rhs the other multiplicand
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
constexpr modint & operator*=(modint rhs) noexcept constexpr modint & operator*=(modint rhs) noexcept
@ -532,6 +546,7 @@ class modint
} }
/// @brief convert this `modint` to the equivalent `integeral_type` /// @brief convert this `modint` to the equivalent `integeral_type`
/// @return the returned `operator`
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -572,7 +587,44 @@ class modint
operator<<(std::basic_ostream<CharT, Traits> & os, operator<<(std::basic_ostream<CharT, Traits> & os,
const modint & i) const modint & i)
{ {
return os << i.reduced_value(); const auto raw = i.reduced_value();
using word = std::remove_cv_t<std::decay_t<decltype(raw)>>;
if constexpr (std::is_same_v<word, unsigned __int128>
|| std::is_same_v<word, __int128>)
{
if (raw == 0)
return os << CharT('0');
CharT buf[40];
int n = 0;
auto v = raw;
while (v != 0)
{
buf[n++] = static_cast<CharT>('0' + static_cast<int>(v % 10));
v /= 10;
}
while (n > 0)
os << buf[--n];
return os;
}
else if constexpr (std::is_integral_v<word>)
return os << raw;
else
{
unsigned char bytes[sizeof(word)];
std::memcpy(bytes, &raw, sizeof(word));
os << "0x";
bool started = false;
constexpr char hex[] = "0123456789abcdef";
for (int b = static_cast<int>(sizeof(word)) - 1; b >= 0; --b)
{
if (!started && bytes[static_cast<std::size_t>(b)] == 0 && b != 0)
continue;
started = true;
const auto byte = bytes[static_cast<std::size_t>(b)];
os << hex[byte >> 4] << hex[byte & 0x0f];
}
return os;
}
} }
template <typename CharT, template <typename CharT,
@ -615,8 +667,10 @@ class modint
}; };
/// @brief Multiplies a `modint<Nbits>` with an `modint::integral_type`. /// @brief Multiplies a `modint<Nbits>` with an `modint::integral_type`.
/// @tparam Nbits width in bits
/// @param lhs the `integral_type` multiplicand /// @param lhs the `integral_type` multiplicand
/// @param rhs the `modint` multiplicand /// @param rhs the `modint` multiplicand
/// @return Multiplies a `modint<Nbits>` with an `modint::integral_type`
template <std::size_t Nbits> template <std::size_t Nbits>
HEDLEY_CONST HEDLEY_CONST
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -628,9 +682,13 @@ constexpr modint<Nbits> operator*(typename modint<Nbits>::integral_type lhs,
} }
/// @brief Compare two `modint<Nbits>`s as if they were regular integers /// @brief Compare two `modint<Nbits>`s as if they were regular integers
/// @tparam Nbits width in bits
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @{ /// @{
/// @brief less-than operator /// @brief less-than operator
/// @return `true` when `lhs < rhs`
template <std::size_t Nbits> template <std::size_t Nbits>
HEDLEY_CONST HEDLEY_CONST
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -642,6 +700,7 @@ constexpr bool operator<(modint<Nbits> lhs, modint<Nbits> rhs) noexcept
} }
/// @brief less-than-or-equal-to operator /// @brief less-than-or-equal-to operator
/// @return `true` when `lhs <= rhs`
template <std::size_t Nbits> template <std::size_t Nbits>
HEDLEY_CONST HEDLEY_CONST
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -653,6 +712,7 @@ constexpr bool operator<=(modint<Nbits> lhs, modint<Nbits> rhs) noexcept
} }
/// @brief greater-than operator /// @brief greater-than operator
/// @return `true` when `lhs > rhs`
template <std::size_t Nbits> template <std::size_t Nbits>
HEDLEY_CONST HEDLEY_CONST
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -664,6 +724,7 @@ constexpr bool operator>(modint<Nbits> lhs, modint<Nbits> rhs) noexcept
} }
/// @brief greater-than-or-equal-to operator /// @brief greater-than-or-equal-to operator
/// @return `true` when `lhs >= rhs`
template <std::size_t Nbits> template <std::size_t Nbits>
HEDLEY_CONST HEDLEY_CONST
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -675,6 +736,7 @@ constexpr bool operator>=(modint<Nbits> lhs, modint<Nbits> rhs) noexcept
} }
/// @brief equality operator /// @brief equality operator
/// @return `true` when `lhs == rhs`
template <std::size_t Nbits> template <std::size_t Nbits>
HEDLEY_CONST HEDLEY_CONST
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -686,6 +748,7 @@ constexpr bool operator==(modint<Nbits> lhs, modint<Nbits> rhs) noexcept
} }
/// @brief inequality operator /// @brief inequality operator
/// @return `true` when `lhs != rhs`
template <std::size_t Nbits> template <std::size_t Nbits>
HEDLEY_CONST HEDLEY_CONST
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -1358,6 +1421,7 @@ namespace std
/// @{ /// @{
/// @details specializes `std::numeric_limits` for `dpf::modint<Nbits>` /// @details specializes `std::numeric_limits` for `dpf::modint<Nbits>`
/// @tparam Nbits width in bits
template<std::size_t Nbits> template<std::size_t Nbits>
class numeric_limits<dpf::modint<Nbits>> class numeric_limits<dpf::modint<Nbits>>
{ {
@ -1409,18 +1473,21 @@ class numeric_limits<dpf::modint<Nbits>>
}; };
/// @details specializes `std::numeric_limits` for `dpf::modint<Nbits> const` /// @details specializes `std::numeric_limits` for `dpf::modint<Nbits> const`
/// @tparam Nbits width in bits
template<std::size_t Nbits> template<std::size_t Nbits>
class numeric_limits<dpf::modint<Nbits> const> class numeric_limits<dpf::modint<Nbits> const>
: public numeric_limits<dpf::modint<Nbits>> {}; : public numeric_limits<dpf::modint<Nbits>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `dpf::modint<Nbits> volatile` /// @brief `dpf::modint<Nbits> volatile`
/// @tparam Nbits width in bits
template<std::size_t Nbits> template<std::size_t Nbits>
class numeric_limits<dpf::modint<Nbits> volatile> class numeric_limits<dpf::modint<Nbits> volatile>
: public numeric_limits<dpf::modint<Nbits>> {}; : public numeric_limits<dpf::modint<Nbits>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `dpf::modint<Nbits> const volatile` /// @brief `dpf::modint<Nbits> const volatile`
/// @tparam Nbits width in bits
template<std::size_t Nbits> template<std::size_t Nbits>
class numeric_limits<dpf::modint<Nbits> const volatile> class numeric_limits<dpf::modint<Nbits> const volatile>
: public numeric_limits<dpf::modint<Nbits>> {}; : public numeric_limits<dpf::modint<Nbits>> {};

560
include/dpf/multipoint.hpp Normal file
View file

@ -0,0 +1,560 @@
/// @file dpf/multipoint.hpp
/// @brief Cuckoo-packed multi-point DPF and verifiable multi-point DPF.
/// @details Packs t distinct points into m ≈ O(t) buckets (de Castro–
/// Polychroniadou, EUROCRYPT 2022, §4). Each bucket is an ordinary
/// point key on a smaller domain — `dpf::verifiable` selects VDPF
/// buckets. Evaluation probes κ = 3 buckets and sums the shares.
/// A batched proof is one 2λ token.
/// @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_MULTIPOINT_HPP__
#define LIBDPF_INCLUDE_DPF_MULTIPOINT_HPP__
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <iterator>
#include <limits>
#include <random>
#include <stdexcept>
#include <type_traits>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "dpf/eval_point.hpp"
#include "dpf/incremental.hpp"
#include "dpf/prg_aes.hpp"
#include "dpf/random.hpp"
#include "dpf/secret_share.hpp"
#include "dpf/verifiable.hpp"
namespace dpf
{
/// @brief Knobs for cuckoo packing. `lambda` is the Remark 1 failure target.
struct multipoint_params
{
std::uint32_t lambda = 40;
std::uint32_t max_evictions = 4096;
int retries = 8;
};
template <typename T>
struct is_multipoint_key : std::false_type
{
};
template <std::size_t Party,
typename InputT,
typename OutputT,
typename BucketKey>
struct multipoint_key
{
static constexpr std::size_t party = Party;
static constexpr bool is_multipoint = true;
static constexpr bool is_verifiable = BucketKey::is_verifiable;
static constexpr std::size_t kappa = 3;
using input_type = InputT;
using output_type = OutputT;
using bucket_key = BucketKey;
using bucket_input = typename BucketKey::input_type;
using share_type = subtractive_share<OutputT, Party>;
simde__m128i sigma{};
std::uint32_t bucket_count = 0;
std::uint64_t bucket_domain = 0;
std::vector<party_key<Party, BucketKey>> buckets{};
};
template <std::size_t Party, typename InputT, typename OutputT, typename BucketKey>
struct is_multipoint_key<multipoint_key<Party, InputT, OutputT, BucketKey>>
: std::true_type
{
};
template <typename T>
inline constexpr bool is_multipoint_key_v =
is_multipoint_key<std::decay_t<T>>::value;
namespace detail
{
namespace mpf
{
struct prp_walk_error : std::runtime_error
{
prp_walk_error()
: std::runtime_error("multipoint PRP cycle walk exceeded its bound")
{
}
};
struct located
{
std::uint32_t bucket = 0;
std::uint64_t index = 0;
};
using wide = unsigned __int128;
inline wide domain_size(std::size_t bits)
{
return wide{1} << bits;
}
/// @brief 4-round Feistel on the next power-of-two square, then cycle-walk
/// into `[0, domain)`. AES-MMO is the round function.
/// @param seed the PRP seed
/// @param x the input, in `[0, domain)`
/// @param domain the domain size
/// @return the permuted value in `[0, domain)`
/// @throws std::invalid_argument if `x` is outside the domain
/// @throws prp_walk_error if the cycle walk exceeds its bound
inline wide permute(simde__m128i seed, wide x, wide domain)
{
if (domain <= 1)
return 0;
if (x >= domain)
throw std::invalid_argument("multipoint PRP input is outside the domain");
int bits = 0;
for (wide v = domain - 1; v > 0; v >>= 1)
++bits;
const int half = (bits + 1) / 2;
const wide mask = (half >= 128)
? ~wide{0}
: (wide{1} << half) - 1;
wide val = x;
for (int guard = 0; guard < 128; ++guard)
{
unsigned __int128 left = (val >> half) & mask;
unsigned __int128 right = val & mask;
for (int round = 0; round < 4; ++round)
{
alignas(16) std::uint64_t lanes[2] = {
static_cast<std::uint64_t>(right),
static_cast<std::uint64_t>(right >> 64)};
auto msg = simde_mm_load_si128(
reinterpret_cast<const simde__m128i *>(lanes));
msg = simde_mm_xor_si128(msg, seed);
msg = simde_mm_xor_si128(msg,
simde_mm_set_epi32(0, 0, 0, round + 1));
const auto out = prg::aes128::eval(msg,
static_cast<psnip_uint32_t>(round + 1));
simde_mm_store_si128(reinterpret_cast<simde__m128i *>(lanes), out);
wide f = lanes[0] | (wide{lanes[1]} << 64);
f &= mask;
left ^= f;
const wide tmp = left;
left = right;
right = tmp;
}
val = (left << half) | right;
if (val < domain)
return val;
}
throw prp_walk_error{};
}
inline located locate(simde__m128i sigma, wide x, int hash,
wide n, wide bucket_domain)
{
constexpr int kappa = 3;
const wide y = permute(sigma,
x + n * static_cast<unsigned>(hash), n * kappa);
located out;
out.bucket = static_cast<std::uint32_t>(y / bucket_domain);
out.index = static_cast<std::uint64_t>(y % bucket_domain);
return out;
}
inline std::uint32_t bucket_count_for(std::uint32_t t, std::uint32_t lambda)
{
const double log2t = (t <= 1) ? 0.0 : std::log2(static_cast<double>(t));
const double e = (static_cast<double>(lambda) + 130.0 + log2t) / 123.5;
auto m = static_cast<std::uint32_t>(std::ceil(e * static_cast<double>(t)));
if (m < t + 1)
m = t + 1;
// Remark 1's simplification wants t ≥ 30. Below that, keep a 2t table.
if (t < 30 && m < t * 2)
m = t * 2;
return m;
}
inline std::uint32_t rng_seed(simde__m128i sigma)
{
const auto block = prg::aes128::eval(sigma, 0xC000u);
alignas(16) std::uint32_t words[4];
simde_mm_store_si128(reinterpret_cast<simde__m128i *>(words), block);
return words[0] ^ (words[1] * 0x9E3779B9u) ^ words[2] ^ words[3];
}
struct slot
{
int item = -1;
int hash = -1;
};
template <typename InputT>
bool insert_cuckoo(simde__m128i sigma, const std::vector<InputT> & alphas,
std::uint32_t m, wide n, wide bucket_domain,
std::uint32_t max_evictions, std::vector<slot> & table)
{
table.assign(m, slot{});
std::mt19937 rng(rng_seed(sigma));
std::uniform_int_distribution<int> pick(0, 2);
const int t = static_cast<int>(alphas.size());
for (int omega = 0; omega < t; ++omega)
{
int cur = omega;
int hash = pick(rng);
std::uint32_t evictions = 0;
for (;;)
{
const auto loc = locate(sigma,
static_cast<wide>(alphas[static_cast<std::size_t>(cur)]),
hash, n, bucket_domain);
if (loc.bucket >= m)
return false;
if (table[loc.bucket].item < 0)
{
table[loc.bucket] = slot{cur, hash};
break;
}
const int evicted = table[loc.bucket].item;
table[loc.bucket] = slot{cur, hash};
cur = evicted;
hash = pick(rng);
if (++evictions > max_evictions)
return false;
}
}
return true;
}
template <bool Verifiable,
typename InteriorPRG,
typename ExteriorPRG,
typename BucketInput,
typename OutputT>
auto make_bucket(BucketInput index, const OutputT & beta)
{
if constexpr (Verifiable)
{
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(index, beta,
dpf::verifiable{});
}
else
{
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(index, beta);
}
}
template <bool Verifiable, typename InteriorPRG, typename ExteriorPRG,
typename BucketInput, typename OutputT>
struct bucket_bare
{
using type = typename decltype(make_bucket<Verifiable, InteriorPRG,
ExteriorPRG>(std::declval<BucketInput>(),
std::declval<const OutputT &>()).first)::key_type;
};
template <bool Verifiable,
typename InteriorPRG,
typename ExteriorPRG,
typename BucketInput,
typename InputT,
typename OutputT>
auto make_impl(std::vector<InputT> alphas, std::vector<OutputT> betas,
multipoint_params params)
{
using bare = typename bucket_bare<Verifiable, InteriorPRG, ExteriorPRG,
BucketInput, OutputT>::type;
using key0 = multipoint_key<0, InputT, OutputT, bare>;
using key1 = multipoint_key<1, InputT, OutputT, bare>;
static_assert(std::is_unsigned_v<InputT> && !std::is_same_v<InputT, bool>,
"make_multipoint: input domain must be an unsigned integer");
static_assert(utils::bitlength_of_v<InputT> <= 32,
"make_multipoint: input domain wider than 32 bits is not supported");
static_assert(std::is_unsigned_v<BucketInput>
&& !std::is_same_v<BucketInput, bool>,
"make_multipoint: BucketInput must be an unsigned integer");
if (alphas.size() != betas.size())
throw std::invalid_argument("make_multipoint: point and payload counts differ");
if (alphas.empty())
throw std::invalid_argument("make_multipoint: no points");
if (alphas.size() > static_cast<std::size_t>(std::numeric_limits<std::uint32_t>::max()))
throw std::invalid_argument("make_multipoint: too many points");
{
auto sorted = alphas;
std::sort(sorted.begin(), sorted.end());
if (std::adjacent_find(sorted.begin(), sorted.end()) != sorted.end())
throw std::invalid_argument("make_multipoint: duplicate points");
}
const auto t = static_cast<std::uint32_t>(alphas.size());
const auto m = bucket_count_for(t, params.lambda);
constexpr std::size_t input_bits = utils::bitlength_of_v<InputT>;
const wide n = domain_size(input_bits);
constexpr int kappa = 3;
const wide b = (n * kappa + m - 1) / m;
constexpr std::size_t bucket_bits = utils::bitlength_of_v<BucketInput>;
const wide bucket_cap = domain_size(bucket_bits);
if (b > bucket_cap)
{
throw std::invalid_argument(
"make_multipoint: bucket domain does not fit in BucketInput");
}
const int attempts = params.retries < 1 ? 1 : params.retries;
for (int attempt = 0; attempt < attempts; ++attempt)
{
try
{
const simde__m128i sigma = dpf::uniform_sample<simde__m128i>();
std::vector<slot> table;
if (!insert_cuckoo(sigma, alphas, m, n, b, params.max_evictions, table))
continue;
key0 left;
key1 right;
left.sigma = sigma;
right.sigma = sigma;
left.bucket_count = m;
right.bucket_count = m;
left.bucket_domain = static_cast<std::uint64_t>(b);
right.bucket_domain = static_cast<std::uint64_t>(b);
left.buckets.reserve(m);
right.buckets.reserve(m);
for (std::uint32_t i = 0; i < m; ++i)
{
BucketInput gamma{};
OutputT beta{};
if (table[i].item >= 0)
{
const auto & alpha = alphas[static_cast<std::size_t>(table[i].item)];
const auto loc = locate(sigma,
static_cast<wide>(alpha), table[i].hash, n, b);
if (loc.bucket != i)
throw prp_walk_error{};
gamma = static_cast<BucketInput>(loc.index);
beta = betas[static_cast<std::size_t>(table[i].item)];
}
auto made = make_bucket<Verifiable, InteriorPRG, ExteriorPRG>(
gamma, beta);
left.buckets.push_back(std::move(made.first));
right.buckets.push_back(std::move(made.second));
}
return std::make_pair(std::move(left), std::move(right));
}
catch (const prp_walk_error &)
{
continue;
}
}
throw std::runtime_error("make_multipoint: cuckoo hashing failed");
}
inline void absorb_proof(proof_token & acc, const proof_token & inner)
{
acc = detail::vdpf::xor_proof(acc, inner);
acc[0] = detail::vdpf::mmo(acc[0], 1);
}
template <typename Key>
typename Key::share_type eval_at(const Key & key, typename Key::input_type x,
proof_token * acc)
{
using input_type = typename Key::input_type;
using bucket_input = typename Key::bucket_input;
constexpr std::size_t input_bits = utils::bitlength_of_v<input_type>;
const wide n = domain_size(input_bits);
const wide b = key.bucket_domain;
typename Key::share_type sum =
Key::share_type::from_raw(typename Key::output_type{});
for (int hash = 0; hash < static_cast<int>(Key::kappa); ++hash)
{
const auto loc = locate(key.sigma, static_cast<wide>(x), hash, n, b);
if (loc.bucket >= key.bucket_count)
throw std::runtime_error("multipoint eval: bucket out of range");
const auto gamma = static_cast<bucket_input>(loc.index);
const auto & bucket = key.buckets[loc.bucket];
if constexpr (Key::is_verifiable)
{
if (acc != nullptr)
{
proof_token inner{};
sum += *dpf::eval_point(bucket, gamma, dpf::prove(inner));
absorb_proof(*acc, inner);
continue;
}
}
sum += *dpf::eval_point(bucket, gamma);
}
return sum;
}
} // namespace mpf
} // namespace detail
/// @brief Cuckoo-pack distinct points into ordinary point-key buckets.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam BucketInput unsigned type of a bucket index. Defaults to `uint32_t`
/// @tparam AlphaRange range of distinct domain points
/// @tparam BetaRange range of payloads, one per point
/// @param alphas the secret points
/// @param betas the payloads
/// @param params packing knobs. `lambda` is the Remark 1 failure target
/// @return the two party keys
/// @throws std::invalid_argument if the lists differ in length, are empty,
/// contain a duplicate, or a bucket index does not fit `BucketInput`
/// @throws std::runtime_error if cuckoo hashing does not succeed
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename BucketInput = std::uint32_t,
typename AlphaRange,
typename BetaRange>
HEDLEY_WARN_UNUSED_RESULT
auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
multipoint_params params = {})
{
using input_type = std::decay_t<decltype(*std::begin(alphas))>;
using output_type = std::decay_t<decltype(*std::begin(betas))>;
return detail::mpf::make_impl<false, InteriorPRG, ExteriorPRG, BucketInput>(
std::vector<input_type>(std::begin(alphas), std::end(alphas)),
std::vector<output_type>(std::begin(betas), std::end(betas)),
params);
}
/// @brief Same packing as `make_multipoint`, with a verifiable bucket key.
/// @see `make_multipoint`
/// @param alphas the secret points
/// @param betas the payloads
/// @param params packing knobs
/// @return the two verifiable party keys
/// @throws std::invalid_argument if the lists differ in length, are empty,
/// contain a duplicate, or a bucket index does not fit `BucketInput`
/// @throws std::runtime_error if cuckoo hashing does not succeed
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename BucketInput = std::uint32_t,
typename AlphaRange,
typename BetaRange>
HEDLEY_WARN_UNUSED_RESULT
auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
verifiable, multipoint_params params = {})
{
using input_type = std::decay_t<decltype(*std::begin(alphas))>;
using output_type = std::decay_t<decltype(*std::begin(betas))>;
return detail::mpf::make_impl<true, InteriorPRG, ExteriorPRG, BucketInput>(
std::vector<input_type>(std::begin(alphas), std::end(alphas)),
std::vector<output_type>(std::begin(betas), std::end(betas)),
params);
}
/// @brief Sum the three bucket shares at `x`.
/// @tparam Key a `multipoint_key`
/// @param key the party key
/// @param x the query point
/// @return the party's share of the payload, or of zero off the packed points
/// @throws std::runtime_error if a located bucket is outside the key
template <typename Key,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
auto eval_multipoint(const Key & key, typename Key::input_type x)
{
return detail::mpf::eval_at(key, x, nullptr);
}
/// @brief Evaluate `x` and fold that query into `pr`.
/// @tparam Key a verifiable `multipoint_key`
/// @param key the party key
/// @param x the query point
/// @param pr proof token replaced with this query's folded proof
/// @return the party's share of the payload
/// @throws std::runtime_error if a located bucket is outside the key
template <typename Key,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
auto eval_multipoint(const Key & key, typename Key::input_type x, prove_ref pr)
{
static_assert(Key::is_verifiable,
"eval_multipoint(..., prove(π)): key must be a verifiable multipoint key");
pr.token = detail::vdpf::zero_proof();
return detail::mpf::eval_at(key, x, &pr.token);
}
/// @brief Evaluate each point of `xs`, writing one share per point.
/// @tparam Key a `multipoint_key`
/// @tparam Range range of query points
/// @tparam OutIt output iterator of shares
/// @param key the party key
/// @param xs the query points
/// @param out where each share is written
/// @throws std::runtime_error if a located bucket is outside the key
template <typename Key, typename Range, typename OutIt,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
void eval_multipoint(const Key & key, const Range & xs, OutIt out)
{
for (const auto & x : xs)
*out++ = eval_multipoint(key, static_cast<typename Key::input_type>(x));
}
/// @brief Evaluate `xs` and fold every query into one proof.
/// @tparam Key a verifiable `multipoint_key`
/// @tparam Range range of query points
/// @tparam OutIt output iterator of shares
/// @param key the party key
/// @param xs the query points
/// @param out where each share is written
/// @param pr proof token replaced with the folded proof of `xs`
/// @throws std::runtime_error if a located bucket is outside the key
template <typename Key, typename Range, typename OutIt,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
void eval_multipoint(const Key & key, const Range & xs, OutIt out, prove_ref pr)
{
static_assert(Key::is_verifiable,
"eval_multipoint(..., prove(π)): key must be a verifiable multipoint key");
pr.token = detail::vdpf::zero_proof();
for (const auto & x : xs)
{
*out++ = detail::mpf::eval_at(key,
static_cast<typename Key::input_type>(x), &pr.token);
}
}
/// @brief Fold a canonical evaluation of every bucket into one proof.
/// @tparam Key a verifiable `multipoint_key`
/// @param key the party key
/// @param pr proof token replaced with the audit proof
template <typename Key,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
void audit_multipoint(const Key & key, prove_ref pr)
{
static_assert(Key::is_verifiable,
"audit_multipoint: key must be a verifiable multipoint key");
pr.token = detail::vdpf::zero_proof();
for (const auto & bucket : key.buckets)
{
proof_token inner{};
(void)*dpf::eval_point(bucket, typename Key::bucket_input{},
dpf::prove(inner));
detail::mpf::absorb_proof(pr.token, inner);
}
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_MULTIPOINT_HPP__

View file

@ -4,6 +4,7 @@
/// packs one lane every four bits, low nibble first. Leaf addition /// packs one lane every four bits, low nibble first. Leaf addition
/// is not XOR and is not `add_epi8`: a carry must not cross into /// is not XOR and is not `add_epi8`: a carry must not cross into
/// the neighbouring nibble. See `packed_lane_arithmetic.hpp`. /// the neighbouring nibble. See `packed_lane_arithmetic.hpp`.
/// @see packed_lane_arithmetic.hpp
#ifndef LIBDPF_INCLUDE_DPF_NYBLE_HPP__ #ifndef LIBDPF_INCLUDE_DPF_NYBLE_HPP__
#define LIBDPF_INCLUDE_DPF_NYBLE_HPP__ #define LIBDPF_INCLUDE_DPF_NYBLE_HPP__
@ -47,6 +48,9 @@ static constexpr dpf::nyble to_nyble(unsigned long long value) noexcept
} }
/// @brief parse one hex digit as a nibble /// @brief parse one hex digit as a nibble
/// @tparam CharT character type
/// @param value the value to convert or store
/// @return the returned `dpf::nyble`
/// @throws std::domain_error if `value` is not `0-9`, `a-f`, or `A-F` /// @throws std::domain_error if `value` is not `0-9`, `a-f`, or `A-F`
template <typename CharT> template <typename CharT>
static constexpr dpf::nyble to_nyble(CharT value) static constexpr dpf::nyble to_nyble(CharT value)
@ -97,6 +101,9 @@ operator>>(std::basic_istream<CharT, Traits> & is, dpf::nyble & value)
} }
/// @brief addition in Z/16Z /// @brief addition in Z/16Z
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return addition in Z/16Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -107,6 +114,9 @@ constexpr dpf::nyble operator+(dpf::nyble lhs, dpf::nyble rhs) noexcept
} }
/// @brief subtraction in Z/16Z /// @brief subtraction in Z/16Z
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return subtraction in Z/16Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -117,6 +127,8 @@ constexpr dpf::nyble operator-(dpf::nyble lhs, dpf::nyble rhs) noexcept
} }
/// @brief additive inverse in Z/16Z /// @brief additive inverse in Z/16Z
/// @param value the value to convert or store
/// @return additive inverse in Z/16Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -126,6 +138,9 @@ constexpr dpf::nyble operator-(dpf::nyble value) noexcept
} }
/// @brief multiplication in Z/16Z /// @brief multiplication in Z/16Z
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return multiplication in Z/16Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE

View file

@ -1,6 +1,5 @@
/// @file dpf/offset_wrapper.hpp /// @file dpf/offset_wrapper.hpp
/// @brief /// @brief An output value shifted by a public offset.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -9,6 +8,9 @@
#ifndef LIBDPF_INCLUDE_DPF_OFFSET_WRAPPER_HPP__ #ifndef LIBDPF_INCLUDE_DPF_OFFSET_WRAPPER_HPP__
#define LIBDPF_INCLUDE_DPF_OFFSET_WRAPPER_HPP__ #define LIBDPF_INCLUDE_DPF_OFFSET_WRAPPER_HPP__
#include <cstdint>
#include <utility>
#include "hedley/hedley.h" #include "hedley/hedley.h"
namespace dpf namespace dpf
@ -43,6 +45,13 @@ struct offset_wrapper final
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr bool is_wildcard() noexcept { return false; } static constexpr bool is_wildcard() noexcept { return false; }
/// @brief Stored share. Concrete wrappers ignore it at evaluation.
/// @return the stored share
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
constexpr const input_type & raw() const noexcept { return offset_; }
private: private:
input_type offset_; // waste an `input_type` to make `sizeof` match up input_type offset_; // waste an `input_type` to make `sizeof` match up
}; };
@ -108,6 +117,34 @@ struct offset_wrapper<dpf::wildcard_value<ConcreteInputT>>
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr bool is_wildcard() noexcept { return true; } static constexpr bool is_wildcard() noexcept { return true; }
/// @brief Party share, including before the offset is marked ready.
/// @return the party share
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
constexpr const input_type & raw() const noexcept { return offset_; }
/// @brief `offset_status` as a byte, for serialization.
/// @return the status byte
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
constexpr std::uint8_t state() const noexcept
{
return static_cast<std::uint8_t>(offset_state_);
}
/// @brief Restore a serialized share and status byte.
/// @param offset the party share
/// @param state the status byte
HEDLEY_ALWAYS_INLINE
void restore(input_type offset, std::uint8_t state) noexcept
{
offset_ = std::move(offset);
offset_state_ = static_cast<offset_status>(state);
}
private: private:
enum class offset_status : psnip_uint8_t { ready = 0, waiting = 1, computing = 2, notset = 3 }; enum class offset_status : psnip_uint8_t { ready = 0, waiting = 1, computing = 2, notset = 3 };

View file

@ -41,8 +41,10 @@
namespace dpf namespace dpf
{ {
/// Buffer element type for leaf eval of `KeyT`: party-tagged subtractive /// @brief Buffer element type for leaf eval of `KeyT`: party-tagged subtractive
/// share when `KeyT` is a `party_key`, otherwise the concrete output. /// share when `KeyT` is a `party_key`, otherwise the concrete output.
/// @tparam KeyT key type
/// @tparam OutputT output type
template <typename KeyT, typename OutputT, bool = is_party_key_v<KeyT>> template <typename KeyT, typename OutputT, bool = is_party_key_v<KeyT>>
struct leaf_buffer_elem struct leaf_buffer_elem
{ {
@ -56,7 +58,9 @@ struct leaf_buffer_elem<KeyT, OutputT, true>
template <typename KeyT, typename OutputT> template <typename KeyT, typename OutputT>
using leaf_buffer_elem_t = typename leaf_buffer_elem<KeyT, OutputT>::type; using leaf_buffer_elem_t = typename leaf_buffer_elem<KeyT, OutputT>::type;
/// Buffer element type for comparison eval of `KeyT`. /// @brief Buffer element type for comparison eval of `KeyT`.
/// @tparam KeyT key type
/// @tparam Beta payload type
template <typename KeyT, typename Beta, bool = is_party_key_v<KeyT>> template <typename KeyT, typename Beta, bool = is_party_key_v<KeyT>>
struct cmp_buffer_elem struct cmp_buffer_elem
{ {
@ -70,9 +74,11 @@ struct cmp_buffer_elem<KeyT, Beta, true>
template <typename KeyT, typename Beta> template <typename KeyT, typename Beta>
using cmp_buffer_elem_t = typename cmp_buffer_elem<KeyT, Beta>::type; using cmp_buffer_elem_t = typename cmp_buffer_elem<KeyT, Beta>::type;
/// `std::vector(n)` value-initializes every slot. Interval / full eval /// @brief `std::vector(n)` value-initializes every slot. Interval / full eval
/// overwrites the whole buffer, so skip default-construction for trivial /// overwrites the whole buffer, so skip default-construction for trivial
/// `T`. Non-trivial outputs still run their default constructor. /// `T`. Non-trivial outputs still run their default constructor.
/// @tparam T value type
/// @tparam Alignment allocation alignment
template <typename T, template <typename T,
std::size_t Alignment> std::size_t Alignment>
class output_buffer_allocator : public aligned_allocator<T, Alignment> class output_buffer_allocator : public aligned_allocator<T, Alignment>
@ -139,8 +145,10 @@ constexpr bool operator!=(const output_buffer_allocator<T, A> & lhs,
return !(lhs == rhs); return !(lhs == rhs);
} }
/// Move-only vector of `T`. Copy construction and copy assignment are /// @brief Move-only vector of `T`. Copy construction and copy assignment are
/// deleted. `at`, `operator[]`, `data`, iterators, and `size` are public. /// deleted. `at`, `operator[]`, `data`, iterators, and `size` are public.
/// @tparam T value type
/// @tparam Alignment allocation alignment
template <typename T, template <typename T,
std::size_t Alignment = utils::max_align_v> std::size_t Alignment = utils::max_align_v>
class output_buffer final class output_buffer final
@ -252,7 +260,7 @@ LIBDPF_PACKED_SHARE_BUFFER(dpf::nyble, 0);
LIBDPF_PACKED_SHARE_BUFFER(dpf::nyble, 1); LIBDPF_PACKED_SHARE_BUFFER(dpf::nyble, 1);
#undef LIBDPF_PACKED_SHARE_BUFFER #undef LIBDPF_PACKED_SHARE_BUFFER
/// Packed bit share buffers reuse the bit-array image; iterators yield shares. /// @brief Packed bit share buffers reuse the bit-array image; iterators yield shares.
#define LIBDPF_BIT_SHARE_BUFFER(PARTY) \ #define LIBDPF_BIT_SHARE_BUFFER(PARTY) \
template <> \ template <> \
class output_buffer<subtractive_share<dpf::bit, PARTY>> \ class output_buffer<subtractive_share<dpf::bit, PARTY>> \
@ -272,8 +280,14 @@ LIBDPF_BIT_SHARE_BUFFER(0);
LIBDPF_BIT_SHARE_BUFFER(1); LIBDPF_BIT_SHARE_BUFFER(1);
#undef LIBDPF_BIT_SHARE_BUFFER #undef LIBDPF_BIT_SHARE_BUFFER
/// Buffer sized for the closed interval `[from, to]` of output `I`. /// @brief Buffer sized for the closed interval `[from, to]` of output `I`.
/// On a `party_key`, elements are subtractive shares of that output. /// @details On a `party_key`, elements are subtractive shares of that output.
/// @tparam DpfKey DPF key type
/// @tparam I output index
/// @tparam InputT input domain type
/// @param from the inclusive start of the range
/// @param to the `to`
/// @return Buffer sized for the closed interval `[from, to]` of output `I`
template <typename DpfKey, template <typename DpfKey,
std::size_t I = 0, std::size_t I = 0,
typename InputT> typename InputT>
@ -318,7 +332,10 @@ inline auto make_output_buffer_for_interval(const DpfKey &, InputT from, InputT
return make_output_buffer_for_interval<DpfKey, I0, I1, Is...>(from, to); return make_output_buffer_for_interval<DpfKey, I0, I1, Is...>(from, to);
} }
/// Buffer sized for every input of output `I`. /// @brief Buffer sized for every input of output `I`.
/// @tparam DpfKey DPF key type
/// @tparam I output index
/// @return Buffer sized for every input of output `I`
template <typename DpfKey, template <typename DpfKey,
std::size_t I = 0> std::size_t I = 0>
auto make_output_buffer_for_full() auto make_output_buffer_for_full()

View file

@ -312,8 +312,10 @@ class dynamic_packed_array
unique_ptr data_{}; unique_ptr data_{};
}; };
/// Packed lane storage whose iterators yield `subtractive_share<LaneT, Party>`. /// @brief Packed lane storage whose iterators yield `subtractive_share<LaneT, Party>`.
/// The bytes are the leaf image (`store_leaf_bytes`); each lane is one share. /// @details The bytes are the leaf image (`store_leaf_bytes`); each lane is one share.
/// @tparam LaneT lane type
/// @tparam Party party index, `0` or `1`
template <typename LaneT, std::size_t Party> template <typename LaneT, std::size_t Party>
class packed_share_output : public dynamic_packed_array<LaneT> class packed_share_output : public dynamic_packed_array<LaneT>
{ {

View file

@ -41,6 +41,11 @@ LaneT extract_lane(const LeafT & leaf, std::size_t lane) noexcept
/// @brief zero `out` is the caller's job; this writes one lane and leaves /// @brief zero `out` is the caller's job; this writes one lane and leaves
/// every other lane untouched. /// every other lane untouched.
/// @tparam LaneT lane type
/// @tparam LeafT leaf type
/// @param leaf the leaf value
/// @param lane the lane index or lane value
/// @param value the value to convert or store
template <typename LaneT, typename LeafT> template <typename LaneT, typename LeafT>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW

View file

@ -131,8 +131,11 @@ inline simde__m256i shuffle_nibbles(simde__m256i table, simde__m256i a) noexcept
simde_mm256_slli_epi16(simde_mm256_and_si256(hi, m), 4)); simde_mm256_slli_epi16(simde_mm256_and_si256(hi, m), 4));
} }
/// Low nibble of every byte, product mod 16. Even and odd bytes are split /// @brief Low nibble of every byte, product mod 16. Even and odd bytes are split
/// so a product in one byte cannot land in the next. /// so a product in one byte cannot land in the next.
/// @param a the `a`
/// @param b the `b`
/// @return Low nibble of every byte, product mod 16
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline simde__m128i mul_low_nibbles(simde__m128i a, simde__m128i b) noexcept inline simde__m128i mul_low_nibbles(simde__m128i a, simde__m128i b) noexcept
{ {

View file

@ -1,7 +1,6 @@
/// @file dpf/parallel_bit_iterable.hpp /// @file dpf/parallel_bit_iterable.hpp
/// @author Christopher Jiang <christopher.jiang@ucalgary.ca> /// @author Christopher Jiang <christopher.jiang@ucalgary.ca>
/// @brief /// @brief SIMD iteration over packed advice or correction bits.
/// @details
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others]{@ref authors} /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others]{@ref authors}
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details. /// see [LICENSE.md](@ref license) for details.

View file

@ -1,7 +1,6 @@
/// @file dpf/parallel_bit_iterable_helpers.hpp /// @file dpf/parallel_bit_iterable_helpers.hpp
/// @author Christopher Jiang <christopher.jiang@ucalgary.ca> /// @author Christopher Jiang <christopher.jiang@ucalgary.ca>
/// @brief /// @brief Loads and masks used by the parallel bit iterators.
/// @details
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details. /// see [LICENSE.md](@ref license) for details.
@ -28,8 +27,13 @@ namespace dpf
namespace namespace
{ {
/// Unaligned 256-bit load of `words_per_vec` words starting at `offset`. /// @brief Unaligned 256-bit load of `words_per_vec` words starting at `offset`.
/// Words past `nwords` are zero so a short batch does not read off the end. /// @details Words past `nwords` are zero so a short batch does not read off the end.
/// @tparam Word word
/// @param words the `words`
/// @param nwords the number of words
/// @param offset the public offset
/// @return Unaligned 256-bit load of `words_per_vec` words starting at `offset`
template <typename Word> template <typename Word>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -56,6 +60,7 @@ template <std::size_t batch_size_log_2, typename ChildT>
struct parallel_bit_iterable_helper; struct parallel_bit_iterable_helper;
/// @brief for batch_size in 1..4 /// @brief for batch_size in 1..4
/// @tparam ChildT CRTP derived type
template <typename ChildT> template <typename ChildT>
struct parallel_bit_iterable_helper<2, ChildT> struct parallel_bit_iterable_helper<2, ChildT>
{ {
@ -88,6 +93,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
}; // struct parallel_bit_iterable_helper<2> }; // struct parallel_bit_iterable_helper<2>
/// @brief for batch_size in 5..8 /// @brief for batch_size in 5..8
/// @tparam ChildT CRTP derived type
template <typename ChildT> template <typename ChildT>
struct parallel_bit_iterable_helper<3, ChildT> struct parallel_bit_iterable_helper<3, ChildT>
{ {
@ -135,6 +141,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
}; // struct parallel_bit_iterable_helper<3> }; // struct parallel_bit_iterable_helper<3>
/// @brief for batch_size in 9..16 /// @brief for batch_size in 9..16
/// @tparam ChildT CRTP derived type
template <typename ChildT> template <typename ChildT>
struct parallel_bit_iterable_helper<4, ChildT> struct parallel_bit_iterable_helper<4, ChildT>
{ {
@ -203,6 +210,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
}; // struct parallel_bit_iterable_helper<4> }; // struct parallel_bit_iterable_helper<4>
/// @brief for batch_size in 17..32 /// @brief for batch_size in 17..32
/// @tparam ChildT CRTP derived type
template <typename ChildT> template <typename ChildT>
struct parallel_bit_iterable_helper<5, ChildT> struct parallel_bit_iterable_helper<5, ChildT>
{ {

View file

@ -34,7 +34,7 @@
namespace dpf namespace dpf
{ {
/// Path memoizers key on the underlying DPF key type. `party_key` wrappers /// @brief Path memoizers key on the underlying DPF key type. `party_key` wrappers
/// share the same tree layout, so a memoizer built for party 0 also accepts /// share the same tree layout, so a memoizer built for party 0 also accepts
/// party 1 (and bare keys). /// party 1 (and bare keys).
template <typename T> template <typename T>
@ -62,10 +62,11 @@ struct path_memoizer_base
virtual return_type end() const noexcept = 0; virtual return_type end() const noexcept = 0;
}; };
/// One interior node per level. `assign_x` returns the first level that the /// @brief One interior node per level. `assign_x` returns the first level that the
/// next walk must recompute. `filled_to` is the deepest level already /// next walk must recompute. `filled_to` is the deepest level already
/// written for the current input. Callers pass this object to `eval_point`; /// written for the current input. Callers pass this object to `eval_point`;
/// they do not call `assign_x` themselves. /// they do not call `assign_x` themselves.
/// @tparam DpfKey DPF key type
template <typename DpfKey> template <typename DpfKey>
struct alignas(alignof(typename path_memoizer_key_t<DpfKey>::interior_node)) struct alignas(alignof(typename path_memoizer_key_t<DpfKey>::interior_node))
basic_path_memoizer final basic_path_memoizer final
@ -143,7 +144,8 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
return std::addressof(arr_[depth+1]); return std::addressof(arr_[depth+1]);
} }
/// Inclusive high-water: `arr_[0..filled_to_]` are valid for the current x. /// @brief Inclusive high-water: `arr_[0..filled_to_]` are valid for the current x.
/// @return Inclusive high-water: `arr_[0..filled_to_]` are valid for the current x
HEDLEY_NO_THROW HEDLEY_NO_THROW
std::size_t filled_to() const noexcept { return filled_to_; } std::size_t filled_to() const noexcept { return filled_to_; }
@ -166,7 +168,8 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
}; };
/// A single interior node. Every `assign_x` restarts at the root. /// @brief A single interior node. Every `assign_x` restarts at the root.
/// @tparam DpfKey DPF key type
template <typename DpfKey> template <typename DpfKey>
struct nonmemoizing_path_memoizer final struct nonmemoizing_path_memoizer final
: public path_memoizer_base<path_memoizer_key_t<DpfKey>> : public path_memoizer_base<path_memoizer_key_t<DpfKey>>
@ -267,7 +270,13 @@ void path_note_filled_to(PathMemoizer & path, std::size_t level)
path.note_filled(level); path.note_filled(level);
} }
/// Walk interior nodes so `path[0..to_level]` is valid for `x`. /// @brief Walk interior nodes so `path[0..to_level]` is valid for `x`.
/// @tparam DpfKey DPF key type
/// @tparam PathMemoizer path memoizer type
/// @param dpf the DPF key
/// @param x the `x`
/// @param path the root-to-leaf path
/// @param to_level the `to_level`
template <typename DpfKey, typename PathMemoizer> template <typename DpfKey, typename PathMemoizer>
void ensure_level(const DpfKey & dpf, typename DpfKey::input_type x, void ensure_level(const DpfKey & dpf, typename DpfKey::input_type x,
PathMemoizer & path, std::size_t to_level) PathMemoizer & path, std::size_t to_level)
@ -279,17 +288,21 @@ void ensure_level(const DpfKey & dpf, typename DpfKey::input_type x,
{ {
bool bit = !!(mask & x); bool bit = !!(mask & x);
auto cw = dpf.correction_word(level_index - 1, bit); auto cw = dpf.correction_word(level_index - 1, bit);
const bool is_last = DpfKey::tree::is_last_level(level_index - 1,
dpf.depth);
path[level_index] = path[level_index] =
DpfKey::traverse_interior(path[level_index - 1], cw, bit); DpfKey::traverse_interior(path[level_index - 1], cw, bit, is_last);
} }
path_note_filled_to(path, to_level); path_note_filled_to(path, to_level);
} }
} // namespace detail } // namespace detail
/// Path workspace for `DpfKey`. A `party_key` argument is unwrapped, and the /// @brief Path workspace for `DpfKey`. A `party_key` argument is unwrapped, and the
/// result accepts both parties. /// result accepts both parties.
/// @snippet evaluation/memoizers.cpp path-memoizer /// @snippet evaluation/memoizers.cpp path-memoizer
/// @tparam DpfKey DPF key type
/// @return Path workspace for `DpfKey`
template <typename DpfKey> template <typename DpfKey>
auto make_basic_path_memoizer() auto make_basic_path_memoizer()
{ {
@ -302,7 +315,9 @@ auto make_basic_path_memoizer(const DpfKey &)
return make_basic_path_memoizer<DpfKey>(); return make_basic_path_memoizer<DpfKey>();
} }
/// Single-node path workspace. Suitable for one query. /// @brief Single-node path workspace. Suitable for one query.
/// @tparam DpfKey DPF key type
/// @return Single-node path workspace
template <typename DpfKey> template <typename DpfKey>
auto make_nonmemoizing_path_memoizer() auto make_nonmemoizing_path_memoizer()
{ {

View file

@ -58,26 +58,31 @@ template <std::size_t N, typename O, typename ...Os>
struct is_at<at_pack<N, O, Os...>> : std::true_type {}; struct is_at<at_pack<N, O, Os...>> : std::true_type {};
template <typename T> inline constexpr bool is_at_v = is_at<T>::value; template <typename T> inline constexpr bool is_at_v = is_at<T>::value;
/// Phantom pack element for a key's comparison (DCF) channel. Not a leaf: it /// @brief Phantom pack element for a key's comparison (DCF) channel. Not a leaf: it
/// only records the cmp prefix depth in the key's type. `Depth` is the number /// only records the cmp prefix depth in the key's type. `Depth` is the number
/// of tree levels the comparison walks (0 is reserved for "no cmp"). /// of tree levels the comparison walks (0 is reserved for "no cmp").
/// `OutBits` is the comparison output group width (bits of the β payload), /// @details `OutBits` is the comparison output group width (bits of the β payload),
/// so the value CWs / addend can be stored at group width instead of a full /// so the value CWs / addend can be stored at group width instead of a full
/// padded `uint64_t` per level. /// padded `uint64_t` per level.
/// @tparam Depth depth
/// @tparam OutBits out bits
/// @tparam Wild whether the payload is a wildcard
/// @tparam BlockWidth checkpoint spacing, in levels
/// @tparam Incremental whether a final correction is stored at every depth
template <std::size_t Depth, std::size_t OutBits = 0, bool Wild = false, template <std::size_t Depth, std::size_t OutBits = 0, bool Wild = false,
std::size_t BlockWidth = 0, bool Incremental = false> std::size_t BlockWidth = 0, bool Incremental = false>
struct cmp_channel_tag struct cmp_channel_tag
{ {
static constexpr std::size_t depth = Depth; static constexpr std::size_t depth = Depth;
static constexpr std::size_t out_bits = OutBits; static constexpr std::size_t out_bits = OutBits;
/// True when the comparison payload (β) is a wildcard to be assigned /// @brief True when the comparison payload (β) is a wildcard to be assigned
/// after keygen. Concrete (non-wildcard) cmp keys keep `Wild == false` /// after keygen. Concrete (non-wildcard) cmp keys keep `Wild == false`
/// so their layout / type name is unchanged. /// so their layout / type name is unchanged.
static constexpr bool wild = Wild; static constexpr bool wild = Wild;
/// 0 keeps the per-level path-sum. `B >= 1` selects blocked checkpoints /// @brief 0 keeps the per-level path-sum. `B >= 1` selects blocked checkpoints
/// of target width `B`. /// of target width `B`.
static constexpr std::size_t block_width = BlockWidth; static constexpr std::size_t block_width = BlockWidth;
/// Save a final correction at every depth (`idcf`). /// @brief Save a final correction at every depth (`idcf`).
static constexpr bool incremental = Incremental; static constexpr bool incremental = Incremental;
}; };
@ -90,6 +95,38 @@ template <typename T>
inline constexpr bool is_cmp_channel_tag_v = inline constexpr bool is_cmp_channel_tag_v =
is_cmp_channel_tag<std::decay_t<T>>::value; is_cmp_channel_tag<std::decay_t<T>>::value;
/// Phantom pack element: key carries per-level VDPF correction seeds.
struct verifiable
{
static constexpr bool is_verifiable_tag = true;
};
/// Phantom pack element: ROM leaf stretch + extractability checks.
struct extractable
{
static constexpr bool is_extractable_tag = true;
};
template <typename T>
struct is_verifiable_tag : std::false_type
{ };
template <>
struct is_verifiable_tag<verifiable> : std::true_type
{ };
template <typename T>
inline constexpr bool is_verifiable_tag_v =
is_verifiable_tag<std::decay_t<T>>::value;
template <typename T>
struct is_extractable_tag : std::false_type
{ };
template <>
struct is_extractable_tag<extractable> : std::true_type
{ };
template <typename T>
inline constexpr bool is_extractable_tag_v =
is_extractable_tag<std::decay_t<T>>::value;
namespace detail namespace detail
{ {
namespace incr namespace incr
@ -99,13 +136,25 @@ namespace incr
// A concrete placed output: an output type `OutputT` planted at prefix `N`. // A concrete placed output: an output type `OutputT` planted at prefix `N`.
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
/// @brief Default: placed payload type is the stored type. Specialized for
/// `arith_beta<T>` in `doerner_shelat.hpp` so the key leaf type is `T`.
/// @tparam T value type
template <typename T>
struct unwrap_placed_output
{
using type = T;
};
template <std::size_t N, typename OutputT> template <std::size_t N, typename OutputT>
struct placed struct placed
{ {
static constexpr std::size_t prefix = N; static constexpr std::size_t prefix = N;
using output_type = OutputT; /// @brief Stored argument type (may be `arith_beta<T>`).
using stored_type = OutputT;
/// @brief Key / leaf payload type (`T` when stored is `arith_beta<T>`).
using output_type = typename unwrap_placed_output<OutputT>::type;
OutputT value; OutputT value;
OutputT addend{}; // public if_false for eq(...); party 0 absorbs at eval output_type addend{}; // public if_false for eq(...); party 0 absorbs at eval
}; };
template <typename T> struct is_placed : std::false_type {}; template <typename T> struct is_placed : std::false_type {};
@ -114,8 +163,10 @@ struct is_placed<placed<N, O>> : std::true_type {};
template <typename T> template <typename T>
inline constexpr bool is_placed_v = is_placed<std::decay_t<T>>::value; inline constexpr bool is_placed_v = is_placed<std::decay_t<T>>::value;
/// Heavy-hitters incremental point function: one payload per prefix length. /// @brief Heavy-hitters incremental point function: one payload per prefix length.
/// `levels[i]` is the bit length of slot `i`. /// @details `levels[i]` is the bit length of slot `i`.
/// @tparam LevelSeq level seq
/// @tparam Betas betas
template <typename LevelSeq, typename ...Betas> template <typename LevelSeq, typename ...Betas>
struct idpf_pack; struct idpf_pack;
@ -416,6 +467,8 @@ struct normalize_one
static constexpr bool cmp_wild = false; static constexpr bool cmp_wild = false;
static constexpr std::size_t cmp_block = 0; static constexpr std::size_t cmp_block = 0;
static constexpr bool cmp_idcf = false; static constexpr bool cmp_idcf = false;
static constexpr bool is_verifiable = false;
static constexpr bool is_extractable = false;
}; };
template <std::size_t BitLen, std::size_t N, typename T> template <std::size_t BitLen, std::size_t N, typename T>
struct normalize_one<BitLen, placed<N, T>> struct normalize_one<BitLen, placed<N, T>>
@ -426,6 +479,8 @@ struct normalize_one<BitLen, placed<N, T>>
static constexpr bool cmp_wild = false; static constexpr bool cmp_wild = false;
static constexpr std::size_t cmp_block = 0; static constexpr std::size_t cmp_block = 0;
static constexpr bool cmp_idcf = false; static constexpr bool cmp_idcf = false;
static constexpr bool is_verifiable = false;
static constexpr bool is_extractable = false;
}; };
template <std::size_t BitLen, std::size_t Depth, std::size_t OutBits, bool Wild, template <std::size_t BitLen, std::size_t Depth, std::size_t OutBits, bool Wild,
std::size_t Block, bool Incremental> std::size_t Block, bool Incremental>
@ -438,6 +493,32 @@ struct normalize_one<BitLen,
static constexpr bool cmp_wild = Wild; static constexpr bool cmp_wild = Wild;
static constexpr std::size_t cmp_block = Block; static constexpr std::size_t cmp_block = Block;
static constexpr bool cmp_idcf = Incremental; static constexpr bool cmp_idcf = Incremental;
static constexpr bool is_verifiable = false;
static constexpr bool is_extractable = false;
};
template <std::size_t BitLen>
struct normalize_one<BitLen, verifiable>
{
using placed_tuple = std::tuple<>;
static constexpr std::size_t cmp_depth = 0;
static constexpr std::size_t cmp_out_bits = 0;
static constexpr bool cmp_wild = false;
static constexpr std::size_t cmp_block = 0;
static constexpr bool cmp_idcf = false;
static constexpr bool is_verifiable = true;
static constexpr bool is_extractable = false;
};
template <std::size_t BitLen>
struct normalize_one<BitLen, extractable>
{
using placed_tuple = std::tuple<>;
static constexpr std::size_t cmp_depth = 0;
static constexpr std::size_t cmp_out_bits = 0;
static constexpr bool cmp_wild = false;
static constexpr std::size_t cmp_block = 0;
static constexpr bool cmp_idcf = false;
static constexpr bool is_verifiable = false;
static constexpr bool is_extractable = true;
}; };
template <std::size_t BitLen, typename ...Elems> template <std::size_t BitLen, typename ...Elems>
@ -448,18 +529,18 @@ struct normalize_pack
std::declval<typename normalize_one<BitLen, Elems>::placed_tuple>()...)); std::declval<typename normalize_one<BitLen, Elems>::placed_tuple>()...));
static constexpr std::size_t cmp_depth = static constexpr std::size_t cmp_depth =
(std::size_t{0} + ... + normalize_one<BitLen, Elems>::cmp_depth); (std::size_t{0} + ... + normalize_one<BitLen, Elems>::cmp_depth);
// At most one comparison channel per key, so the sum is that channel's
// output width (0 when there is no cmp channel).
static constexpr std::size_t cmp_out_bits = static constexpr std::size_t cmp_out_bits =
(std::size_t{0} + ... + normalize_one<BitLen, Elems>::cmp_out_bits); (std::size_t{0} + ... + normalize_one<BitLen, Elems>::cmp_out_bits);
// At most one comparison channel per key, so the OR is that channel's
// wildcard flag (false when there is no cmp channel).
static constexpr bool cmp_wild = static constexpr bool cmp_wild =
(false || ... || normalize_one<BitLen, Elems>::cmp_wild); (false || ... || normalize_one<BitLen, Elems>::cmp_wild);
static constexpr std::size_t cmp_block = static constexpr std::size_t cmp_block =
(std::size_t{0} + ... + normalize_one<BitLen, Elems>::cmp_block); (std::size_t{0} + ... + normalize_one<BitLen, Elems>::cmp_block);
static constexpr bool cmp_idcf = static constexpr bool cmp_idcf =
(false || ... || normalize_one<BitLen, Elems>::cmp_idcf); (false || ... || normalize_one<BitLen, Elems>::cmp_idcf);
static constexpr bool is_verifiable =
(false || ... || normalize_one<BitLen, Elems>::is_verifiable);
static constexpr bool is_extractable =
(false || ... || normalize_one<BitLen, Elems>::is_extractable);
}; };
template <std::size_t BitLen> template <std::size_t BitLen>
@ -471,19 +552,26 @@ struct normalize_pack<BitLen>
static constexpr bool cmp_wild = false; static constexpr bool cmp_wild = false;
static constexpr std::size_t cmp_block = 0; static constexpr std::size_t cmp_block = 0;
static constexpr bool cmp_idcf = false; static constexpr bool cmp_idcf = false;
static constexpr bool is_verifiable = false;
static constexpr bool is_extractable = false;
}; };
/// True iff the pack is "classic-shaped": every element is a bare output (no /// @brief True iff the pack is "classic-shaped": every element is a bare output (no
/// `placed<>` from `at<>` and no `cmp_channel_tag<>`). /// `placed<>` from `at<>`, no `cmp_channel_tag<>`, and no verifiable/extractable).
template <typename ...Elems> template <typename ...Elems>
inline constexpr bool is_classic_pack_v = inline constexpr bool is_classic_pack_v =
!((is_placed_v<Elems> || is_cmp_channel_tag_v<Elems>) || ...); !((is_placed_v<Elems> || is_cmp_channel_tag_v<Elems>
|| is_verifiable_tag_v<Elems> || is_extractable_tag_v<Elems>) || ...);
} // namespace incr } // namespace incr
} // namespace detail } // namespace detail
/// Sparse heavy-hitters IDPF. Slot `i` is the point function on prefix /// @brief Sparse heavy-hitters IDPF. Slot `i` is the point function on prefix
/// `Levels[i]`, evaluated with `out<i>`. /// `Levels[i]`, evaluated with `out<i>`.
/// @tparam Levels levels
/// @tparam Betas betas
/// @param betas the `betas`
/// @return Sparse heavy-hitters IDPF
template <std::size_t ...Levels, typename ...Betas> template <std::size_t ...Levels, typename ...Betas>
auto idpf_at(Betas ...betas) auto idpf_at(Betas ...betas)
{ {
@ -500,7 +588,10 @@ auto idpf_from_seq(std::index_sequence<I...>, Betas ...betas)
return idpf_at<(I + 1)...>(std::move(betas)...); return idpf_at<(I + 1)...>(std::move(betas)...);
} }
/// Consecutive prefixes of length 1, 2, …, `sizeof...(Betas)`. /// @brief Consecutive prefixes of length 1, 2, …, `sizeof...(Betas)`.
/// @tparam Betas betas
/// @param betas the `betas`
/// @return Consecutive prefixes of length 1, 2, …, `sizeof...(Betas)`
template <typename ...Betas> template <typename ...Betas>
auto idpf(Betas ...betas) auto idpf(Betas ...betas)
{ {

View file

@ -1,6 +1,5 @@
/// @file dpf/prg.hpp /// @file dpf/prg.hpp
/// @brief /// @brief PRG aliases, the share expander, and the call counter.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -17,6 +16,7 @@
#include <type_traits> #include <type_traits>
#include "dpf/prg_aes.hpp" #include "dpf/prg_aes.hpp"
#include "dpf/prg_aes_ccr.hpp"
#include "dpf/prg_chacha.hpp" #include "dpf/prg_chacha.hpp"
#include "dpf/prg_dummy.hpp" #include "dpf/prg_dummy.hpp"
#include "dpf/prg_lowmc.hpp" #include "dpf/prg_lowmc.hpp"
@ -31,7 +31,12 @@ namespace prg
namespace detail namespace detail
{ {
/// Fill `T` from consecutive PRG blocks starting at `pos` (low bytes first). /// @brief Fill `T` from consecutive PRG blocks starting at `pos` (low bytes first).
/// @tparam PRG pseudorandom generator
/// @tparam T value type
/// @param seed the PRG seed
/// @param pos the 0-based index
/// @return the returned `T`
template <typename PRG, typename T> template <typename PRG, typename T>
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -85,6 +90,13 @@ auto lowmc128::expand(block_type seed, psnip_uint32_t pos) noexcept
return detail::expand_as_share<lowmc128, T, Party>(seed, pos); return detail::expand_as_share<lowmc128, T, Party>(seed, pos);
} }
template <typename T, std::size_t Party>
HEDLEY_NO_THROW
auto aes128_ccr::expand(block_type seed, psnip_uint32_t pos) noexcept
{
return detail::expand_as_share<aes128_ccr, T, Party>(seed, pos);
}
template <unsigned Rounds> template <unsigned Rounds>
template <typename T, std::size_t Party> template <typename T, std::size_t Party>
HEDLEY_NO_THROW HEDLEY_NO_THROW

View file

@ -1,6 +1,5 @@
/// @file dpf/prg_aes.hpp /// @file dpf/prg_aes.hpp
/// @brief /// @brief Fixed-key AES Matyas–Meyer–Oseas PRG.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -100,10 +99,15 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
} }
/// Round-major multi-block MMO. Positions use the same lane as /// @brief Round-major multi-block MMO. Positions use the same lane as
/// `eval` / `eval01` (`set_epi64x(0, pos)`). The first AddRoundKey /// `eval` / `eval01` (`set_epi64x(0, pos)`). The first AddRoundKey
/// includes `rd_key[0]` so this matches the one-block `eval` for any /// includes `rd_key[0]` so this matches the one-block `eval` for any
/// key, not only the all-zero key this PRG currently installs. /// key, not only the all-zero key this PRG currently installs.
/// @param seed the PRG seed
/// @param output the destination. Unused when `count` is 0
/// @param count the number of blocks
/// @param pos the 0-based index
/// @throws std::invalid_argument if `prg lane index is out of range`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
static void eval(block_type seed, block_type * HEDLEY_RESTRICT output, static void eval(block_type seed, block_type * HEDLEY_RESTRICT output,
psnip_uint32_t count, psnip_uint32_t pos = 0) psnip_uint32_t count, psnip_uint32_t pos = 0)
@ -164,8 +168,11 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
} }
} }
/// Four independent `eval01` calls as one 8-block round-major AES. /// @brief Four independent `eval01` calls as one 8-block round-major AES.
/// `left[i] == eval(seeds[i], 0)`, `right[i] == eval(seeds[i], 1)`. /// @details `left[i] == eval(seeds[i], 0)`, `right[i] == eval(seeds[i], 1)`.
/// @param seeds the root seeds
/// @param left the `left`
/// @param right the `right`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2, 3) HEDLEY_NON_NULL(1, 2, 3)
@ -198,7 +205,10 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
} }
} }
/// Four independent `eval(seed, pos)` as one 4-block round-major AES. /// @brief Four independent `eval(seed, pos)` as one 4-block round-major AES.
/// @param seeds the root seeds
/// @param output the destination. Unused when `count` is 0
/// @param pos the 0-based index
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2) HEDLEY_NON_NULL(1, 2)
@ -229,7 +239,10 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
} }
} }
/// Eight independent `eval(seed, pos)` as one 8-block round-major AES. /// @brief Eight independent `eval(seed, pos)` as one 8-block round-major AES.
/// @param seeds the root seeds
/// @param output the destination. Unused when `count` is 0
/// @param pos the 0-based index
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2) HEDLEY_NON_NULL(1, 2)
@ -260,7 +273,13 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
} }
} }
/// Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`). /// @brief Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`).
/// @tparam T value type
/// @tparam Party party index, `0` or `1`
/// @param seed the PRG seed
/// @param pos the 0-based index
/// @return Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`)
/// @see `prg.hpp`
template <typename T, std::size_t Party> template <typename T, std::size_t Party>
HEDLEY_NO_THROW HEDLEY_NO_THROW
static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept; static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept;
@ -268,8 +287,10 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
private: private:
static const AesKey key; static const AesKey key;
/// `blk[i]` is already `seed[i] XOR rd_key[0] XOR pos_i`. Runs AES /// @brief `blk[i]` is already `seed[i] XOR rd_key[0] XOR pos_i`. Runs AES
/// rounds 1..last and the MMO feed-forward `XOR seed[i]`. /// rounds 1..last and the MMO feed-forward `XOR seed[i]`.
/// @param blk the `blk`
/// @param seed the PRG seed
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2) HEDLEY_NON_NULL(1, 2)

254
include/dpf/prg_aes_ccr.hpp Normal file
View file

@ -0,0 +1,254 @@
/// @file dpf/prg_aes_ccr.hpp
/// @brief Circular correlation-robust (CCR) hash from fixed-key AES.
/// @details Implements the GKWY / Half-Tree CCR construction
/// `H(x) = π(σ(x)) ⊕ σ(x)` where `π` is the library's fixed-key AES
/// (same schedule as `prg::aes128`) and `σ` is the bitstring linear
/// orthomorphism `σ(xL∥xR) = (xL⊕xR)∥xL`.
///
/// `eval01(s)` returns the Half-Tree children `{H(s), H(s)⊕s}`. That
/// expand is **not** a drop-in for BGI `make_dpf`; use it only with
/// Half-Tree `tree_traits` (see `dpf/tree_traits.hpp`).
/// @see dpf/tree_traits.hpp
/// @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_PRG_AES_CCR_HPP__
#define LIBDPF_INCLUDE_DPF_PRG_AES_CCR_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "portable-snippets/exact-int/exact-int.h"
#include "dpf/prg_aes.hpp"
#include "dpf/twiddle.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
namespace prg
{
namespace ccr_detail
{
/// @brief Linear orthomorphism on 128-bit strings: `σ(xL∥xR) = (xL⊕xR)∥xL`.
/// @param x the `x`
/// @return Linear orthomorphism on 128-bit strings: `σ(xL∥xR) = (xL⊕xR)∥xL`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
simde__m128i sigma(simde__m128i x) noexcept
{
const std::uint64_t lo = static_cast<std::uint64_t>(
simde_mm_cvtsi128_si64(x));
const std::uint64_t hi = static_cast<std::uint64_t>(
simde_mm_extract_epi64(x, 1));
return simde_mm_set_epi64x(static_cast<long long>(lo),
static_cast<long long>(lo ^ hi));
}
/// @brief Inverse: if `σ(x)=(a,b)=(xL⊕xR, xL)` then `xL=b`, `xR=a⊕b`.
/// @param y the `y`
/// @return Inverse: if `σ(x)=(a,b)=(xL⊕xR, xL)` then `xL=b`, `xR=a⊕b`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
simde__m128i sigma_inv(simde__m128i y) noexcept
{
const std::uint64_t a = static_cast<std::uint64_t>(
simde_mm_cvtsi128_si64(y));
const std::uint64_t b = static_cast<std::uint64_t>(
simde_mm_extract_epi64(y, 1));
return simde_mm_set_epi64x(static_cast<long long>(a ^ b),
static_cast<long long>(b));
}
/// @brief `σ'(x) = σ(x) ⊕ x`.
/// @param x the `x`
/// @return `σ'(x) = σ(x) ⊕ x`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
simde__m128i sigma_prime(simde__m128i x) noexcept
{
return simde_mm_xor_si128(sigma(x), x);
}
} // namespace ccr_detail
/// @brief CCR hash + Half-Tree expand over fixed-key AES-128.
/// @details Selecting this as an *interior* PRG opts into Half-Tree via `half_tree_tag`.
struct aes128_ccr final
{
using block_type = simde__m128i;
using half_tree_tag = void;
using underlying_aes = aes128;
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static void require_block_aligned(const void * p) noexcept
{
underlying_aes::require_block_aligned(p);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static block_type sigma(block_type x) noexcept
{
return ccr_detail::sigma(x);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static block_type sigma_inv(block_type y) noexcept
{
return ccr_detail::sigma_inv(y);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static block_type sigma_prime(block_type x) noexcept
{
return ccr_detail::sigma_prime(x);
}
/// @brief `H(x) = π(σ(x)) ⊕ σ(x)` with `π` the fixed-key AES permutation
/// (implemented as the library MMO at position 0).
/// @param x the `x`
/// @return `H(x) = π(σ(x)) ⊕ σ(x)` with `π` the fixed-key AES permutation (implemented as the
/// library MMO at position 0)
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
static block_type hash(block_type x) noexcept
{
return underlying_aes::eval(sigma(x), 0);
}
/// @brief Alias for `hash`.
/// @param x the `x`
/// @return Alias for `hash`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
static block_type H(block_type x) noexcept
{
return hash(x);
}
/// @brief Position-tweaked CCR hash: `H(x; pos) = AES_MMO(σ(x), pos)`.
/// @param seed the PRG seed
/// @param pos the 0-based index
/// @return Position-tweaked CCR hash: `H(x; pos) = AES_MMO(σ(x), pos)`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
static block_type eval(block_type seed, psnip_uint32_t pos) noexcept
{
return underlying_aes::eval(sigma(seed), pos);
}
/// @brief Half-Tree children: left = `H(s)`, right = `H(s) ⊕ s`.
/// @param seed the PRG seed
/// @return Half-Tree children: left = `H(s)`, right = `H(s) ⊕ s`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
static auto eval01(block_type seed) noexcept
{
const block_type h = hash(seed);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::array<block_type, 2>{h, simde_mm_xor_si128(h, seed)};
HEDLEY_PRAGMA(GCC diagnostic pop)
}
/// @brief Last-level two-tweak stretch: `{H(s|0), H(s|1)}` (LSB forced).
/// @param seed the PRG seed
/// @return Last-level two-tweak stretch: `{H(s|0), H(s|1)}` (LSB forced)
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
static auto eval01_twotweak(block_type seed) noexcept
{
const block_type base = dpf::unset_lo_bit(seed);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::array<block_type, 2>{
hash(base),
hash(dpf::set_lo_bit(base))
};
HEDLEY_PRAGMA(GCC diagnostic pop)
}
/// @brief `count` CCR blocks starting at lane `pos`.
/// @details `output` is unused when `count` is 0. Each block is
/// `AES_MMO(σ(seed), pos + i)`.
/// @param seed the PRG seed
/// @param output the destination. Unused when `count` is 0
/// @param count the number of blocks
/// @param pos the first lane index
/// @throws std::invalid_argument if `pos + count` wraps `uint32_t`
HEDLEY_ALWAYS_INLINE
static void eval(block_type seed, block_type * HEDLEY_RESTRICT output,
psnip_uint32_t count, psnip_uint32_t pos = 0)
{
underlying_aes::eval(sigma(seed), output, count, pos);
}
/// @brief CCR hash of four seeds.
/// @param inputs four seeds
/// @param output four hashes, `H(inputs[i])`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2)
static void hash_x4(const block_type * HEDLEY_RESTRICT inputs,
block_type * HEDLEY_RESTRICT output) noexcept
{
require_block_aligned(inputs);
require_block_aligned(output);
alignas(block_type) block_type sx[4];
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
sx[i] = sigma(inputs[i]);
underlying_aes::eval_x4(sx, output, 0);
}
/// @brief Half-Tree children of four seeds.
/// @param seeds four seeds
/// @param left `H(seeds[i])`
/// @param right `H(seeds[i]) ⊕ seeds[i]`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2, 3)
static void eval01_x4(const block_type * HEDLEY_RESTRICT seeds,
block_type * HEDLEY_RESTRICT left,
block_type * HEDLEY_RESTRICT right) noexcept
{
require_block_aligned(seeds);
require_block_aligned(left);
require_block_aligned(right);
hash_x4(seeds, left);
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
right[i] = simde_mm_xor_si128(left[i], seeds[i]);
}
template <typename T, std::size_t Party>
HEDLEY_NO_THROW
static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept;
};
} // namespace prg
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_PRG_AES_CCR_HPP__

View file

@ -37,7 +37,7 @@ namespace chacha_detail
inline constexpr std::uint32_t zero_nonce[3] = {0, 0, 0}; inline constexpr std::uint32_t zero_nonce[3] = {0, 0, 0};
/// ASCII `"dpf-chacha-prg"` plus two zero bytes. Public second half of the key. /// @brief ASCII `"dpf-chacha-prg"` plus two zero bytes. Public second half of the key.
inline constexpr std::uint8_t domain[16] = { inline constexpr std::uint8_t domain[16] = {
'd', 'p', 'f', '-', 'c', 'h', 'a', 'c', 'd', 'p', 'f', '-', 'c', 'h', 'a', 'c',
'h', 'a', '-', 'p', 'r', 'g', 0, 0 'h', 'a', '-', 'p', 'r', 'g', 0, 0
@ -77,7 +77,10 @@ simde__m128i load_block(const std::uint8_t * p) noexcept
return out; return out;
} }
/// 128-bit seed in the low half, `domain` in the high half, both little-endian. /// @brief 128-bit seed in the low half, `domain` in the high half, both little-endian.
/// @param seed the PRG seed
/// @param key the `key`
HEDLEY_NON_NULL(2)
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
void seed_key(simde__m128i seed, std::uint32_t key[8]) noexcept void seed_key(simde__m128i seed, std::uint32_t key[8]) noexcept
@ -92,8 +95,9 @@ void seed_key(simde__m128i seed, std::uint32_t key[8]) noexcept
} }
template <int N> template <int N>
HEDLEY_ALWAYS_INLINE HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr std::uint32_t rotl(std::uint32_t x) noexcept constexpr std::uint32_t rotl(std::uint32_t x) noexcept
{ {
static_assert(N > 0 && N < 32, "ChaCha rotation is between 1 and 31"); static_assert(N > 0 && N < 32, "ChaCha rotation is between 1 and 31");
@ -112,8 +116,9 @@ void quarter(std::uint32_t & a, std::uint32_t & b,
} }
template <int N> template <int N>
HEDLEY_ALWAYS_INLINE HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
simde__m128i rotl_epi32(simde__m128i v) noexcept simde__m128i rotl_epi32(simde__m128i v) noexcept
{ {
static_assert(N > 0 && N < 32, "ChaCha rotation is between 1 and 31"); static_assert(N > 0 && N < 32, "ChaCha rotation is between 1 and 31");
@ -121,6 +126,7 @@ simde__m128i rotl_epi32(simde__m128i v) noexcept
simde_mm_srli_epi32(v, 32 - N)); simde_mm_srli_epi32(v, 32 - N));
} }
HEDLEY_NON_NULL(1)
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
void quarter(simde__m128i x[], int a, int b, int c, int d) noexcept void quarter(simde__m128i x[], int a, int b, int c, int d) noexcept
@ -135,6 +141,7 @@ void quarter(simde__m128i x[], int a, int b, int c, int d) noexcept
x[b] = rotl_epi32<7>(simde_mm_xor_si128(x[b], x[c])); x[b] = rotl_epi32<7>(simde_mm_xor_si128(x[b], x[c]));
} }
HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
std::uint32_t epi32_lane(simde__m128i v, int lane) noexcept std::uint32_t epi32_lane(simde__m128i v, int lane) noexcept
@ -150,8 +157,21 @@ std::uint32_t epi32_lane(simde__m128i v, int lane) noexcept
return static_cast<std::uint32_t>(simde_mm_cvtsi128_si32(v)); return static_cast<std::uint32_t>(simde_mm_cvtsi128_si32(v));
} }
/// One ChaCha block. `key` is 8 little-endian words. `nonce` is 3 words. /// @name ChaCha blocks
/// @tparam Rounds ChaCha round count. Must be positive and even
/// @param key the ChaCha key words
/// @param counter the ChaCha block counter
/// @param out the output buffer
/// @{
/// @brief One ChaCha block.
/// @details `key` is 8 little-endian words. `nonce` is 3 words.
/// @param key the ChaCha key words
/// @param counter the ChaCha block counter
/// @param nonce the ChaCha nonce
/// @param out the output buffer
template <unsigned Rounds> template <unsigned Rounds>
HEDLEY_NON_NULL(1, 3, 4)
HEDLEY_NO_THROW HEDLEY_NO_THROW
void block(const std::uint32_t key[8], std::uint32_t counter, void block(const std::uint32_t key[8], std::uint32_t counter,
const std::uint32_t nonce[3], std::uint8_t out[64]) noexcept const std::uint32_t nonce[3], std::uint8_t out[64]) noexcept
@ -187,9 +207,10 @@ HEDLEY_PRAGMA(GCC unroll 16)
} }
} }
/// Four independent ChaCha blocks. Lane `i` uses `key[i]` and `counter[i]`. /// @brief Four independent ChaCha blocks.
/// Nonce is zero. Each `out[i]` receives 64 bytes. /// @details Lane `i` uses `key[i]` and `counter[i]`. Nonce is zero. Each `out[i]` receives 64 bytes.
template <unsigned Rounds> template <unsigned Rounds>
HEDLEY_NON_NULL(1, 2, 3)
HEDLEY_NO_THROW HEDLEY_NO_THROW
void block4(const std::uint32_t key[][8], const std::uint32_t counter[4], void block4(const std::uint32_t key[][8], const std::uint32_t counter[4],
std::uint8_t out[][64]) noexcept std::uint8_t out[][64]) noexcept
@ -248,9 +269,12 @@ HEDLEY_PRAGMA(GCC unroll 16)
} }
} }
/// @}
} // namespace chacha_detail } // namespace chacha_detail
/// ChaCha stream PRG with `Rounds` rounds (20 is RFC 8439). /// @brief ChaCha stream PRG with `Rounds` rounds (20 is RFC 8439).
/// @tparam Rounds ChaCha round count. Must be positive and even
template <unsigned Rounds = 20> template <unsigned Rounds = 20>
struct chacha final struct chacha final
{ {
@ -272,7 +296,9 @@ struct chacha final
return chacha_detail::load_block(buf + 16 * (pos & 3u)); return chacha_detail::load_block(buf + 16 * (pos & 3u));
} }
/// Positions 0 and 1, one ChaCha block (the first 32 keystream bytes). /// @brief Positions 0 and 1, one ChaCha block (the first 32 keystream bytes).
/// @param seed the PRG seed
/// @return Positions 0 and 1, one ChaCha block (the first 32 keystream bytes)
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
static auto eval01(block_type seed) noexcept static auto eval01(block_type seed) noexcept
@ -290,6 +316,12 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
} }
/// @brief `count` blocks starting at lane `pos`. `output` is unused when `count` is 0.
/// @param seed the PRG seed
/// @param output the destination. Unused when `count` is 0
/// @param count the number of blocks
/// @param pos the 0-based index
/// @throws std::invalid_argument if `pos + count` wraps `uint32_t`.
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
static void eval(block_type seed, block_type * HEDLEY_RESTRICT output, static void eval(block_type seed, block_type * HEDLEY_RESTRICT output,
psnip_uint32_t count, psnip_uint32_t pos = 0) psnip_uint32_t count, psnip_uint32_t pos = 0)
@ -436,19 +468,25 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
eval_x4(seeds + 4, output + 4, pos); eval_x4(seeds + 4, output + 4, pos);
} }
/// Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`). /// @brief Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`).
/// @tparam T value type
/// @tparam Party party index, `0` or `1`
/// @param seed the PRG seed
/// @param pos the 0-based index
/// @return Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`)
/// @see `prg.hpp`
template <typename T, std::size_t Party> template <typename T, std::size_t Party>
HEDLEY_NO_THROW HEDLEY_NO_THROW
static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept; static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept;
}; // struct chacha }; // struct chacha
/// RFC 8439 ChaCha20. /// @brief RFC 8439 ChaCha20.
using chacha20 = chacha<20>; using chacha20 = chacha<20>;
/// ChaCha12. Same keying as `chacha20`, 12 rounds. /// @brief ChaCha12. Same keying as `chacha20`, 12 rounds.
using chacha12 = chacha<12>; using chacha12 = chacha<12>;
/// ChaCha8. Same keying as `chacha20`, 8 rounds. /// @brief ChaCha8. Same keying as `chacha20`, 8 rounds.
using chacha8 = chacha<8>; using chacha8 = chacha<8>;
} // namespace prg } // namespace prg

View file

@ -88,7 +88,13 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
std::copy_n(seeds, 8, output); std::copy_n(seeds, 8, output);
} }
/// Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`). /// @brief Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`).
/// @tparam T value type
/// @tparam Party party index, `0` or `1`
/// @param seed the PRG seed
/// @param pos the 0-based index
/// @return Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`)
/// @see `prg.hpp`
template <typename T, std::size_t Party> template <typename T, std::size_t Party>
HEDLEY_NO_THROW HEDLEY_NO_THROW
static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept; static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept;

View file

@ -26,8 +26,8 @@ namespace dpf
namespace prg namespace prg
{ {
/// LowMCv3, 128-bit block and key, 10 S-boxes, 32 rounds, all-zero key. /// @brief LowMCv3, 128-bit block and key, 10 S-boxes, 32 rounds, all-zero key.
/// `eval(seed, pos)` is `E(seed ⊕ pos) ⊕ seed`, with `pos` in the low lane. /// @details `eval(seed, pos)` is `E(seed ⊕ pos) ⊕ seed`, with `pos` in the low lane.
struct lowmc128 final struct lowmc128 final
{ {
using block_type = simde__m128i; using block_type = simde__m128i;
@ -104,7 +104,13 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
} }
} }
/// Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`). /// @brief Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`).
/// @tparam T value type
/// @tparam Party party index, `0` or `1`
/// @param seed the PRG seed
/// @param pos the 0-based index
/// @return Raw-bit subtractive share of `T` for party `Party` (see `prg.hpp`)
/// @see `prg.hpp`
template <typename T, std::size_t Party> template <typename T, std::size_t Party>
HEDLEY_NO_THROW HEDLEY_NO_THROW
static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept; static auto expand(block_type seed, psnip_uint32_t pos = 0) noexcept;

View file

@ -1,6 +1,5 @@
/// @file dpf/random.hpp /// @file dpf/random.hpp
/// @brief /// @brief Entropy source and uniform sampling.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -32,7 +31,7 @@ namespace dpf
namespace detail namespace detail
{ {
/// When set, `uniform_fill` copies from this hook and does not read the /// @brief When set, `uniform_fill` copies from this hook and does not read the
/// system RNG. Used to feed the same beaver coins to dealer `make_dpf` and /// system RNG. Used to feed the same beaver coins to dealer `make_dpf` and
/// Doerner–Shelat gen. Null in normal use. /// Doerner–Shelat gen. Null in normal use.
inline thread_local void (*uniform_bytes_hook)(void *, std::size_t) = nullptr; inline thread_local void (*uniform_bytes_hook)(void *, std::size_t) = nullptr;
@ -50,8 +49,10 @@ bool fill_from_hook(T & buf) noexcept
return true; return true;
} }
/// `bool` and `enum : bool` (including `dpf::bit`) have only two valid /// @brief `bool` and `enum : bool` (including `dpf::bit`) have only two valid
/// representations. Filling them with a raw entropy byte is undefined. /// representations. Filling them with a raw entropy byte is undefined.
/// @tparam T value type
/// @return `bool` and `enum : bool` (including `dpf::bit`) have only two valid representations
template <typename T> template <typename T>
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr bool is_boolean_representation() noexcept constexpr bool is_boolean_representation() noexcept
@ -73,8 +74,8 @@ constexpr bool is_boolean_representation() noexcept
#if !defined(LIBDPF_USE_ARC4RANDOM) #if !defined(LIBDPF_USE_ARC4RANDOM)
/// One unbuffered, exclusively locked read of the entropy device. /// @brief One unbuffered, exclusively locked read of the entropy device.
/// Buffering would copy unread bytes into a `fork()` child, so parent and /// @details Buffering would copy unread bytes into a `fork()` child, so parent and
/// child would repeat the same key material. The lock keeps concurrent /// child would repeat the same key material. The lock keeps concurrent
/// `fread` calls off the shared `FILE`. /// `fread` calls off the shared `FILE`.
struct entropy_source struct entropy_source

View file

@ -1,245 +0,0 @@
/// @file dpf/rotated_iterable.hpp
/// @brief Retired container rotation view. The live type is `rotation_iterable`.
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license.
// /// @file dpf/rotated_iterable.hpp
// /// @author Ryan Henry <ryan.henry@ucalgary.ca>
// /// @brief defines `dpf::rotated_iterable` and associated helpers
// /// @details
// /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
// /// @license Released under a GNU General Public v2.0 (GPLv2) license;
// /// see [LICENSE.md](@ref GPLv2) for details.
// #ifndef LIBDPF_INCLUDE_DPF_ROTATED_VIEW_HPP__
// #define LIBDPF_INCLUDE_DPF_ROTATED_VIEW_HPP__
// namespace dpf
// {
// template <typename ContainerT>
// struct rotated_iterable_iterator; // forward declaration
// template <typename ContainerT>
// struct rotated_iterable_const_iterator; // forward declaration
// template <typename ContainerT>
// struct rotated_iterable
// {
// using container_type = ContainerT;
// using value_type = typename container_type::value_type;
// using size_type = typename container_type::size_type;
// using difference_type = typename container_type::difference_type;
// using reference = typename container_type::reference;
// using const_reference = typename container_type::const_reference;
// using pointer = typename container_type::pointer;
// using const_pointer = typename container_type::const_pointer;
// using iterator = rotated_iterable_iterator<container_type>;
// using const_iterator = rotated_iterable_const_iterator<container_type>;
// using wrapped_iterator = typename ContainerT::iterator;
// rotated_iterable(const ContainerT & container, difference_type distance)
// : container_{container},
// distance_{distance >= 0 ? distance % container_.size() : (distance % container_.size()) + container_.size()},
// wrap_to{std::begin(container)},
// wrap_after{std::next(std::end(container), -1)},
// end_after{std::next(wrap_to, distance-1)}
// {
// distance_ %= container_.size();
// if (distance_ < 0)
// {
// distance_ += container_.size();
// }
// }
// HEDLEY_ALWAYS_INLINE
// reference operator[](size_type index)
// {
// index += distance_;
// if (index > container_.size())
// {
// index -= container_.size();
// }
// return container_[index];
// }
// HEDLEY_ALWAYS_INLINE
// const_reference operator[](size_type index) const
// {
// index += distance_;
// if (index > container_.size())
// {
// index -= container_.size();
// }
// return container_[index];
// }
// HEDLEY_NO_THROW
// HEDLEY_ALWAYS_INLINE
// iterator begin() noexcept
// {
// return iterator{*this, std::next(end_after, 1)};
// }
// HEDLEY_NO_THROW
// HEDLEY_ALWAYS_INLINE
// const_iterator begin() const noexcept
// {
// return const_iterator{*this, std::next(end_after, 1)};
// }
// HEDLEY_NO_THROW
// HEDLEY_ALWAYS_INLINE
// const_iterator cbegin() const noexcept
// {
// return begin();
// }
// HEDLEY_NO_THROW
// HEDLEY_ALWAYS_INLINE
// iterator end() noexcept
// {
// return iterator{*this, std::next(wrap_after, 1)};
// }
// HEDLEY_NO_THROW
// HEDLEY_ALWAYS_INLINE
// const_iterator end() const noexcept
// {
// return const_iterator{*this, std::next(wrap_after, 1)};
// }
// HEDLEY_NO_THROW
// HEDLEY_ALWAYS_INLINE
// const_iterator cend() const noexcept
// {
// return end();
// }
// auto distance() const
// {
// return distance_;
// }
// private:
// container_type & container_;
// difference_type distance_;
// wrapped_iterator wrap_to;
// wrapped_iterator wrap_after;
// wrapped_iterator end_after;
// }; // rotated_iterable
// template <typename ContainerT>
// struct rotated_iterator
// {
// using wrapped_iterable_type = rotated_iterable<ContainerT>;
// using wrapped_iterator = typename wrapped_iterable_type::iterator;
// using size_type = typename wrapped_iterable_type::size_type;
// using reference = typename wrapped_iterable_type::reference;
// rotated_iterable<ContainerT> & v;
// wrapped_iterator it;
// rotated_iterator & operator++()
// {
// if (it == v.wrap_after)
// {
// it = v.wrap_to;
// }
// else if (it == v.end_after)
// {
// it = std::next(v.wrap_after, 1);
// }
// else
// {
// ++it;
// }
// return *this;
// }
// rotated_iterator & operator--()
// {
// --it;
// if (it == v.wrap_after)
// {
// it = v.end_after;
// }
// else if (it == v.end_after)
// {
// it = std::next(v.wrap_to, -1);
// }
// return *this;
// }
// reference operator*() const { return *it; }
// bool operator!=(const rotated_iterator other) const
// { return &v != &other.v || it != other.it; }
// };
// template <typename ContainerT>
// struct rotated_const_iterator
// {
// using wrapped_iterable_type = rotated_iterable<ContainerT>;
// using wrapped_iterator = typename wrapped_iterable_type::iterator;
// using size_type = typename wrapped_iterable_type::size_type;
// using const_reference = typename wrapped_iterable_type::const_reference;
// const rotated_iterable<ContainerT> & v;
// wrapped_iterator it;
// rotated_const_iterator & operator++()
// {
// if (it == v.wrap_after)
// {
// it = v.wrap_to;
// }
// else if (it == v.end_after)
// {
// it = std::next(v.wrap_after, 1);
// }
// else
// {
// ++it;
// }
// return *this;
// }
// rotated_const_iterator & operator--()
// {
// --it;
// if (it == v.wrap_after)
// {
// it = v.end_after;
// }
// else if (it == v.end_after)
// {
// it = std::next(v.wrap_to, -1);
// }
// return *this;
// }
// const_reference operator*() const { return *it; }
// bool operator!=(const rotated_const_iterator other) const
// { return &v != &other.v || it != other.it; }
// };
// template <typename ContainerT>
// auto rotated_by(const ContainerT & container,
// typename ContainerT::size_type rotate_by)
// {
// return rotated_iterable{container, rotate_by};
// }
// template <typename ContainerT,
// typename UnaryFunction>
// auto for_each_rotated_by(const ContainerT & container,
// typename ContainerT::size_type rotate_by, UnaryFunction && f)
// {
// for (auto i = rotate_by; i < container.size(); ++i) f(container[i]);
// for (auto i = 0; i < rotate_by; ++i) f(container[i]);
// }
// } // namespace dpf
// #endif // LIBDPF_INCLUDE_DPF_ROTATED_VIEW_HPP__

View file

@ -26,7 +26,7 @@
namespace dpf namespace dpf
{ {
/// Sharing scheme tag. /// @brief Sharing scheme tag.
enum class sharing : unsigned char enum class sharing : unsigned char
{ {
additive = 0, additive = 0,
@ -94,8 +94,14 @@ using share_value_type_t = typename share_value_type<std::decay_t<T>>::type;
namespace detail namespace detail
{ {
/// Party coefficient of the secret for this scheme: additive always +1; /// @brief Party coefficient of the secret for this scheme: additive always +1;
/// subtractive is +1 for party 0 and −1 for party 1. /// subtractive is +1 for party 0 and −1 for party 1.
/// @tparam Scheme scheme
/// @tparam Party party index, `0` or `1`
/// @tparam T value type
/// @param v the `v`
/// @return Party coefficient of the secret for this scheme: additive always +1; subtractive is +1
/// for party 0 and −1 for party 1
template <sharing Scheme, std::size_t Party, typename T> template <sharing Scheme, std::size_t Party, typename T>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
@ -141,7 +147,9 @@ struct secret_share
secret_share & operator=(secret_share &&) noexcept = default; secret_share & operator=(secret_share &&) noexcept = default;
~secret_share() = default; ~secret_share() = default;
/// Bit-preserving construction. Does not apply a party coefficient. /// @brief Bit-preserving construction. Does not apply a party coefficient.
/// @param v the `v`
/// @return Bit-preserving construction
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_CONST HEDLEY_CONST
@ -162,7 +170,8 @@ struct secret_share
HEDLEY_PURE HEDLEY_PURE
constexpr T & raw() noexcept { return value; } constexpr T & raw() noexcept { return value; }
/// Secret-preserving conversion to an additive share of the same party. /// @brief Secret-preserving conversion to an additive share of the same party.
/// @return Secret-preserving conversion to an additive share of the same party
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
@ -175,7 +184,8 @@ struct secret_share
detail::party_coeff_times<sharing::subtractive, Party>(value)); detail::party_coeff_times<sharing::subtractive, Party>(value));
} }
/// Secret-preserving conversion to a subtractive share of the same party. /// @brief Secret-preserving conversion to a subtractive share of the same party.
/// @return Secret-preserving conversion to a subtractive share of the same party
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
@ -188,7 +198,10 @@ struct secret_share
detail::party_coeff_times<sharing::additive, Party>(value)); detail::party_coeff_times<sharing::additive, Party>(value));
} }
/// Bit-preserving retag (no secret-preserving sign fix). /// @brief Bit-preserving retag (no secret-preserving sign fix).
/// @tparam NewScheme new scheme
/// @tparam NewParty new party
/// @return Bit-preserving retag (no secret-preserving sign fix)
template <sharing NewScheme, std::size_t NewParty = Party> template <sharing NewScheme, std::size_t NewParty = Party>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
@ -239,7 +252,11 @@ struct secret_share
return *this; return *this;
} }
/// Absorb a public plaintext on party 0 only. /// @brief Absorb a public plaintext on party 0 only.
/// @tparam Plain plain
/// @tparam T value type
/// @param c the `c`
/// @return `*this`
template <typename Plain, template <typename Plain,
std::enable_if_t<!is_secret_share_v<Plain> std::enable_if_t<!is_secret_share_v<Plain>
&& std::is_convertible_v<Plain, T>, int> = 0> && std::is_convertible_v<Plain, T>, int> = 0>
@ -525,7 +542,8 @@ struct party_key : Key
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
const Key & key() const noexcept { return static_cast<const Key &>(*this); } const Key & key() const noexcept { return static_cast<const Key &>(*this); }
/// Party-tagged additive share of the comparison absorb addend. /// @brief Party-tagged additive share of the comparison absorb addend.
/// @return Party-tagged additive share of the comparison absorb addend
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
auto cmp_addend() const noexcept auto cmp_addend() const noexcept
@ -555,9 +573,10 @@ struct party_of<party_key<Party, Key>>
template <typename T> template <typename T>
inline constexpr std::size_t party_of_v = party_of<std::decay_t<T>>::value; inline constexpr std::size_t party_of_v = party_of<std::decay_t<T>>::value;
/// Strip a `party_key` wrapper; bare keys are unchanged. Memoizers and other /// @brief Strip a `party_key` wrapper; bare keys are unchanged. Memoizers and other
/// tree-layout helpers key on the underlying DPF key type so a memoizer built /// tree-layout helpers key on the underlying DPF key type so a memoizer built
/// for party 0 also accepts party 1. /// for party 0 also accepts party 1.
/// @tparam T value type
template <typename T> template <typename T>
struct unwrap_party_key struct unwrap_party_key
{ {

View file

@ -285,8 +285,10 @@ struct pointer_facade
} // namespace detail } // namespace detail
/// One level. The buffer is traversed in the opposite direction on /// @brief One level. The buffer is traversed in the opposite direction on
/// alternate levels. /// alternate levels.
/// @tparam DpfKey DPF key type
/// @tparam Allocator allocator type
template <typename DpfKey, template <typename DpfKey,
typename Allocator = aligned_allocator<typename DpfKey::interior_node>> typename Allocator = aligned_allocator<typename DpfKey::interior_node>>
struct inplace_reversing_sequence_memoizer final struct inplace_reversing_sequence_memoizer final
@ -381,18 +383,20 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
unique_ptr buf; unique_ptr buf;
}; };
/// Two levels, so a level can be built while the previous level is still /// @brief Two levels, so a level can be built while the previous level is still
/// intact. Default workspace for `eval_sequence` on a recipe. /// intact. Default workspace for `eval_sequence` on a recipe.
/// @tparam DpfKey DPF key type
/// @tparam Allocator allocator type
template <typename DpfKey, template <typename DpfKey,
typename Allocator = aligned_allocator<typename DpfKey::interior_node>> typename Allocator = aligned_allocator<typename DpfKey::interior_node>>
struct double_space_sequence_memoizer final struct double_space_sequence_memoizer final
: public sequence_recipe_memoizer_base<DpfKey> : public sequence_recipe_memoizer_base<DpfKey>
{ {
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
private: private:
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using parent = sequence_recipe_memoizer_base<DpfKey>; using parent = sequence_recipe_memoizer_base<DpfKey>;
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
public: public:
using unique_ptr = typename Allocator::unique_ptr; using unique_ptr = typename Allocator::unique_ptr;
using return_type = typename DpfKey::interior_node *; using return_type = typename DpfKey::interior_node *;
@ -435,17 +439,19 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
unique_ptr buf; unique_ptr buf;
}; };
/// Every level of the recipe's traversal. /// @brief Every level of the recipe's traversal.
/// @tparam DpfKey DPF key type
/// @tparam Allocator allocator type
template <typename DpfKey, template <typename DpfKey,
typename Allocator = aligned_allocator<typename DpfKey::interior_node>> typename Allocator = aligned_allocator<typename DpfKey::interior_node>>
struct full_tree_sequence_memoizer final struct full_tree_sequence_memoizer final
: public sequence_recipe_memoizer_base<DpfKey> : public sequence_recipe_memoizer_base<DpfKey>
{ {
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
private: private:
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using parent = sequence_recipe_memoizer_base<DpfKey>; using parent = sequence_recipe_memoizer_base<DpfKey>;
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
public: public:
using unique_ptr = typename Allocator::unique_ptr; using unique_ptr = typename Allocator::unique_ptr;
using return_type = typename DpfKey::interior_node *; using return_type = typename DpfKey::interior_node *;
@ -502,10 +508,12 @@ auto make_sequence_memoizer(const sequence_recipe & recipe)
HEDLEY_PRAGMA(GCC diagnostic push) HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes") HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
/// One-level sequence workspace bound to `recipe`. /// @brief One-level sequence workspace bound to `recipe`.
/// @tparam DpfKey DPF key type
/// @param recipe The object later passed to `eval_sequence`. The memoizer /// @param recipe The object later passed to `eval_sequence`. The memoizer
/// holds a reference to it. /// holds a reference to it.
/// @snippet evaluation/memoizers.cpp sequence-memoizer /// @snippet evaluation/memoizers.cpp sequence-memoizer
/// @return One-level sequence workspace bound to `recipe`
template <typename DpfKey> template <typename DpfKey>
inline auto make_inplace_reversing_sequence_memoizer(const sequence_recipe & recipe) inline auto make_inplace_reversing_sequence_memoizer(const sequence_recipe & recipe)
{ {
@ -519,8 +527,11 @@ inline auto make_inplace_reversing_sequence_memoizer(const DpfKey &, const seque
return make_inplace_reversing_sequence_memoizer<DpfKey>(recipe); return make_inplace_reversing_sequence_memoizer<DpfKey>(recipe);
} }
/// Two-level sequence workspace bound to `recipe`. /// @brief Two-level sequence workspace bound to `recipe`.
/// @snippet evaluation/eval_sequence.cpp eval-sequence-recipe /// @snippet evaluation/eval_sequence.cpp eval-sequence-recipe
/// @tparam DpfKey DPF key type
/// @param recipe the sequence recipe the memoizer was built from
/// @return Two-level sequence workspace bound to `recipe`
template <typename DpfKey> template <typename DpfKey>
inline auto make_double_space_sequence_memoizer(const sequence_recipe & recipe) inline auto make_double_space_sequence_memoizer(const sequence_recipe & recipe)
{ {
@ -534,20 +545,23 @@ inline auto make_double_space_sequence_memoizer(const DpfKey &, const sequence_r
return make_double_space_sequence_memoizer<DpfKey>(recipe); return make_double_space_sequence_memoizer<DpfKey>(recipe);
} }
/// Full-tree sequence workspace bound to `recipe`. /// @brief Full-tree sequence workspace bound to `recipe`.
/// @tparam DpfKey DPF key type
/// @param recipe the sequence recipe the memoizer was built from
/// @return Full-tree sequence workspace bound to `recipe`
template <typename DpfKey> template <typename DpfKey>
inline auto make_full_tree_sequence_memoizer(const sequence_recipe & recipe) inline auto make_full_tree_sequence_memoizer(const sequence_recipe & recipe)
{ {
using key_t = unwrap_party_key_t<DpfKey>; using key_t = unwrap_party_key_t<DpfKey>;
return detail::make_sequence_memoizer<full_tree_sequence_memoizer<key_t>>(recipe); return detail::make_sequence_memoizer<full_tree_sequence_memoizer<key_t>>(recipe);
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
template <typename DpfKey> template <typename DpfKey>
inline auto make_full_tree_sequence_memoizer(const DpfKey &, const sequence_recipe & recipe) inline auto make_full_tree_sequence_memoizer(const DpfKey &, const sequence_recipe & recipe)
{ {
return make_full_tree_sequence_memoizer<DpfKey>(recipe); return make_full_tree_sequence_memoizer<DpfKey>(recipe);
} }
HEDLEY_PRAGMA(GCC diagnostic pop)
} // namespace dpf } // namespace dpf

View file

@ -26,7 +26,7 @@
namespace dpf namespace dpf
{ {
/// Steps, leaf count, and per-level endpoints for one sorted point list. /// @brief Steps, leaf count, and per-level endpoints for one sorted point list.
struct sequence_recipe struct sequence_recipe
{ {
public: public:
@ -52,7 +52,8 @@ struct sequence_recipe
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr const std::vector<std::size_t> & level_endpoints() const noexcept { return level_endpoints_; } constexpr const std::vector<std::size_t> & level_endpoints() const noexcept { return level_endpoints_; }
/// `level_endpoints().size() - 1`. Not `constexpr`: `std::vector::size` is not a constant expression in C++17. /// @brief `level_endpoints().size() - 1`. Not `constexpr`: `std::vector::size` is not a constant expression in C++17.
/// @return `level_endpoints().size() - 1`
HEDLEY_PURE HEDLEY_PURE
HEDLEY_NO_THROW HEDLEY_NO_THROW
std::size_t depth() const noexcept { return level_endpoints_.size()-1; } std::size_t depth() const noexcept { return level_endpoints_.size()-1; }
@ -139,9 +140,13 @@ auto make_sequence_recipe(ForwardIterator begin, ForwardIterator end)
} // namespace detail } // namespace detail
/// Compile `[begin, end)` into a recipe for `DpfKey`'s input type. /// @brief Compile `[begin, end)` into a recipe for `DpfKey`'s input type.
/// @tparam DpfKey Key type, or a `party_key` of that key. Only the input /// @tparam DpfKey Key type, or a `party_key` of that key. Only the input
/// type and depth are used. /// type and depth are used.
/// @tparam ForwardIterator forward iterator type
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @return Compile `[begin, end)` into a recipe for `DpfKey`'s input type
/// @throws std::runtime_error if the range is not sorted nondecreasing. /// @throws std::runtime_error if the range is not sorted nondecreasing.
template <typename DpfKey, template <typename DpfKey,
typename ForwardIterator> typename ForwardIterator>
@ -157,8 +162,17 @@ auto make_sequence_recipe(const DpfKey &, ForwardIterator begin, ForwardIterator
return make_sequence_recipe<DpfKey>(begin, end); return make_sequence_recipe<DpfKey>(begin, end);
} }
/// Build a sequence recipe that stops at `StopLevel` with packing `LgOpl` /// @brief Build a sequence recipe that stops at `StopLevel` with packing `LgOpl`
/// (multi-level / `out<I>` slots). Lane points are in the slot's prefix domain. /// (multi-level / `out<I>` slots). Lane points are in the slot's prefix domain.
/// @tparam StopLevel stop level
/// @tparam LgOpl lg opl
/// @tparam InputT input domain type
/// @tparam ForwardIterator forward iterator type
/// @param msb_mask the `msb_mask`
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @return the constructed object
/// @throws std::runtime_error if `list must be sorted`
template <std::size_t StopLevel, std::size_t LgOpl, typename InputT, template <std::size_t StopLevel, std::size_t LgOpl, typename InputT,
typename ForwardIterator> typename ForwardIterator>
auto make_sequence_recipe_at(InputT msb_mask, ForwardIterator begin, auto make_sequence_recipe_at(InputT msb_mask, ForwardIterator begin,

View file

@ -9,15 +9,15 @@
namespace dpf namespace dpf
{ {
/// Tag base for `eval_sequence` storage layout. /// @brief Tag base for `eval_sequence` storage layout.
struct return_type_tag_{}; struct return_type_tag_{};
/// Store whole leaves. Default for `eval_sequence`. The iterable still /// @brief Store whole leaves. Default for `eval_sequence`. The iterable still
/// yields one share per listed point. /// yields one share per listed point.
struct return_entire_node_tag_ final : public return_type_tag_ {}; struct return_entire_node_tag_ final : public return_type_tag_ {};
// static constexpr auto return_entire_node_tag = return_entire_node_tag_{}; // static constexpr auto return_entire_node_tag = return_entire_node_tag_{};
/// Store one share per listed point. /// @brief Store one share per listed point.
struct return_output_only_tag_ final : public return_type_tag_ {}; struct return_output_only_tag_ final : public return_type_tag_ {};
// static constexpr auto return_output_only_tag = return_output_only_tag_{}; // static constexpr auto return_output_only_tag = return_output_only_tag_{};

View file

@ -1,6 +1,5 @@
/// @file dpf/setbit_index_iterable.hpp /// @file dpf/setbit_index_iterable.hpp
/// @brief /// @brief Iterates the positions of set bits in a bit array.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;

View file

@ -1,7 +1,7 @@
/// @file dpf/subsequence_iterable.hpp /// @file dpf/subsequence_iterable.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief defines `dpf::subsequence_iterable` and associated helpers /// @brief defines `dpf::subsequence_iterable` and associated helpers
/// @details /// @details Yields a listed subset of another iterable without copying it.
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details. /// see [LICENSE.md](@ref license) for details.

385
include/dpf/tree_traits.hpp Normal file
View file

@ -0,0 +1,385 @@
/// @file dpf/tree_traits.hpp
/// @brief BGI vs Half-Tree walk policy, selected by the interior PRG.
/// @details Default traits match today's Boyle–Gilboa–Ishai tree. A PRG that
/// defines `half_tree_tag` (e.g. `prg::aes128_ccr`) opts into the
/// Guo et al. Half-Tree mid-level expand / CW / advance, with a
/// two-tweak last level that keeps BGI-style advice packing.
/// @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_TREE_TRAITS_HPP__
#define LIBDPF_INCLUDE_DPF_TREE_TRAITS_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <type_traits>
#include <utility>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "portable-snippets/exact-int/exact-int.h"
#include "dpf/twiddle.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// @brief Walk policy for interior DPF levels. Specialized when `PRG::half_tree_tag`
/// exists.
/// @tparam PRG pseudorandom generator
template <typename PRG, typename = void>
struct tree_traits
{
using prg = PRG;
using node = typename PRG::block_type;
static constexpr bool is_half_tree = false;
static constexpr bool stores_mid_advice = true;
static constexpr bool last_level_differs = false;
template <typename Sampler>
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1)
static void root_init(node out[2], Sampler && sample)
{
out[0] = dpf::unset_lo_bit(static_cast<node>(sample()));
out[1] = dpf::set_lo_bit(static_cast<node>(sample()));
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static bool is_last_level(std::size_t level, std::size_t depth) noexcept
{
(void)level;
(void)depth;
return false;
}
/// @brief BGI expand: clear lo-2bits, then `PRG::eval01`.
/// @param s the `s`
/// @return BGI expand: clear lo-2bits, then `PRG::eval01`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static auto expand(node s, bool /*is_last*/ = false) noexcept
{
return PRG::eval01(dpf::unset_lo_2bits(s));
}
/// @brief Convert / value-CW stretch. BGI: identical to `expand`.
/// @param s the node to stretch
/// @return the same value as `expand`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static auto expand_value(node s) noexcept
{
return expand(s, false);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2, 3)
static void expand_x4(const node * HEDLEY_RESTRICT seeds,
node * HEDLEY_RESTRICT left, node * HEDLEY_RESTRICT right,
bool /*is_last*/ = false) noexcept
{
alignas(node) node cleared[4];
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
cleared[i] = dpf::unset_lo_2bits(seeds[i]);
PRG::eval01_x4(cleared, left, right);
}
/// @brief Pack CW for eval: embed advice bit `dir` into the lo-bit.
/// @param cw the `cw`
/// @param advice the advice bit
/// @param dir the `dir`
/// @return the returned `node`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static node pack_cw(node cw, psnip_uint8_t advice, bool dir,
bool /*is_last*/ = false) noexcept
{
return dpf::set_lo_bit(cw, (advice >> dir) & 1);
}
/// @brief Off-path child XOR + packed advice `t0|t1` (BGI).
/// @param cw_out the `cw_out`
/// @param advice_out the `advice_out`
/// @param kids0 the `kids0`
/// @param kids1 the `kids1`
/// @param bit the bit value or bit index
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static void make_cw(node & cw_out, psnip_uint8_t & advice_out,
const std::array<node, 2> & kids0, const std::array<node, 2> & kids1,
const node & /*s0*/, const node & /*s1*/, bool bit,
bool /*is_last*/ = false) noexcept
{
const node child[2] = {
simde_mm_xor_si128(kids0[0], kids1[0]),
simde_mm_xor_si128(kids0[1], kids1[1])
};
const bool t0 = static_cast<bool>(dpf::get_lo_bit(child[0]) ^ !bit);
const bool t1 = static_cast<bool>(dpf::get_lo_bit(child[1]) ^ bit);
cw_out = child[!bit];
advice_out = static_cast<psnip_uint8_t>((t1 << 1) | t0);
}
/// @brief `xor_if(child[dir], pack_cw(...), parent_control)`.
/// @param parent the parent node
/// @param kids the `kids`
/// @param cw the `cw`
/// @param advice the advice bit
/// @param dir the `dir`
/// @param parent_control the `parent_control`
/// @return `xor_if(child[dir], pack_cw(...), parent_control)`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
static node advance(node parent, const std::array<node, 2> & kids,
node cw, psnip_uint8_t advice, bool dir,
bool parent_control, bool /*is_last*/ = false) noexcept
{
const node packed = pack_cw(cw, advice, dir);
return dpf::xor_if(kids[dir ? 1u : 0u], packed, parent_control);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static node traverse(node parent, node cw_packed, bool dir,
bool /*is_last*/ = false) noexcept
{
auto kids = expand(parent, false);
return dpf::xor_if_lo_bit(kids[dir ? 1u : 0u], cw_packed, parent);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static auto traverse01(node parent, node cw0, node cw1,
bool is_last = false) noexcept
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto kids = expand(parent, is_last);
return std::array<node, 2>{
dpf::xor_if_lo_bit(kids[0], cw0, parent),
dpf::xor_if_lo_bit(kids[1], cw1, parent)
};
HEDLEY_PRAGMA(GCC diagnostic pop)
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 4, 5)
static void traverse01_x4(const node * HEDLEY_RESTRICT parents,
node cw0, node cw1, node * HEDLEY_RESTRICT left,
node * HEDLEY_RESTRICT right, bool is_last = false) noexcept
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
expand_x4(parents, left, right, is_last);
HEDLEY_PRAGMA(GCC diagnostic pop)
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
left[i] = dpf::xor_if_lo_bit(left[i], cw0, parents[i]);
right[i] = dpf::xor_if_lo_bit(right[i], cw1, parents[i]);
}
}
};
/// @brief Half-Tree specialization (PRG advertises `half_tree_tag`).
/// @tparam PRG pseudorandom generator
template <typename PRG>
struct tree_traits<PRG, std::void_t<typename PRG::half_tree_tag>>
{
using prg = PRG;
using node = typename PRG::block_type;
static constexpr bool is_half_tree = true;
static constexpr bool stores_mid_advice = false;
static constexpr bool last_level_differs = true;
template <typename Sampler>
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1)
static void root_init(node out[2], Sampler && sample)
{
// Shares of a fixed Δ with lsb(Δ)=1.
const node s = dpf::unset_lo_bit(static_cast<node>(sample()));
const node delta = dpf::set_lo_bit(static_cast<node>(sample()));
out[0] = s;
out[1] = simde_mm_xor_si128(s, delta);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static bool is_last_level(std::size_t level, std::size_t depth) noexcept
{
return depth != 0 && level + 1 == depth;
}
/// @brief Mid: `{H(s), H(s)⊕s}` (keep control bit). Last: two-tweak stretch.
/// @param s the `s`
/// @param is_last the `is_last`
/// @return Mid: `{H(s), H(s)⊕s}` (keep control bit)
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static auto expand(node s, bool is_last = false) noexcept
{
if (is_last)
{
// `{H(s|0), H(s|1)}` — LSB forced, matching Guo et al. / myl7.
const node base = dpf::unset_lo_bit(s);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::array<node, 2>{
PRG::hash(base),
PRG::hash(dpf::set_lo_bit(base))
};
HEDLEY_PRAGMA(GCC diagnostic pop)
}
return PRG::eval01(s);
}
/// @brief Convert / value-CW stretch: always two-tweak `{H(s|0), H(s|1)}`.
/// @details Seed walk mid levels keep using half-style `expand`.
/// @param s the node to stretch
/// @return the pair `{H(s|0), H(s|1)}`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static auto expand_value(node s) noexcept
{
return expand(s, true);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 2, 3)
static void expand_x4(const node * HEDLEY_RESTRICT seeds,
node * HEDLEY_RESTRICT left, node * HEDLEY_RESTRICT right,
bool is_last = false) noexcept
{
if (is_last)
{
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
auto kids = expand(seeds[i], true);
left[i] = kids[0];
right[i] = kids[1];
}
return;
}
PRG::eval01_x4(seeds, left, right);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
static node pack_cw(node cw, psnip_uint8_t advice, bool dir,
bool is_last = false) noexcept
{
if (!is_last)
return cw; // mid: full CW, no advice packing
return dpf::set_lo_bit(cw, (advice >> dir) & 1);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static void make_cw(node & cw_out, psnip_uint8_t & advice_out,
const std::array<node, 2> & kids0, const std::array<node, 2> & kids1,
const node & s0, const node & s1, bool bit, bool is_last = false) noexcept
{
if (is_last)
{
// BGI-style advice on the leaf step only.
const node child[2] = {
simde_mm_xor_si128(kids0[0], kids1[0]),
simde_mm_xor_si128(kids0[1], kids1[1])
};
const bool t0 = static_cast<bool>(dpf::get_lo_bit(child[0]) ^ !bit);
const bool t1 = static_cast<bool>(dpf::get_lo_bit(child[1]) ^ bit);
cw_out = child[!bit];
advice_out = static_cast<psnip_uint8_t>((t1 << 1) | t0);
(void)s0;
(void)s1;
return;
}
// CW = H(s0)⊕H(s1)⊕ᾱ·Δ = off-path children XOR (Half-Tree identity).
const node child[2] = {
simde_mm_xor_si128(kids0[0], kids1[0]),
simde_mm_xor_si128(kids0[1], kids1[1])
};
cw_out = child[!bit];
advice_out = 0;
(void)s0;
(void)s1;
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
static node advance(node parent, const std::array<node, 2> & kids,
node cw, psnip_uint8_t advice, bool dir, bool parent_control,
bool is_last = false) noexcept
{
// Mid and last: select child[dir], XOR packed CW if parent control set.
// Mid Half-Tree: child[1]=H⊕s so this is `h ⊕ (dir?s:0) ⊕ (t?cw:0)`.
const node packed = pack_cw(cw, advice, dir, is_last);
return dpf::xor_if(kids[dir ? 1u : 0u], packed, parent_control);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static node traverse(node parent, node cw_packed, bool dir,
bool is_last = false) noexcept
{
auto kids = expand(parent, is_last);
return dpf::xor_if_lo_bit(kids[dir ? 1u : 0u], cw_packed, parent);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
static auto traverse01(node parent, node cw0, node cw1,
bool is_last = false) noexcept
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto kids = expand(parent, is_last);
return std::array<node, 2>{
dpf::xor_if_lo_bit(kids[0], cw0, parent),
dpf::xor_if_lo_bit(kids[1], cw1, parent)
};
HEDLEY_PRAGMA(GCC diagnostic pop)
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_NON_NULL(1, 4, 5)
static void traverse01_x4(const node * HEDLEY_RESTRICT parents,
node cw0, node cw1, node * HEDLEY_RESTRICT left,
node * HEDLEY_RESTRICT right, bool is_last = false) noexcept
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
expand_x4(parents, left, right, is_last);
HEDLEY_PRAGMA(GCC diagnostic pop)
DPF_UNROLL_LOOP
for (std::size_t i = 0; i < 4; ++i)
{
left[i] = dpf::xor_if_lo_bit(left[i], cw0, parents[i]);
right[i] = dpf::xor_if_lo_bit(right[i], cw1, parents[i]);
}
}
};
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_TREE_TRAITS_HPP__

View file

@ -1,6 +1,5 @@
/// @file dpf/twiddle.hpp /// @file dpf/twiddle.hpp
/// @brief /// @brief Low-bit extract, sibling nodes, and small bit masks.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -79,9 +78,9 @@ auto get_if_lo_bit(std::array<simde__m128i, N> a, simde__m128i b) noexcept
std::transform(std::begin(a), std::end(a), std::begin(a), [mask](simde__m128i & a){return simde_mm_and_si128(a, mask);}); std::transform(std::begin(a), std::end(a), std::begin(a), [mask](simde__m128i & a){return simde_mm_and_si128(a, mask);});
return a; return a;
} }
HEDLEY_PRAGMA(GCC diagnostic pop) HEDLEY_PRAGMA(GCC diagnostic pop)
// if low bit of c is set, then return xor of a and b, else return a // if low bit of c is set, then return xor of a and b, else return a
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE

View file

@ -4,6 +4,7 @@
/// packs one lane every two bits, low lane in the low bits of the /// packs one lane every two bits, low lane in the low bits of the
/// first byte, matching `dpf::bit`. Leaf addition is not XOR: a /// first byte, matching `dpf::bit`. Leaf addition is not XOR: a
/// carry stays inside the 2-bit lane. See `packed_lane_arithmetic.hpp`. /// carry stays inside the 2-bit lane. See `packed_lane_arithmetic.hpp`.
/// @see packed_lane_arithmetic.hpp
#ifndef LIBDPF_INCLUDE_DPF_TWOBIT_HPP__ #ifndef LIBDPF_INCLUDE_DPF_TWOBIT_HPP__
#define LIBDPF_INCLUDE_DPF_TWOBIT_HPP__ #define LIBDPF_INCLUDE_DPF_TWOBIT_HPP__
@ -50,7 +51,10 @@ static constexpr dpf::twobit to_twobit(unsigned long long value) noexcept
} }
/// @brief parse one character as a 2-bit digit /// @brief parse one character as a 2-bit digit
/// @tparam CharT character type
/// @param zero character for 0 (default `'0'`) /// @param zero character for 0 (default `'0'`)
/// @param value the value to convert or store
/// @return the returned `dpf::twobit`
/// @throws std::domain_error if `value` is not one of the four digits /// @throws std::domain_error if `value` is not one of the four digits
template <typename CharT> template <typename CharT>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -92,6 +96,9 @@ operator>>(std::basic_istream<CharT, Traits> & is, dpf::twobit & value)
} }
/// @brief addition in Z/4Z /// @brief addition in Z/4Z
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return addition in Z/4Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -102,6 +109,9 @@ constexpr dpf::twobit operator+(dpf::twobit lhs, dpf::twobit rhs) noexcept
} }
/// @brief subtraction in Z/4Z /// @brief subtraction in Z/4Z
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return subtraction in Z/4Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -112,6 +122,8 @@ constexpr dpf::twobit operator-(dpf::twobit lhs, dpf::twobit rhs) noexcept
} }
/// @brief additive inverse in Z/4Z /// @brief additive inverse in Z/4Z
/// @param value the value to convert or store
/// @return additive inverse in Z/4Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -121,6 +133,9 @@ constexpr dpf::twobit operator-(dpf::twobit value) noexcept
} }
/// @brief multiplication in Z/4Z /// @brief multiplication in Z/4Z
/// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return multiplication in Z/4Z
HEDLEY_CONST HEDLEY_CONST
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE

View file

@ -1,6 +1,6 @@
/// @file dpf/utils.hpp /// @file dpf/utils.hpp
/// @brief miscellaneous helper functions, structs, preprocessor directives /// @brief miscellaneous helper functions, structs, preprocessor directives
/// @details /// @details Type traits, bit lengths, and small tuple helpers shared by the headers.
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -95,13 +95,13 @@ class numeric_limits<uint128_t const>
: public numeric_limits<uint128_t> {}; : public numeric_limits<uint128_t> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `uint128_t volatile` /// @brief `uint128_t volatile`
template<> template<>
class numeric_limits<uint128_t volatile> class numeric_limits<uint128_t volatile>
: public numeric_limits<uint128_t> {}; : public numeric_limits<uint128_t> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `uint128_t const volatile` /// @brief `uint128_t const volatile`
template<> template<>
class numeric_limits<uint128_t const volatile> class numeric_limits<uint128_t const volatile>
: public numeric_limits<uint128_t> {}; : public numeric_limits<uint128_t> {};
@ -162,13 +162,13 @@ class numeric_limits<uint256_t const>
: public numeric_limits<uint256_t> {}; : public numeric_limits<uint256_t> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `uint256_t volatile` /// @brief `uint256_t volatile`
template<> template<>
class numeric_limits<uint256_t volatile> class numeric_limits<uint256_t volatile>
: public numeric_limits<uint256_t> {}; : public numeric_limits<uint256_t> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `uint256_t const volatile` /// @brief `uint256_t const volatile`
template<> template<>
class numeric_limits<uint256_t const volatile> class numeric_limits<uint256_t const volatile>
: public numeric_limits<uint256_t> {}; : public numeric_limits<uint256_t> {};
@ -184,6 +184,11 @@ namespace utils
{ {
/// @brief Ugly hack to implement `constexpr`-frien`dly conditional `throw` /// @brief Ugly hack to implement `constexpr`-frien`dly conditional `throw`
/// @tparam Exception exception
/// @param b the `b`
/// @param what the diagnostic message
/// @return Ugly hack to implement `constexpr`-frien`dly conditional `throw`
/// @throws Exception
template <typename Exception> template <typename Exception>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
static constexpr auto constexpr_maybe_throw(bool b, std::string_view what) -> void static constexpr auto constexpr_maybe_throw(bool b, std::string_view what) -> void
@ -213,6 +218,11 @@ template <typename T>
static constexpr bool is_quotient_integer_v = is_quotient_integer<T>::value; static constexpr bool is_quotient_integer_v = is_quotient_integer<T>::value;
/// @brief Integer overflow-proof ceiling of division /// @brief Integer overflow-proof ceiling of division
/// @tparam T value type
/// @tparam T value type
/// @param numerator the `numerator`
/// @param denominator the `denominator`
/// @return Integer overflow-proof ceiling of division
template <typename T, template <typename T,
std::enable_if_t<is_quotient_integer_v<T>, bool> = false> std::enable_if_t<is_quotient_integer_v<T>, bool> = false>
HEDLEY_CONST HEDLEY_CONST
@ -225,6 +235,11 @@ static constexpr T quotient_ceiling(T numerator, T denominator) noexcept
} }
/// @brief Integer overflow-proof floor of division /// @brief Integer overflow-proof floor of division
/// @tparam T value type
/// @tparam T value type
/// @param numerator the `numerator`
/// @param denominator the `denominator`
/// @return Integer overflow-proof floor of division
template <typename T, template <typename T,
std::enable_if_t<is_quotient_integer_v<T>, bool> = false> std::enable_if_t<is_quotient_integer_v<T>, bool> = false>
HEDLEY_CONST HEDLEY_CONST
@ -242,11 +257,12 @@ struct is_signed_integral
template <typename T> template <typename T>
static constexpr bool is_signed_integral_v = is_signed_integral<T>::value; static constexpr bool is_signed_integral_v = is_signed_integral<T>::value;
/// Whether DPF keygen/eval flip the input MSB (two's-complement domains). /// @brief Whether DPF keygen/eval flip the input MSB (two's-complement domains).
/// Distinct from `is_signed_integral`: wrappers such as signed `fixedpoint` /// @details Distinct from `is_signed_integral`: wrappers such as signed `fixedpoint`
/// are not `std::is_integral`, and treating them as such would break /// are not `std::is_integral`, and treating them as such would break
/// `make_unsigned`. Sequence recipe construction and breadth-first eval /// `make_unsigned`. Sequence recipe construction and breadth-first eval
/// must use this trait, not `is_signed_integral_v`. /// must use this trait, not `is_signed_integral_v`.
/// @tparam T value type
template <typename T> template <typename T>
struct uses_signed_msb struct uses_signed_msb
: std::bool_constant< : std::bool_constant<
@ -274,6 +290,9 @@ template <typename T>
using make_unsigned_t = typename make_unsigned<T>::type; using make_unsigned_t = typename make_unsigned<T>::type;
/// @brief Make an `std::bitset` from a variadic list of `bool`s /// @brief Make an `std::bitset` from a variadic list of `bool`s
/// @tparam Bools bools
/// @param bs the `bs`
/// @return Make an `std::bitset` from a variadic list of `bool`s
template <typename ...Bools> template <typename ...Bools>
auto make_bitset(Bools ...bs) auto make_bitset(Bools ...bs)
{ {
@ -356,6 +375,7 @@ struct bitlength_of<simde__m128i>
template <> template <>
struct bitlength_of<simde__m256i> struct bitlength_of<simde__m256i>
: public std::integral_constant<std::size_t, 256> { }; : public std::integral_constant<std::size_t, 256> { };
HEDLEY_PRAGMA(GCC diagnostic pop)
// template <> // template <>
// struct bitlength_of<simde__m512i> // struct bitlength_of<simde__m512i>
@ -364,7 +384,6 @@ struct bitlength_of<simde__m256i>
template <typename T, std::size_t N> template <typename T, std::size_t N>
struct bitlength_of<std::array<T, N>> struct bitlength_of<std::array<T, N>>
: public std::integral_constant<std::size_t, bitlength_of_v<T> * N> { }; : public std::integral_constant<std::size_t, bitlength_of_v<T> * N> { };
HEDLEY_PRAGMA(GCC diagnostic pop)
template <typename OutputT, template <typename OutputT,
typename NodeT> typename NodeT>
@ -385,6 +404,9 @@ template <typename OutputT,
static constexpr std::size_t bitlength_of_output_v = bitlength_of_output<OutputT, NodeT>::value; static constexpr std::size_t bitlength_of_output_v = bitlength_of_output<OutputT, NodeT>::value;
/// @brief the primitive integral type used to represent non integral types /// @brief the primitive integral type used to represent non integral types
/// @tparam Nbits width in bits
/// @tparam MinBits min bits
/// @tparam MaxBits max bits
template <std::size_t Nbits, template <std::size_t Nbits,
std::size_t MinBits = Nbits, std::size_t MinBits = Nbits,
std::size_t MaxBits = std::max(Nbits, MinBits)> std::size_t MaxBits = std::max(Nbits, MinBits)>
@ -413,6 +435,9 @@ template <std::size_t Nbits,
using integral_type_from_bitlength_t = typename integral_type_from_bitlength<Nbits, MinBits, MaxBits>::type; using integral_type_from_bitlength_t = typename integral_type_from_bitlength<Nbits, MinBits, MaxBits>::type;
/// @brief the primitive integral type used to represent non integral types /// @brief the primitive integral type used to represent non integral types
/// @tparam Nbits width in bits
/// @tparam MinBits min bits
/// @tparam MaxBits max bits
template <std::size_t Nbits, template <std::size_t Nbits,
std::size_t MinBits = Nbits, std::size_t MinBits = Nbits,
std::size_t MaxBits = std::max(std::size_t(256), MinBits)> std::size_t MaxBits = std::max(std::size_t(256), MinBits)>
@ -475,9 +500,13 @@ struct make_from_integral_value
} }
}; };
/// Reconstruct `x0 XOR x1` via the integral bridge. Prefer this over /// @brief Reconstruct `x0 XOR x1` via the integral bridge. Prefer this over
/// `static_cast<T>(x0 ^ x1)`: for `keyword`, `operator^` yields the parent /// `static_cast<T>(x0 ^ x1)`: for `keyword`, `operator^` yields the parent
/// `modint`, which cannot convert back through the private keyword ctor. /// `modint`, which cannot convert back through the private keyword ctor.
/// @tparam T value type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @return Reconstruct `x0 XOR x1` via the integral bridge
template <typename T> template <typename T>
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr T xor_input_shares(T x0, T x1) noexcept constexpr T xor_input_shares(T x0, T x1) noexcept
@ -507,8 +536,12 @@ static constexpr IntegralT get_node_mask(InputT mask, std::size_t level_index)
return static_cast<IntegralT>(to_int(mask) >> (level_index-1 + dpf_type::lg_outputs_per_leaf)); return static_cast<IntegralT>(to_int(mask) >> (level_index-1 + dpf_type::lg_outputs_per_leaf));
} }
/// Logical right shift. Offsets at or past the width yield 0 (a `>>` of that /// @brief Logical right shift. Offsets at or past the width yield 0 (a `>>` of that
/// width is undefined for the native unsigned types). /// width is undefined for the native unsigned types).
/// @tparam IntegralT integral type
/// @param value the value to convert or store
/// @param offset the public offset
/// @return Logical right shift
template <typename IntegralT> template <typename IntegralT>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -519,7 +552,11 @@ constexpr IntegralT shift_right(IntegralT value, std::size_t offset) noexcept
return static_cast<IntegralT>(value >> offset); return static_cast<IntegralT>(value >> offset);
} }
/// Floor of `from_inclusive / 2^lg_opl`. `lg_opl` is `log2(outputs_per_leaf)`. /// @brief Floor of `from_inclusive / 2^lg_opl`. `lg_opl` is `log2(outputs_per_leaf)`.
/// @tparam IntegralT integral type
/// @param from_inclusive the `from_inclusive`
/// @param lg_opl the `lg_opl`
/// @return Floor of `from_inclusive / 2^lg_opl`
template <typename IntegralT> template <typename IntegralT>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -530,11 +567,15 @@ constexpr IntegralT leaf_node_floor(IntegralT from_inclusive, std::size_t lg_opl
return shift_right(from_inclusive, lg_opl); return shift_right(from_inclusive, lg_opl);
} }
/// Exclusive leaf index of an inclusive input `to_inclusive`. /// @brief Exclusive leaf index of an inclusive input `to_inclusive`.
/// `2^lg_opl` outputs share a leaf. When `to_inclusive + 1` does not fit in /// @details `2^lg_opl` outputs share a leaf. When `to_inclusive + 1` does not fit in
/// `IntegralT`, the exclusive node index is `2^(width - lg_opl)`. That value /// `IntegralT`, the exclusive node index is `2^(width - lg_opl)`. That value
/// itself does not fit when `lg_opl == 0`; the returned 0 is that saturated /// itself does not fit when `lg_opl == 0`; the returned 0 is that saturated
/// end (`[from, 2^width)`), which `split_leaf_nodes` interprets. /// end (`[from, 2^width)`), which `split_leaf_nodes` interprets.
/// @tparam IntegralT integral type
/// @param to_inclusive the `to_inclusive`
/// @param lg_opl the `lg_opl`
/// @return Exclusive leaf index of an inclusive input `to_inclusive`
template <typename IntegralT> template <typename IntegralT>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -553,9 +594,15 @@ constexpr IntegralT leaf_node_ceil_exclusive(IntegralT to_inclusive, std::size_t
return quotient_ceiling(next, opl); return quotient_ceiling(next, opl);
} }
/// Multi-level flavor: caller passes the slot's `lg(outputs-per-leaf)` (and, /// @brief Multi-level flavor: caller passes the slot's `lg(outputs-per-leaf)` (and,
/// for interval splitting, its `tree_level`) explicitly. The classic wrappers /// for interval splitting, its `tree_level`) explicitly. The classic wrappers
/// below forward the deepest-slot packing (`DpfKey::lg_outputs_per_leaf`). /// below forward the deepest-slot packing (`DpfKey::lg_outputs_per_leaf`).
/// @tparam InputT input domain type
/// @tparam size_t size type
/// @param from the inclusive start of the range
/// @param lg_opl the `lg_opl`
/// @return Multi-level flavor: caller passes the slot's `lg(outputs-per-leaf)` (and, for interval
/// splitting, its `tree_level`) explicitly
template <typename InputT, template <typename InputT,
typename IntegralT = integral_type_from_bitlength_t< typename IntegralT = integral_type_from_bitlength_t<
bitlength_of_v<InputT>, bitlength_of_v<std::size_t>>> bitlength_of_v<InputT>, bitlength_of_v<std::size_t>>>
@ -591,8 +638,9 @@ static constexpr IntegralT get_to_node(InputT to)
return get_to_node_at<InputT, IntegralT>(to, DpfKey::lg_outputs_per_leaf); return get_to_node_at<InputT, IntegralT>(to, DpfKey::lg_outputs_per_leaf);
} }
/// One half-open leaf-node range. `to_node == 0` with a nonzero `count` is the /// @brief One half-open leaf-node range. `to_node == 0` with a nonzero `count` is the
/// saturated end `[from_node, 2^width)`. /// saturated end `[from_node, 2^width)`.
/// @tparam IntegralT integral type
template <typename IntegralT> template <typename IntegralT>
struct node_segment struct node_segment
{ {
@ -609,9 +657,14 @@ struct node_segments
std::size_t total = 0; std::size_t total = 0;
}; };
/// True when the inclusive walk `[from, to]` wraps the low `bits` of the /// @brief True when the inclusive walk `[from, to]` wraps the low `bits` of the
/// domain. Comparison is on the post-MSB-flip bit pattern. Leaf ids alone /// domain. Comparison is on the post-MSB-flip bit pattern. Leaf ids alone
/// cannot carry this: packing can put a wrapping pair into `from_node <= to_node`. /// cannot carry this: packing can put a wrapping pair into `from_node <= to_node`.
/// @tparam IntegralT integral type
/// @param from the inclusive start of the range
/// @param to the `to`
/// @param bits the packed bits
/// @return True when the inclusive walk `[from, to]` wraps the low `bits` of the domain
template <typename IntegralT> template <typename IntegralT>
inline bool interval_wraps(IntegralT from, IntegralT to, std::size_t bits) inline bool interval_wraps(IntegralT from, IntegralT to, std::size_t bits)
{ {
@ -626,18 +679,25 @@ inline bool interval_wraps(IntegralT from, IntegralT to, std::size_t bits)
return from > to; return from > to;
} }
/// Split an inclusive output interval, already reduced to leaf ids, into one /// @brief Split an inclusive output interval, already reduced to leaf ids, into one
/// or two half-open walks. A linearized `from_node > to_node` wraps the node /// or two half-open walks. A linearized `from_node > to_node` wraps the node
/// id space `[0, 2^depth)`. A saturated `to_node == 0` means the exclusive end /// id space `[0, 2^depth)`. A saturated `to_node == 0` means the exclusive end
/// is `2^{bitwidth(IntegralT)}`, which is the whole id space when `depth` is /// is `2^{bitwidth(IntegralT)}`, which is the whole id space when `depth` is
/// that width. /// that width.
/// ///
/// `input_wraps` is the order of the original inputs, before leaf coarsening. /// `input_wraps` is the order of the original inputs, before leaf coarsening.
/// The buffer is still two runs, `[from_node, 2^depth)` then `[0, to_node)`, /// @details The buffer is still two runs, `[from_node, 2^depth)` then `[0, to_node)`,
/// even when packing makes `from_node <= to_node`. In that case the runs /// even when packing makes `from_node <= to_node`. In that case the runs
/// overlap on the shared leaf: the iterable's preclip consumes the start of /// overlap on the shared leaf: the iterable's preclip consumes the start of
/// the first copy and its length stops inside the second. Collapsing the /// the first copy and its length stops inside the second. Collapsing the
/// overlap into one forward segment writes the wrong leaves. /// overlap into one forward segment writes the wrong leaves.
/// @tparam IntegralT integral type
/// @param from_node the `from_node`
/// @param to_node the `to_node`
/// @param depth the tree depth
/// @param input_wraps the `input_wraps`
/// @return the returned `node_segments<IntegralT>`
/// @throws std::length_error if `DPF leaf domain does not fit in size_t`
template <typename IntegralT> template <typename IntegralT>
inline node_segments<IntegralT> split_leaf_nodes(IntegralT from_node, inline node_segments<IntegralT> split_leaf_nodes(IntegralT from_node,
IntegralT to_node, std::size_t depth, bool input_wraps = false) IntegralT to_node, std::size_t depth, bool input_wraps = false)
@ -745,7 +805,13 @@ static std::size_t get_leafnodes_in_output_interval(InputT from, InputT to)
static_cast<std::size_t>(DpfKey::depth), wraps).total; static_cast<std::size_t>(DpfKey::depth), wraps).total;
} }
/// Historical name used by the test suite. /// @brief Historical name used by the test suite.
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @tparam IntegralT integral type
/// @param from the inclusive start of the range
/// @param to the `to`
/// @return Historical name used by the test suite
template <typename DpfKey, template <typename DpfKey,
typename InputT = typename DpfKey::input_type, typename InputT = typename DpfKey::input_type,
typename IntegralT = typename DpfKey::integral_type> typename IntegralT = typename DpfKey::integral_type>
@ -1092,6 +1158,7 @@ struct countr_zero<simde__m256i>
return suffix_len; return suffix_len;
} }
}; };
HEDLEY_PRAGMA(GCC diagnostic pop)
// template <> // template <>
// struct countl_zero<simde__m512i> // struct countl_zero<simde__m512i>
@ -1113,7 +1180,6 @@ struct countr_zero<simde__m256i>
// return prefix_len; // return prefix_len;
// } // }
// }; // };
HEDLEY_PRAGMA(GCC diagnostic pop)
template <typename T> template <typename T>
struct is_xor_wrapper : std::false_type {}; struct is_xor_wrapper : std::false_type {};
@ -1121,9 +1187,10 @@ struct is_xor_wrapper : std::false_type {};
template <typename T> template <typename T>
static constexpr bool is_xor_wrapper_v = is_xor_wrapper<T>::value; static constexpr bool is_xor_wrapper_v = is_xor_wrapper<T>::value;
/// Sub-byte DPF outputs whose lanes are packed inside a leaf node /// @brief Sub-byte DPF outputs whose lanes are packed inside a leaf node
/// (`dpf::bit` is 1, `dpf::twobit` is 2, `dpf::nyble` is 4). The leaf /// (`dpf::bit` is 1, `dpf::twobit` is 2, `dpf::nyble` is 4). The leaf
/// image is the buffer image: interval eval memcpy's the node. /// image is the buffer image: interval eval memcpy's the node.
/// @tparam T value type
template <typename T> template <typename T>
struct is_packed_subbyte : std::false_type {}; struct is_packed_subbyte : std::false_type {};
@ -1147,7 +1214,10 @@ constexpr auto data(T & bar) noexcept // NOLINT(runtime/references)
return std::data(bar); return std::data(bar);
} }
/// Pointer overload. Constness of `bar` is the constness of `T`. /// @brief Pointer overload. Constness of `bar` is the constness of `T`.
/// @tparam T value type
/// @param bar the `bar`
/// @return Pointer overload
template <typename T> template <typename T>
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
@ -1269,6 +1339,39 @@ auto get_common_part_hash(const std::array<InteriorNodeT, Depth> & correction_wo
return digest; return digest;
} }
template <typename InteriorNodeT,
std::size_t Depth,
typename LeafTupleT,
typename WildcardMaskT,
typename ExtraT>
auto get_common_part_hash(const std::array<InteriorNodeT, Depth> & correction_words,
const std::array<psnip_uint8_t, Depth> & correction_advice,
const LeafTupleT & leaf_tuple,
const WildcardMaskT & wildcard_mask,
const ExtraT & extra)
{
using zero_type = unsigned char;
static constexpr zero_type zero{};
SHA256 h;
digest_type digest;
h.add(&correction_words, sizeof(correction_words));
h.add(&correction_advice, sizeof(correction_advice));
if constexpr (std::tuple_size_v<ExtraT> > 0)
h.add(&extra, sizeof(extra));
std::apply([&h, &wildcard_mask](auto const & ...leaf)
{
std::apply([&h, &leaf...](auto ...is_wildcard)
{
(h.add(!is_wildcard ? reinterpret_cast<const zero_type*>(&leaf.get()) : &zero, !is_wildcard ? sizeof(leaf.get()) : sizeof(zero)), ...);
}, wildcard_mask);
}, leaf_tuple);
h.getHash(digest.data());
return digest;
}
template <typename DpfKey> template <typename DpfKey>
auto get_common_part_hash(const DpfKey & dpf) auto get_common_part_hash(const DpfKey & dpf)
{ {
@ -1284,6 +1387,7 @@ struct has_operators_plus_minus : public std::false_type { };
/// @brief True when `a + b` and `a - b` are valid expressions. /// @brief True when `a + b` and `a - b` are valid expressions.
/// Overload sets are accepted; taking the address of `operator+` /// Overload sets are accepted; taking the address of `operator+`
/// is not, because that fails when `+` or `-` is overloaded. /// is not, because that fails when `+` or `-` is overloaded.
/// @tparam OutputT output type
template <typename OutputT> template <typename OutputT>
struct has_operators_plus_minus<OutputT, struct has_operators_plus_minus<OutputT,
std::void_t< std::void_t<

277
include/dpf/vec.hpp Normal file
View file

@ -0,0 +1,277 @@
/// @file dpf/vec.hpp
/// @brief `dpf::vec<T, N>`, one output of `N` lanes added componentwise.
/// @details Each lane lives in the ring of `T` (`+` and `-` wrap the way `T`
/// wraps). A leaf adds every lane and does not carry from one lane
/// into the next. `T` is an ordinary output type: an integer,
/// `modint`, `fixedpoint`, `twobit`, `nyble`, or `xor_wrapper`.
/// @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_DPF_VEC_HPP__
#define LIBDPF_INCLUDE_DPF_VEC_HPP__
#include <array>
#include <cstddef>
#include <cstring>
#include <type_traits>
#include <utility>
#include "hedley/hedley.h"
#include "dpf/leaf_arithmetic.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// @brief `N` lanes of `T`, combined componentwise with no carry between lanes.
/// @tparam T lane type. Its `+`, `-`, and `*` are used per lane
/// @tparam N number of lanes
template <typename T, std::size_t N>
struct vec
{
static_assert(N > 0, "dpf::vec needs at least one lane");
static constexpr bool dpf_vec = true;
/// @brief Number of lanes.
static constexpr std::size_t lane_count = N;
/// @brief Type of one lane.
using lane_type = T;
/// @brief Lane storage, index 0 in the least-significant lane.
std::array<T, N> lanes{};
/// @brief Value-initialize every lane.
HEDLEY_NO_THROW
constexpr vec() noexcept = default;
/// @brief Mutable lane `i`.
/// @param i the lane index
/// @return a reference to that lane
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr T & operator[](std::size_t i) noexcept { return lanes[i]; }
/// @brief Lane `i`.
/// @param i the lane index
/// @return a reference to that lane
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr const T & operator[](std::size_t i) const noexcept { return lanes[i]; }
/// @brief Negate every lane.
/// @return the negated vector
HEDLEY_ALWAYS_INLINE
constexpr vec operator-() const
noexcept(noexcept(static_cast<T>(-std::declval<const T &>())))
{
vec out;
for (std::size_t i = 0; i < N; ++i)
out.lanes[i] = static_cast<T>(-lanes[i]);
return out;
}
/// @brief Add each lane, with no carry into the next lane.
/// @param rhs the right-hand vector
/// @return the lane-wise sum
HEDLEY_ALWAYS_INLINE
constexpr vec operator+(const vec & rhs) const
noexcept(noexcept(std::declval<const T &>() + std::declval<const T &>()))
{
vec out;
for (std::size_t i = 0; i < N; ++i)
out.lanes[i] = static_cast<T>(lanes[i] + rhs.lanes[i]);
return out;
}
/// @brief Subtract each lane, with no borrow from the next lane.
/// @param rhs the right-hand vector
/// @return the lane-wise difference
HEDLEY_ALWAYS_INLINE
constexpr vec operator-(const vec & rhs) const
noexcept(noexcept(std::declval<const T &>() - std::declval<const T &>()))
{
vec out;
for (std::size_t i = 0; i < N; ++i)
out.lanes[i] = static_cast<T>(lanes[i] - rhs.lanes[i]);
return out;
}
/// @brief Multiply each lane.
/// @param rhs the right-hand vector
/// @return the lane-wise product
HEDLEY_ALWAYS_INLINE
constexpr vec operator*(const vec & rhs) const
noexcept(noexcept(std::declval<const T &>() * std::declval<const T &>()))
{
vec out;
for (std::size_t i = 0; i < N; ++i)
out.lanes[i] = static_cast<T>(lanes[i] * rhs.lanes[i]);
return out;
}
/// @brief Lane-wise equality.
/// @param rhs the right-hand vector
/// @return `true` when every lane matches
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr bool operator==(const vec & rhs) const noexcept
{
for (std::size_t i = 0; i < N; ++i)
{
if (!(lanes[i] == rhs.lanes[i]))
return false;
}
return true;
}
/// @brief Lane-wise inequality.
/// @param rhs the right-hand vector
/// @return `true` when some lane differs
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
constexpr bool operator!=(const vec & rhs) const noexcept
{
return !(*this == rhs);
}
};
namespace utils
{
template <typename T, std::size_t N>
struct bitlength_of<dpf::vec<T, N>>
: std::integral_constant<std::size_t, bitlength_of_v<T> * N> {};
} // namespace utils
namespace leaf_arithmetic
{
namespace vec_detail
{
template <typename U>
struct is_std_array : std::false_type {};
template <typename U, std::size_t M>
struct is_std_array<std::array<U, M>> : std::true_type {};
/// @brief Add or subtract every stored lane. A node wider than the vector is a
/// sequence of vectors, then a tail of leftover lanes. Bytes that do not
/// fill a lane are XOR-combined so a random pad still cancels.
/// @tparam T lane type
/// @tparam N number of lanes
/// @tparam Op lane operation
/// @param dst the destination
/// @param a the `a`
/// @param b the `b`
/// @param nbytes the number of bytes
/// @param op the `op`
template <typename T, std::size_t N, typename Op>
HEDLEY_ALWAYS_INLINE
void apply_bytes(unsigned char * dst, const unsigned char * a,
const unsigned char * b, std::size_t nbytes, Op op)
{
using V = dpf::vec<T, N>;
std::size_t off = 0;
const std::size_t nvec = nbytes / sizeof(V);
for (std::size_t i = 0; i < nvec; ++i)
{
V va, vb;
std::memcpy(&va, a + off, sizeof(V));
std::memcpy(&vb, b + off, sizeof(V));
V vc = op(va, vb);
std::memcpy(dst + off, &vc, sizeof(V));
off += sizeof(V);
}
while (off + sizeof(T) <= nbytes)
{
T va, vb;
std::memcpy(&va, a + off, sizeof(T));
std::memcpy(&vb, b + off, sizeof(T));
T vc = op(va, vb);
std::memcpy(dst + off, &vc, sizeof(T));
off += sizeof(T);
}
for (; off < nbytes; ++off)
dst[off] = static_cast<unsigned char>(a[off] ^ b[off]);
}
template <typename T, std::size_t N, typename NodeT, typename Op>
HEDLEY_ALWAYS_INLINE
NodeT apply_node(const NodeT & a, const NodeT & b, Op op)
{
alignas(NodeT) unsigned char ca[sizeof(NodeT)];
alignas(NodeT) unsigned char cb[sizeof(NodeT)];
alignas(NodeT) unsigned char cc[sizeof(NodeT)];
std::memcpy(ca, &a, sizeof(NodeT));
std::memcpy(cb, &b, sizeof(NodeT));
apply_bytes<T, N>(cc, ca, cb, sizeof(NodeT), op);
NodeT out;
std::memcpy(&out, cc, sizeof(NodeT));
return out;
}
} // namespace vec_detail
template <typename T, std::size_t N, typename NodeT>
struct add_t<dpf::vec<T, N>, NodeT,
std::enable_if_t<!std::is_void_v<NodeT>
&& !vec_detail::is_std_array<NodeT>::value>>
{
HEDLEY_ALWAYS_INLINE
auto operator()(const NodeT & a, const NodeT & b) const
{
return vec_detail::apply_node<T, N>(a, b,
[](const auto & x, const auto & y) { return x + y; });
}
};
template <typename T, std::size_t N, typename Elem, std::size_t K>
struct add_t<dpf::vec<T, N>, std::array<Elem, K>>
{
auto operator()(const std::array<Elem, K> & a, const std::array<Elem, K> & b) const
{
std::array<Elem, K> c{};
auto * dst = reinterpret_cast<unsigned char *>(c.data());
vec_detail::apply_bytes<T, N>(dst,
reinterpret_cast<const unsigned char *>(a.data()),
reinterpret_cast<const unsigned char *>(b.data()),
sizeof(c),
[](const auto & x, const auto & y) { return x + y; });
return c;
}
};
template <typename T, std::size_t N, typename NodeT>
struct subtract_t<dpf::vec<T, N>, NodeT,
std::enable_if_t<!std::is_void_v<NodeT>
&& !vec_detail::is_std_array<NodeT>::value>>
{
HEDLEY_ALWAYS_INLINE
auto operator()(const NodeT & a, const NodeT & b) const
{
return vec_detail::apply_node<T, N>(a, b,
[](const auto & x, const auto & y) { return x - y; });
}
};
template <typename T, std::size_t N, typename Elem, std::size_t K>
struct subtract_t<dpf::vec<T, N>, std::array<Elem, K>>
{
auto operator()(const std::array<Elem, K> & a, const std::array<Elem, K> & b) const
{
std::array<Elem, K> c{};
auto * dst = reinterpret_cast<unsigned char *>(c.data());
vec_detail::apply_bytes<T, N>(dst,
reinterpret_cast<const unsigned char *>(a.data()),
reinterpret_cast<const unsigned char *>(b.data()),
sizeof(c),
[](const auto & x, const auto & y) { return x - y; });
return c;
}
};
} // namespace leaf_arithmetic
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_VEC_HPP__

366
include/dpf/verifiable.hpp Normal file
View file

@ -0,0 +1,366 @@
/// @file dpf/verifiable.hpp
/// @brief Verifiable evaluation tokens and extractable-key helpers.
/// @details VDPF proof fold follows de Castro–Polychroniadou (hash-based
/// correction seeds, 2λ-bit tokens, equality Verify). Extractable
/// checks are public-part equality, ROM-style leaf XOF, and an
/// \(\mathbb{F}_{2^{61}-1}\) weight-1 subset sketch. Phantom tags
/// `dpf::verifiable` / `dpf::extractable` live in placement.hpp.
/// @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_VERIFIABLE_HPP__
#define LIBDPF_INCLUDE_DPF_VERIFIABLE_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <iterator>
#include <type_traits>
#include <utility>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "portable-snippets/exact-int/exact-int.h"
#include "dpf/placement.hpp"
#include "dpf/prg_aes.hpp"
#include "dpf/fp61.hpp"
#include "dpf/xor_wrapper.hpp"
#include "dpf/twiddle.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// 4λ = 64-byte correction seed (four AES blocks).
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using cs_block = std::array<simde__m128i, 4>;
/// 2λ = 32-byte proof token (two AES blocks).
using proof_token = std::array<simde__m128i, 2>;
HEDLEY_PRAGMA(GCC diagnostic pop)
namespace detail
{
namespace vdpf
{
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
simde__m128i mmo(simde__m128i seed, psnip_uint32_t pos) noexcept
{
return prg::aes128::eval(seed, pos);
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
cs_block hash_level_seed(std::size_t level, simde__m128i seed) noexcept
{
const simde__m128i tagged = simde_mm_xor_si128(seed,
simde_mm_set_epi64x(static_cast<psnip_int64_t>(0x56),
static_cast<psnip_int64_t>(level)));
return cs_block{
mmo(tagged, 0),
mmo(tagged, 1),
mmo(tagged, 2),
mmo(tagged, 3)};
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
cs_block hash_node(std::size_t level, psnip_uint64_t x_bits,
simde__m128i seed) noexcept
{
const simde__m128i tagged = simde_mm_xor_si128(seed,
simde_mm_set_epi64x(static_cast<psnip_int64_t>(0x5600 | (level & 0xff)),
static_cast<psnip_int64_t>(x_bits)));
return cs_block{
mmo(tagged, 0),
mmo(tagged, 1),
mmo(tagged, 2),
mmo(tagged, 3)};
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
cs_block make_cs(std::size_t level, psnip_uint64_t prefix_bits,
simde__m128i s0, simde__m128i s1) noexcept
{
const auto h0v = hash_node(level, prefix_bits, s0);
const auto h1v = hash_node(level, prefix_bits, s1);
return cs_block{
simde_mm_xor_si128(h0v[0], h1v[0]),
simde_mm_xor_si128(h0v[1], h1v[1]),
simde_mm_xor_si128(h0v[2], h1v[2]),
simde_mm_xor_si128(h0v[3], h1v[3])};
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
cs_block correct(cs_block pi_tilde, const cs_block & cs,
bool t) noexcept
{
if (!t)
return pi_tilde;
return cs_block{
simde_mm_xor_si128(pi_tilde[0], cs[0]),
simde_mm_xor_si128(pi_tilde[1], cs[1]),
simde_mm_xor_si128(pi_tilde[2], cs[2]),
simde_mm_xor_si128(pi_tilde[3], cs[3])};
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
proof_token h0(const cs_block & in) noexcept
{
const simde__m128i a = simde_mm_xor_si128(in[0], in[2]);
const simde__m128i b = simde_mm_xor_si128(in[1], in[3]);
return proof_token{mmo(a, 0x48), mmo(b, 0x48)};
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
proof_token xor_proof(proof_token a, proof_token b) noexcept
{
return proof_token{
simde_mm_xor_si128(a[0], b[0]),
simde_mm_xor_si128(a[1], b[1])};
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
proof_token zero_proof() noexcept
{
return proof_token{simde_mm_setzero_si128(), simde_mm_setzero_si128()};
}
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
void fold_node(proof_token & pi, std::size_t level,
psnip_uint64_t x_bits, simde__m128i seed, const cs_block & cs) noexcept
{
const bool t = static_cast<bool>(dpf::get_lo_bit(seed));
const cs_block tilde = hash_node(level, x_bits, seed);
const cs_block corrected = correct(tilde, cs, t);
cs_block mixed{
simde_mm_xor_si128(pi[0], corrected[0]),
simde_mm_xor_si128(pi[1], corrected[1]),
corrected[2],
corrected[3]};
pi = xor_proof(pi, h0(mixed));
}
HEDLEY_NO_THROW
inline void leaf_xof(simde__m128i seed, simde__m128i * HEDLEY_RESTRICT out,
psnip_uint32_t count, psnip_uint32_t pos = 0) noexcept
{
const simde__m128i tagged = simde_mm_xor_si128(seed,
simde_mm_set_epi64x(0x45, 0));
for (psnip_uint32_t i = 0; i < count; ++i)
out[i] = mmo(tagged, pos + i);
}
/// Drop-in exterior PRG for leaf stretch under `dpf::extractable`.
template <typename BasePRG>
struct extractable_leaf_prg
{
using block_type = typename BasePRG::block_type;
HEDLEY_NO_THROW
static void eval(block_type seed, block_type * HEDLEY_RESTRICT out,
psnip_uint32_t count, psnip_uint32_t pos = 0) noexcept
{
leaf_xof(seed, out, count, pos);
}
};
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
bool proof_equal(const proof_token & a, const proof_token & b) noexcept
{
return std::memcmp(&a, &b, sizeof(proof_token)) == 0;
}
template <typename KeyT>
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
void init_proof(proof_token & pi, const KeyT & /*key*/) noexcept
{
// Running proof starts at 0; each fold mixes in corrected leaf digests.
pi = zero_proof();
}
} // namespace vdpf
} // namespace detail
/// @brief A proof token the caller owns, passed into evaluation.
struct prove_ref
{
/// @brief The token updated by the evaluation.
proof_token & token;
/// @brief Bind `t`.
/// @param t the token to update
HEDLEY_NO_THROW
explicit prove_ref(proof_token & t) noexcept : token{t} { }
};
/// @brief Bind `t` as the proof accumulator for one evaluation.
/// @param t the token to update
/// @return a `prove_ref` bound to `t`
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
prove_ref prove(proof_token & t) noexcept
{
return prove_ref{t};
}
/// @brief Whether two proof tokens are identical.
/// @param a the first token
/// @param b the second token
/// @return `true` when every byte matches
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
bool verify(const proof_token & a, const proof_token & b) noexcept
{
return detail::vdpf::proof_equal(a, b);
}
/// @brief Fold two batches of proof tokens and compare them.
/// @tparam Range0 range of `proof_token` for party 0
/// @tparam Range1 range of `proof_token` for party 1
/// @param left party 0 tokens, in evaluation order
/// @param right party 1 tokens, in the same order
/// @return `false` when the ranges differ in length or the folded tokens differ
template <typename Range0, typename Range1>
bool verify_batch(Range0 && left, Range1 && right)
{
proof_token a = detail::vdpf::zero_proof();
proof_token b = detail::vdpf::zero_proof();
auto it0 = std::begin(left);
auto it1 = std::begin(right);
const auto end0 = std::end(left);
const auto end1 = std::end(right);
for (; it0 != end0 && it1 != end1; ++it0, ++it1)
{
a = detail::vdpf::xor_proof(a, *it0);
b = detail::vdpf::xor_proof(b, *it1);
a[0] = detail::vdpf::mmo(a[0], 1);
b[0] = detail::vdpf::mmo(b[0], 1);
}
if (it0 != end0 || it1 != end1)
return false;
return verify(a, b);
}
/// @brief Whether two keys publish the same correction words, advice, and hash.
/// @tparam KeyT0 key type of party 0
/// @tparam KeyT1 key type of party 1
/// @param k0 party 0 key
/// @param k1 party 1 key
/// @return `false` when a public field differs
template <typename KeyT0, typename KeyT1>
bool same_public_part(const KeyT0 & k0, const KeyT1 & k1)
{
static_assert(KeyT0::is_verifiable == KeyT1::is_verifiable,
"same_public_part: mismatched verifiable flags");
if (std::memcmp(k0.correction_words().data(), k1.correction_words().data(),
sizeof(typename KeyT0::correction_words_array)) != 0)
return false;
if (std::memcmp(k0.correction_advice().data(), k1.correction_advice().data(),
sizeof(typename KeyT0::correction_advice_array)) != 0)
return false;
if constexpr (KeyT0::is_verifiable)
{
if (std::memcmp(k0.correction_seeds().data(),
k1.correction_seeds().data(),
sizeof(typename KeyT0::correction_seeds_array)) != 0)
return false;
}
return std::memcmp(&k0.common_part_hash(), &k1.common_part_hash(),
sizeof(digest_type)) == 0;
}
struct sketch_share
{
fp61 z1{};
fp61 z2{};
fp61 z3{};
};
/// @brief Weight-1 subset sketch of payloads `ys` against challenges `rs`.
/// @tparam YRange range of integers convertible to `fp61`
/// @tparam RRange range of challenges, one per payload
/// @param ys the payloads
/// @param rs the challenges
/// @return the three folded moments. A short range stops at the shorter end
template <typename YRange, typename RRange>
sketch_share sketch_fold(YRange && ys, RRange && rs)
{
sketch_share out{};
auto iy = std::begin(ys);
auto ir = std::begin(rs);
const auto ey = std::end(ys);
const auto er = std::end(rs);
for (; iy != ey && ir != er; ++iy, ++ir)
{
const fp61 y{*iy};
const fp61 r{*ir};
const fp61 r2 = r * r;
out.z1 = out.z1 + y;
out.z2 = out.z2 + y * r;
out.z3 = out.z3 + y * r2;
}
return out;
}
/// @brief Whether `s0 - s1` is a weight-1 subset sketch.
/// @param s0 party 0's folded sketch
/// @param s1 party 1's folded sketch
/// @return `true` when `z2² = z1 · z3` after the shares are opened
HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
bool sketch_verify(sketch_share s0, sketch_share s1) noexcept
{
const fp61 z1 = s0.z1 - s1.z1;
const fp61 z2 = s0.z2 - s1.z2;
const fp61 z3 = s0.z3 - s1.z3;
return (z2 * z2) == (z1 * z3);
}
template <typename T, typename = void>
struct has_dpf_fp61 : std::false_type
{ };
template <typename T>
struct has_dpf_fp61<T, std::void_t<decltype(std::decay_t<T>::dpf_fp61)>>
: std::bool_constant<std::decay_t<T>::dpf_fp61>
{ };
template <typename T, typename = void>
struct extractable_codomain_ok
: std::bool_constant<
(utils::bitlength_of_v<std::decay_t<T>> >= 128)
|| has_dpf_fp61<T>::value>
{ };
template <typename T>
struct extractable_codomain_ok<xor_wrapper<T>, void>
: extractable_codomain_ok<T>
{ };
template <typename T>
inline constexpr bool extractable_codomain_ok_v =
extractable_codomain_ok<std::decay_t<T>>::value;
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_VERIFIABLE_HPP__

View file

@ -82,6 +82,7 @@ template <typename T> static constexpr wildcard_value<T> wildcard{};
/// @details A trait class that provides the member constant `value` which is /// @details A trait class that provides the member constant `value` which is
/// equal to `true`, if `T` is a specialization of the `wildcard_value` /// equal to `true`, if `T` is a specialization of the `wildcard_value`
/// template and `false` otherwise. /// template and `false` otherwise.
/// @tparam T value type
/// @see dpf::is_wildcard_v /// @see dpf::is_wildcard_v
template <typename T> struct is_wildcard : std::false_type { }; template <typename T> struct is_wildcard : std::false_type { };
template <typename T> struct is_wildcard<wildcard_value<T>> : std::true_type { }; template <typename T> struct is_wildcard<wildcard_value<T>> : std::true_type { };
@ -194,6 +195,7 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
static constexpr auto m256i = wildcard<simde__m256i>; static constexpr auto m256i = wildcard<simde__m256i>;
using m256d_t = wildcard_value<simde__m256d>; using m256d_t = wildcard_value<simde__m256d>;
static constexpr auto m256d = wildcard<simde__m256d>; static constexpr auto m256d = wildcard<simde__m256d>;
HEDLEY_PRAGMA(GCC diagnostic pop)
// using m512_t = wildcard_value<simde__m512>; // using m512_t = wildcard_value<simde__m512>;
// static constexpr auto m512 = wildcard<simde__m512>; // static constexpr auto m512 = wildcard<simde__m512>;
@ -201,7 +203,6 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
// static constexpr auto m512i = wildcard<simde__m512i>; // static constexpr auto m512i = wildcard<simde__m512i>;
// using m512d_t = wildcard_value<simde__m512d>; // using m512d_t = wildcard_value<simde__m512d>;
// static constexpr auto m512d = wildcard<simde__m512d>; // static constexpr auto m512d = wildcard<simde__m512d>;
HEDLEY_PRAGMA(GCC diagnostic pop)
using ieee_float_t = wildcard_value<float>; using ieee_float_t = wildcard_value<float>;
/// @brief Placeholder for a `float` whose leaf group is bitwise XOR /// @brief Placeholder for a `float` whose leaf group is bitwise XOR
@ -257,12 +258,15 @@ namespace utils
{ {
/// @brief specializes `dpf::utils::bitlength_of` for `dpf::wildcard_value` /// @brief specializes `dpf::utils::bitlength_of` for `dpf::wildcard_value`
/// @tparam T value type
template <typename T> template <typename T>
struct bitlength_of<wildcard_value<T>> struct bitlength_of<wildcard_value<T>>
: public bitlength_of<T> : public bitlength_of<T>
{ }; { };
/// @brief specializes `dpf::utils::bitlength_of_output` for `dpf::wildcard_value` /// @brief specializes `dpf::utils::bitlength_of_output` for `dpf::wildcard_value`
/// @tparam T value type
/// @tparam NodeT GGM node type
template <typename T, template <typename T,
typename NodeT> typename NodeT>
struct bitlength_of_output<wildcard_value<T>, NodeT> struct bitlength_of_output<wildcard_value<T>, NodeT>

View file

@ -57,6 +57,7 @@ struct xor_wrapper
constexpr xor_wrapper(xor_wrapper &&) noexcept = default; constexpr xor_wrapper(xor_wrapper &&) noexcept = default;
/// @brief Value c'tor /// @brief Value c'tor
/// @param v the `v`
// cppcheck-suppress noExplicitConstructor // cppcheck-suppress noExplicitConstructor
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr xor_wrapper(value_type v) noexcept : value{v} { } // NOLINT(runtime/explicit) constexpr xor_wrapper(value_type v) noexcept : value{v} { } // NOLINT(runtime/explicit)
@ -896,6 +897,7 @@ template <typename T>
struct is_xor_wrapper<xor_wrapper<T>> : std::true_type {}; struct is_xor_wrapper<xor_wrapper<T>> : std::true_type {};
/// @brief specializes `dpf::utils::bitlength_of` for `xor_wrapper` /// @brief specializes `dpf::utils::bitlength_of` for `xor_wrapper`
/// @tparam T value type
template <typename T> template <typename T>
struct bitlength_of<xor_wrapper<T>> struct bitlength_of<xor_wrapper<T>>
: public bitlength_of<T> : public bitlength_of<T>
@ -953,23 +955,27 @@ namespace std
/// @{ /// @{
/// @details specializes `std::numeric_limits` for `xor_wrapper<T>` /// @details specializes `std::numeric_limits` for `xor_wrapper<T>`
/// @tparam T value type
template<typename T> template<typename T>
class numeric_limits<dpf::xor_wrapper<T>> class numeric_limits<dpf::xor_wrapper<T>>
: public numeric_limits<dpf::utils::make_unsigned_t<T>> {}; : public numeric_limits<dpf::utils::make_unsigned_t<T>> {};
/// @details specializes `std::numeric_limits` for `xor_wrapper<T> const` /// @details specializes `std::numeric_limits` for `xor_wrapper<T> const`
/// @tparam T value type
template<typename T> template<typename T>
class numeric_limits<dpf::xor_wrapper<T> const> class numeric_limits<dpf::xor_wrapper<T> const>
: public numeric_limits<dpf::xor_wrapper<T>> {}; : public numeric_limits<dpf::xor_wrapper<T>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `xor_wrapper<T> volatile` /// @brief `xor_wrapper<T> volatile`
/// @tparam T value type
template<typename T> template<typename T>
class numeric_limits<dpf::xor_wrapper<T> volatile> class numeric_limits<dpf::xor_wrapper<T> volatile>
: public numeric_limits<dpf::xor_wrapper<T>> {}; : public numeric_limits<dpf::xor_wrapper<T>> {};
/// @details specializes `std::numeric_limits` for /// @details specializes `std::numeric_limits` for
/// `xor_wrapper<T> const volatile` /// @brief `xor_wrapper<T> const volatile`
/// @tparam T value type
template<typename T> template<typename T>
class numeric_limits<dpf::xor_wrapper<T> const volatile> class numeric_limits<dpf::xor_wrapper<T> const volatile>
: public numeric_limits<dpf::xor_wrapper<T>> {}; : public numeric_limits<dpf::xor_wrapper<T>> {};

View file

@ -1,6 +1,6 @@
/// @file dpf/zip_iterable.hpp /// @file dpf/zip_iterable.hpp
/// @brief defines the `dpf::zip_itrable` class and associated helpers /// @brief defines the `dpf::zip_itrable` class and associated helpers
/// @details /// @details Walks several iterables in lockstep and yields one tuple per step.
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors) /// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;

View file

@ -25,7 +25,7 @@
namespace grotto namespace grotto
{ {
/// Appendix D gadgets whose polynomial degree is 0 and whose max error is 0. /// @brief Appendix D gadgets whose polynomial degree is 0 and whose max error is 0.
enum class exact_constant enum class exact_constant
{ {
signum, signum,
@ -36,7 +36,7 @@ enum class exact_constant
zero, zero,
nonzero, nonzero,
ilogb, ilogb,
/// `ceil(log2(|x|))`. Exact powers of two agree with `ilogb`; every other /// @brief `ceil(log2(|x|))`. Exact powers of two agree with `ilogb`; every other
/// positive magnitude is one larger. Zero uses the same `-64` sentinel. /// positive magnitude is one larger. Zero uses the same `-64` sentinel.
ceil_ilogb, ceil_ilogb,
ilog10, ilog10,
@ -51,20 +51,21 @@ struct constant_lut
using raw_type = Raw; using raw_type = Raw;
/// Signed piece starts. `bounds.front()` is `numeric_limits<Raw>::min()`, /// @brief Signed piece starts. `bounds.front()` is `numeric_limits<Raw>::min()`,
/// and the starts are strictly increasing. /// and the starts are strictly increasing.
std::vector<Raw> bounds; std::vector<Raw> bounds;
/// `values[i]` is the function on `[bounds[i], next)`, where `next` is /// @brief `values[i]` is the function on `[bounds[i], next)`, where `next` is
/// `bounds[i + 1]` or one past `numeric_limits<Raw>::max()` for the last piece. /// `bounds[i + 1]` or one past `numeric_limits<Raw>::max()` for the last piece.
std::vector<std::int64_t> values; std::vector<std::int64_t> values;
HEDLEY_NO_THROW HEDLEY_NO_THROW
std::size_t linear_parts() const noexcept { return values.size(); } std::size_t linear_parts() const noexcept { return values.size(); }
/// Pieces after joining the first and last when they carry the same value. /// @brief Pieces after joining the first and last when they carry the same value.
/// Those two meet across the signed wrap, which is how the paper counts /// @details Those two meet across the signed wrap, which is how the paper counts
/// parts for `zero` and `nonzero` (2, not 3). /// parts for `zero` and `nonzero` (2, not 3).
/// @return Pieces after joining the first and last when they carry the same value
HEDLEY_NO_THROW HEDLEY_NO_THROW
std::size_t wrapped_parts() const noexcept std::size_t wrapped_parts() const noexcept
{ {
@ -115,7 +116,7 @@ constexpr bool shift_fits(u128 value, unsigned shift) noexcept
return shift < 128 && value <= (~u128{0} >> shift); return shift < 128 && value <= (~u128{0} >> shift);
} }
/// 10^0 .. 10^19. Every ilog10 projection reads this one table. /// @brief 10^0 .. 10^19. Every ilog10 projection reads this one table.
inline constexpr std::uint64_t pow10[] = { inline constexpr std::uint64_t pow10[] = {
1ull, 1ull,
10ull, 10ull,
@ -159,7 +160,10 @@ inline bool magnitude_ge_pow10(u128 mag, int k, unsigned fractional_bits) noexce
return mag * scale >= (u128{1} << fractional_bits); return mag * scale >= (u128{1} << fractional_bits);
} }
/// Smallest positive magnitude whose base-10 log is at least `k`. /// @brief Smallest positive magnitude whose base-10 log is at least `k`.
/// @param k the `k`
/// @param fractional_bits the number of fractional bits
/// @return Smallest positive magnitude whose base-10 log is at least `k`
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline u128 first_magnitude_at_least_pow10(int k, unsigned fractional_bits) noexcept inline u128 first_magnitude_at_least_pow10(int k, unsigned fractional_bits) noexcept
{ {
@ -315,8 +319,8 @@ void push_pow2_cuts(std::vector<Raw> & cuts)
detail::push_both_signs<Raw>(detail::u128{1} << k, cuts); detail::push_both_signs<Raw>(detail::u128{1} << k, cuts);
} }
/// floor(log2(|raw|)) - F, with -64 on the |x| <= 2^{-64} class (including 0). /// @brief floor(log2(|raw|)) - F, with -64 on the |x| <= 2^{-64} class (including 0).
/// One exponent program; fractional precision only shifts the stored exponent. /// @details One exponent program; fractional precision only shifts the stored exponent.
template <> template <>
struct exact_lut<exact_constant::ilogb> struct exact_lut<exact_constant::ilogb>
{ {
@ -347,7 +351,7 @@ struct exact_lut<exact_constant::ilogb>
} }
}; };
/// ceil(log2(|x|)). Same powers-of-two cuts as `ilogb`; exact powers keep the /// @brief ceil(log2(|x|)). Same powers-of-two cuts as `ilogb`; exact powers keep the
/// floor exponent and every other magnitude steps up by one. /// floor exponent and every other magnitude steps up by one.
template <> template <>
struct exact_lut<exact_constant::ceil_ilogb> struct exact_lut<exact_constant::ceil_ilogb>
@ -387,7 +391,7 @@ struct exact_lut<exact_constant::ceil_ilogb>
} }
}; };
/// floor(log10(|x|)), with -19 on |x| <= 10^{-19}. Thresholds come from `pow10`. /// @brief floor(log10(|x|)), with -19 on |x| <= 10^{-19}. Thresholds come from `pow10`.
template <> template <>
struct exact_lut<exact_constant::ilog10> struct exact_lut<exact_constant::ilog10>
{ {
@ -424,8 +428,8 @@ struct exact_lut<exact_constant::ilog10>
} }
}; };
/// 64-bit leading-zero count of trunc(x). Negatives are 0; a zero integer part is 64. /// @brief 64-bit leading-zero count of trunc(x). Negatives are 0; a zero integer part is 64.
/// Exponent k of the integer part maps to 63-k after a shift of `fractional_bits`. /// @details Exponent k of the integer part maps to 63-k after a shift of `fractional_bits`.
template <> template <>
struct exact_lut<exact_constant::clz> struct exact_lut<exact_constant::clz>
{ {
@ -461,8 +465,8 @@ struct exact_lut<exact_constant::clz>
} }
}; };
/// 64-bit redundant sign bits of trunc(x) toward zero. /// @brief 64-bit redundant sign bits of trunc(x) toward zero.
/// Positive q uses 62-floor(log2(q)); negative q uses 62-floor(log2(q-1)). /// @details Positive q uses 62-floor(log2(q)); negative q uses 62-floor(log2(q-1)).
template <> template <>
struct exact_lut<exact_constant::clrsb> struct exact_lut<exact_constant::clrsb>
{ {
@ -564,7 +568,7 @@ constant_lut<Raw> make_exact_constant_lut(exact_constant which, unsigned fractio
throw std::invalid_argument("unknown exact constant"); throw std::invalid_argument("unknown exact constant");
} }
/// Comparison against a public threshold. Two pieces; the cut sits on `bound` /// @brief Comparison against a public threshold. Two pieces; the cut sits on `bound`
/// (`lt` / `geq`) or just after it (`leq` / `gt`). /// (`lt` / `geq`) or just after it (`leq` / `gt`).
enum class threshold_cmp enum class threshold_cmp
{ {
@ -600,7 +604,12 @@ constant_lut<Raw> make_threshold_lut(Raw bound, threshold_cmp kind)
[=](Raw raw) { return evaluate_threshold(raw, bound, kind); }); [=](Raw raw) { return evaluate_threshold(raw, bound, kind); });
} }
/// `1` on the inclusive clip window `[low, high]`, `0` outside it. /// @brief `1` on the inclusive clip window `[low, high]`, `0` outside it.
/// @tparam Raw underlying representation
/// @param low the lower endpoint
/// @param high the upper endpoint
/// @return `1` on the inclusive clip window `[low, high]`, `0` outside it
/// @throws std::invalid_argument if `low > high`
template <typename Raw> template <typename Raw>
constant_lut<Raw> make_interval_lut(Raw low, Raw high) constant_lut<Raw> make_interval_lut(Raw low, Raw high)
{ {
@ -614,13 +623,20 @@ constant_lut<Raw> make_interval_lut(Raw low, Raw high)
[=](Raw raw) { return raw >= low && raw <= high ? std::int64_t{1} : std::int64_t{0}; }); [=](Raw raw) { return raw >= low && raw <= high ? std::int64_t{1} : std::int64_t{0}; });
} }
/// `floor(min(max(raw, low), high) / modulus)`, division toward -infinity. /// @brief `floor(min(max(raw, low), high) / modulus)`, division toward -infinity.
/// ///
/// `modulus`, `low`, and `high` are in the same raw units as the domain, so /// `modulus`, `low`, and `high` are in the same raw units as the domain, so
/// one program covers every fractional precision: a mathematical step `M` /// one program covers every fractional precision: a mathematical step `M`
/// with `F` fractional bits is the raw modulus `M << F`. The paper's /// with `F` fractional bits is the raw modulus `M << F`. The paper's
/// `quot(M, T1, T2)` is this function. Piece count is about `(high-low)/modulus`; /// `quot(M, T1, T2)` is this function. Piece count is about `(high-low)/modulus`;
/// the build rejects windows that would need more than 2^16 pieces. /// the build rejects windows that would need more than 2^16 pieces.
/// @tparam Raw underlying representation
/// @param raw the underlying integer
/// @param modulus the public modulus
/// @param low the lower endpoint
/// @param high the upper endpoint
/// @return `floor(min(max(raw, low), high) / modulus)`, division toward -infinity
/// @throws std::invalid_argument if `modulus must be positive`
template <typename Raw> template <typename Raw>
std::int64_t evaluate_clipped_quotient(Raw raw, Raw modulus, Raw low, Raw high) std::int64_t evaluate_clipped_quotient(Raw raw, Raw modulus, Raw low, Raw high)
{ {

View file

@ -26,12 +26,12 @@
namespace grotto namespace grotto
{ {
/// Sentinel raw value for `ilogb(0)` and `ilog10(0)`. /// @brief Sentinel raw value for `ilogb(0)` and `ilog10(0)`.
inline constexpr std::int64_t ilog_of_zero = inline constexpr std::int64_t ilog_of_zero =
std::numeric_limits<std::int64_t>::min(); std::numeric_limits<std::int64_t>::min();
/// `make_msb_lut(i)` allows `i` in `[0, msb_bit_limit)`. /// @brief `make_msb_lut(i)` allows `i` in `[0, msb_bit_limit)`.
/// Bit 0 is two intervals; bit 7 is 256. /// @details Bit 0 is two intervals; bit 7 is 256.
inline constexpr unsigned msb_bit_limit = 8; inline constexpr unsigned msb_bit_limit = 8;
namespace detail namespace detail
@ -162,7 +162,11 @@ inline u128 pow10_u128(int exponent)
return p; return p;
} }
/// `mag / 2^k >= 10^e`. /// @brief `mag / 2^k >= 10^e`.
/// @param mag the magnitude
/// @param fractional_bits the number of fractional bits
/// @param exponent the exponent
/// @return `mag / 2^k >= 10^e`
inline bool magnitude_ge_pow10(u128 mag, unsigned fractional_bits, int exponent) inline bool magnitude_ge_pow10(u128 mag, unsigned fractional_bits, int exponent)
{ {
if (mag == 0) if (mag == 0)
@ -393,9 +397,14 @@ easy_lut<Raw> make_ilog10_lut(unsigned fractional_bits = 0)
}); });
} }
/// Bit `index` counting down from the most significant bit of `Raw`. /// @brief Bit `index` counting down from the most significant bit of `Raw`.
/// Index 0 is the sign bit. Larger indexes are refused: the bit is constant /// @details Index 0 is the sign bit. Larger indexes are refused: the bit is constant
/// on `2^{index+1}` intervals. /// on `2^{index+1}` intervals.
/// @tparam Raw underlying representation
/// @param index the index
/// @param fractional_bits the number of fractional bits
/// @return Bit `index` counting down from the most significant bit of `Raw`
/// @throws std::invalid_argument if `only the most significant bits are piecewise-cheap`
template <typename Raw> template <typename Raw>
HEDLEY_WARN_UNUSED_RESULT HEDLEY_WARN_UNUSED_RESULT
easy_lut<Raw> make_msb_lut(unsigned index, unsigned fractional_bits = 0) easy_lut<Raw> make_msb_lut(unsigned index, unsigned fractional_bits = 0)

View file

@ -23,7 +23,8 @@
namespace grotto namespace grotto
{ {
/// Piece `y_raw = round((c0 + c1·raw + c2·raw²) / den)`. /// @brief Piece `y_raw = round((c0 + c1·raw + c2·raw²) / den)`.
/// @tparam Raw underlying representation
template <typename Raw> template <typename Raw>
struct easy_lut struct easy_lut
{ {
@ -176,8 +177,11 @@ easy_lut<Raw> make_relu_lut(unsigned fractional_bits = 0)
}); });
} }
/// Negative side is `x / 2^shift`, rounded to nearest, ties away from zero. /// @brief Negative side is `x / 2^shift`, rounded to nearest, ties away from zero.
/// `shift == 0` is the identity. The slope does not depend on fractional width. /// @details `shift == 0` is the identity. The slope does not depend on fractional width.
/// @tparam Raw underlying representation
/// @param shift the bit shift
/// @return Negative side is `x / 2^shift`, rounded to nearest, ties away from zero
template <typename Raw> template <typename Raw>
easy_lut<Raw> make_leaky_relu_lut(unsigned shift) easy_lut<Raw> make_leaky_relu_lut(unsigned shift)
{ {
@ -248,7 +252,10 @@ easy_lut<Raw> make_hardtanh_lut(unsigned fractional_bits)
return make_clip_lut<Raw>(fractional_bits, -1, 1); return make_clip_lut<Raw>(fractional_bits, -1, 1);
} }
/// `0` on `[-1, 1]`, `x - 1` above, `x + 1` below. /// @brief `0` on `[-1, 1]`, `x - 1` above, `x + 1` below.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` on `[-1, 1]`, `x - 1` above, `x + 1` below
template <typename Raw> template <typename Raw>
easy_lut<Raw> make_softshrink_lut(unsigned fractional_bits) easy_lut<Raw> make_softshrink_lut(unsigned fractional_bits)
{ {
@ -270,7 +277,10 @@ easy_lut<Raw> make_softshrink_lut(unsigned fractional_bits)
}); });
} }
/// `0` on `[-1, 1]`, identity outside. Lambda is the integer 1. /// @brief `0` on `[-1, 1]`, identity outside. Lambda is the integer 1.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` on `[-1, 1]`, identity outside
template <typename Raw> template <typename Raw>
easy_lut<Raw> make_hardshrink_lut(unsigned fractional_bits) easy_lut<Raw> make_hardshrink_lut(unsigned fractional_bits)
{ {
@ -290,7 +300,11 @@ easy_lut<Raw> make_hardshrink_lut(unsigned fractional_bits)
}); });
} }
/// `0` left of `-3`, `1` right of `3`, `(x + 3) / 6` between, rounded. /// @brief `0` left of `-3`, `1` right of `3`, `(x + 3) / 6` between, rounded.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` left of `-3`, `1` right of `3`, `(x + 3) / 6` between, rounded
/// @throws std::invalid_argument if `fractional width does not fit`
template <typename Raw> template <typename Raw>
easy_lut<Raw> make_hardsigmoid_lut(unsigned fractional_bits) easy_lut<Raw> make_hardsigmoid_lut(unsigned fractional_bits)
{ {
@ -316,7 +330,11 @@ easy_lut<Raw> make_hardsigmoid_lut(unsigned fractional_bits)
}); });
} }
/// `0` left of `-3`, `x` right of `3`, `x(x + 3) / 6` between, rounded. /// @brief `0` left of `-3`, `x` right of `3`, `x(x + 3) / 6` between, rounded.
/// @tparam Raw underlying representation
/// @param fractional_bits the number of fractional bits
/// @return `0` left of `-3`, `x` right of `3`, `x(x + 3) / 6` between, rounded
/// @throws std::invalid_argument if `fractional width does not fit the denominator`
template <typename Raw> template <typename Raw>
easy_lut<Raw> make_hardswish_lut(unsigned fractional_bits) easy_lut<Raw> make_hardswish_lut(unsigned fractional_bits)
{ {

View file

@ -1,6 +1,5 @@
/// @file grotto/fixedpoint.hpp /// @file grotto/fixedpoint.hpp
/// @brief /// @brief Fixed-point values stored in an integer backend.
/// @details
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
@ -40,6 +39,8 @@ namespace detail
/// @brief Integer value of an already-rounded finite double, as a 256-bit word. /// @brief Integer value of an already-rounded finite double, as a 256-bit word.
/// Values that do not fit saturate to all-ones. /// Values that do not fit saturate to all-ones.
/// @param rounded the `rounded`
/// @return Integer value of an already-rounded finite double, as a 256-bit word
HEDLEY_NO_THROW HEDLEY_NO_THROW
inline uint256_t uint256_from_rounded_double(double rounded) noexcept inline uint256_t uint256_from_rounded_double(double rounded) noexcept
{ {
@ -98,7 +99,11 @@ struct is_static_castable<To, From,
std::void_t<decltype(static_cast<To>(std::declval<From>()))>> std::void_t<decltype(static_cast<To>(std::declval<From>()))>>
: std::true_type {}; : std::true_type {};
/// Low `bits` of `wide`, saturated to all-ones when `wide` does not fit. /// @brief Low `bits` of `wide`, saturated to all-ones when `wide` does not fit.
/// @tparam Raw underlying representation
/// @tparam Bits bits
/// @param wide the `wide`
/// @return Low `bits` of `wide`, saturated to all-ones when `wide` does not fit
template <typename Raw, std::size_t Bits> template <typename Raw, std::size_t Bits>
HEDLEY_NO_THROW HEDLEY_NO_THROW
Raw saturate_low_bits(uint256_t wide) noexcept Raw saturate_low_bits(uint256_t wide) noexcept
@ -202,8 +207,13 @@ inline IntegralType rounded_double_to_integral(double rounded) noexcept
} }
} }
/// Shift an integer into fixed-point raw form: `value * 2^FractionalBits`, /// @brief Shift an integer into fixed-point raw form: `value * 2^FractionalBits`,
/// wrapping in the backend's two's-complement encoding. One shift; no `double`. /// wrapping in the backend's two's-complement encoding. One shift; no `double`.
/// @tparam IntegralType underlying integral type
/// @tparam FractionalBits number of fractional bits
/// @tparam T value type
/// @param integer_value the `integer_value`
/// @return the returned `IntegralType`
template <typename IntegralType, template <typename IntegralType,
unsigned FractionalBits, unsigned FractionalBits,
typename T> typename T>
@ -227,8 +237,11 @@ inline constexpr bool is_signed_rep_v =
std::is_signed_v<IntegralType> std::is_signed_v<IntegralType>
|| std::is_same_v<IntegralType, simde_int128>; || std::is_same_v<IntegralType, simde_int128>;
/// Two's-complement negate via the unsigned width. Defined for the /// @brief Two's-complement negate via the unsigned width. Defined for the
/// most-negative value (wraps); signed `-x` would be UB there. /// most-negative value (wraps); signed `-x` would be UB there.
/// @tparam IntegralType underlying integral type
/// @param x the `x`
/// @return Two's-complement negate via the unsigned width
template <typename IntegralType> template <typename IntegralType>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_CONST HEDLEY_CONST
@ -253,8 +266,12 @@ constexpr IntegralType raw_abs(IntegralType x) noexcept
return x; return x;
} }
/// Remainder with the sign of `a` and magnitude `< |b|` (C++ `%` / /// @brief Remainder with the sign of `a` and magnitude `< |b|` (C++ `%` /
/// `std::fmod`). Zero divisor → 0; this type has no NaN. /// `std::fmod`). Zero divisor → 0; this type has no NaN.
/// @tparam IntegralType underlying integral type
/// @param a the `a`
/// @param b the `b`
/// @return Remainder with the sign of `a` and magnitude `< |b|` (C++ `%` / `std::fmod`)
template <typename IntegralType> template <typename IntegralType>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_CONST HEDLEY_CONST
@ -276,7 +293,7 @@ HEDLEY_NO_THROW
auto constexpr make_fixed_from_integral_type(IntegralType value) noexcept; auto constexpr make_fixed_from_integral_type(IntegralType value) noexcept;
/// @tparam FractionalBits Number of fractional bits used in the fixed-point /// @tparam FractionalBits Number of fractional bits used in the fixed-point
/// representation. /// @brief representation.
/// @tparam IntegralType The underlying integral type used for the fixed-point /// @tparam IntegralType The underlying integral type used for the fixed-point
/// representation. /// representation.
template <unsigned FractionalBits, template <unsigned FractionalBits,
@ -310,18 +327,21 @@ public:
/// @brief Copy c'tor /// @brief Copy c'tor
/// @details Constructs a fixed-point with the value copied from `other`. /// @details Constructs a fixed-point with the value copied from `other`.
/// @param other the value to compare or copy
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint(const fixedpoint & other) noexcept = default; constexpr fixedpoint(const fixedpoint & other) noexcept = default;
/// @brief Move c'tor /// @brief Move c'tor
/// @details Constructs a fixed-point with the value copied from `other` using move semantics. /// @details Constructs a fixed-point with the value copied from `other` using move semantics.
/// @param other the value to compare or copy
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint(fixedpoint && other) noexcept = default; constexpr fixedpoint(fixedpoint && other) noexcept = default;
/// @brief Value c'tor /// @brief Value c'tor
/// @details Initializes the fixed-point with the value determined by `desired`, using the <a href="https://en.cppreference.com/w/cpp/numeric/fenv/FE_round">current rounding mode</a> for the least-significant bit. /// @details Initializes the fixed-point with the value determined by `desired`, using the <a href="https://en.cppreference.com/w/cpp/numeric/fenv/FE_round">current rounding mode</a> for the least-significant bit.
/// @param desired the `desired`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint(double desired) noexcept // NOLINT (implicit c'tor) constexpr fixedpoint(double desired) noexcept // NOLINT (implicit c'tor)
@ -333,6 +353,9 @@ public:
/// @details `fixedpoint(3)` is the mathematical value 3 (raw encoding /// @details `fixedpoint(3)` is the mathematical value 3 (raw encoding
/// `3 << fractional_bits`), not a raw word. One shift; no `double`. /// `3 << fractional_bits`), not a raw word. One shift; no `double`.
/// Use `from_raw` for a bit-exact encoding. /// Use `from_raw` for a bit-exact encoding.
/// @tparam T value type
/// @tparam T value type
/// @param integer_value the `integer_value`
template <typename T, template <typename T,
std::enable_if_t< std::enable_if_t<
std::is_integral_v<T> std::is_integral_v<T>
@ -345,6 +368,8 @@ public:
{ } { }
/// @brief Bit-exact construction from the backend integer encoding. /// @brief Bit-exact construction from the backend integer encoding.
/// @param raw the underlying integer
/// @return Bit-exact construction from the backend integer encoding
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr fixedpoint from_raw(integral_type raw) noexcept static constexpr fixedpoint from_raw(integral_type raw) noexcept
@ -354,24 +379,30 @@ public:
/// @} /// @}
/// @name Assignment operators /// @name Assignment operators
/// @brief Assign a new value to a fixed-point number /// @brief Assign a new value to a fixed-point number
/// {@ /// @{
/// @brief Copy assignment /// @brief Copy assignment
/// @details Assigns the fixed-point with a copy of `other` /// @details Assigns the fixed-point with a copy of `other`
/// @param other the value to compare or copy
/// @return `*this`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint & operator=(const fixedpoint & other) noexcept = default; constexpr fixedpoint & operator=(const fixedpoint & other) noexcept = default;
/// @brief Move assignment /// @brief Move assignment
/// @details Assigns the fixed-point with a copy of `other` using move semantics. /// @details Assigns the fixed-point with a copy of `other` using move semantics.
/// @param other the value to compare or copy
/// @return `*this`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint & operator=(fixedpoint && other) noexcept = default; constexpr fixedpoint & operator=(fixedpoint && other) noexcept = default;
/// @brief Value assignment /// @brief Value assignment
/// @details Assigns the fixed-point with a value determined by `desired`, using the <a href="https://en.cppreference.com/w/cpp/numeric/fenv/FE_round">current rounding mode</a> for the least-significant bit.. /// @details Assigns the fixed-point with a value determined by `desired`, using the <a href="https://en.cppreference.com/w/cpp/numeric/fenv/FE_round">current rounding mode</a> for the least-significant bit..
/// @param desired the `desired`
/// @return `*this`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint & operator=(const double & desired) noexcept constexpr fixedpoint & operator=(const double & desired) noexcept
@ -386,6 +417,7 @@ public:
~fixedpoint() = default; ~fixedpoint() = default;
/// @brief Cast to `double` /// @brief Cast to `double`
/// @return Cast to `double`
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_PURE HEDLEY_PURE
@ -402,8 +434,11 @@ public:
return static_cast<bool>(this->integral_representation() & mask); return static_cast<bool>(this->integral_representation() & mask);
} }
/// Bit test against another encoding (DPF writes `mask & x` with both /// @brief Bit test against another encoding (DPF writes `mask & x` with both
/// sides the input type when `msb_mask` is a `fixedpoint`). /// sides the input type when `msb_mask` is a `fixedpoint`).
/// @param mask the bit mask
/// @return Bit test against another encoding (DPF writes `mask & x` with both sides the input
/// type when `msb_mask` is a `fixedpoint`)
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -412,7 +447,8 @@ public:
return static_cast<bool>(value & mask.value); return static_cast<bool>(value & mask.value);
} }
/// Bitwise complement of the encoding. `std::bit_not` uses this. /// @brief Bitwise complement of the encoding. `std::bit_not` uses this.
/// @return Bitwise complement of the encoding
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -423,7 +459,8 @@ public:
~static_cast<unsigned_type>(value))); ~static_cast<unsigned_type>(value)));
} }
/// Next / previous representable encoding (one ULP). /// @brief Next / previous representable encoding (one ULP).
/// @return `*this`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint & operator++() noexcept constexpr fixedpoint & operator++() noexcept
@ -458,7 +495,7 @@ public:
return tmp; return tmp;
} }
/// Logical shift of the encoding. DPF walks `msb_mask` with `>>`; a /// @brief Logical shift of the encoding. DPF walks `msb_mask` with `>>`; a
/// signed arithmetic shift would sign-extend the MSB and break that. /// signed arithmetic shift would sign-extend the MSB and break that.
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -497,6 +534,7 @@ public:
/// @brief Access underlying integral representation /// @brief Access underlying integral representation
/// @details If the represented fixed-point number is `x`, then this /// @details If the represented fixed-point number is `x`, then this
/// function returns an `integral_type` whose value is `x*2**fractional_bits`. /// function returns an `integral_type` whose value is `x*2**fractional_bits`.
/// @return Access underlying integral representation
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -506,6 +544,7 @@ public:
} }
/// @brief Unary negation operator /// @brief Unary negation operator
/// @return Unary negation operator
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -516,6 +555,8 @@ public:
/// @brief Binary addition operator /// @brief Binary addition operator
/// @details Computes the sum of two fixed-point numbers /// @details Computes the sum of two fixed-point numbers
/// @param rhs the right-hand operand
/// @return Binary addition operator
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -525,6 +566,8 @@ public:
} }
/// @brief Binary addition assignment operator /// @brief Binary addition assignment operator
/// @param rhs the right-hand operand
/// @return `*this`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint & operator+=(fixedpoint rhs) noexcept constexpr fixedpoint & operator+=(fixedpoint rhs) noexcept
@ -534,6 +577,8 @@ public:
} }
/// @brief Binary subtraction operator /// @brief Binary subtraction operator
/// @param rhs the right-hand operand
/// @return Binary subtraction operator
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
HEDLEY_PURE HEDLEY_PURE
@ -543,6 +588,8 @@ public:
} }
/// @brief Binary addition assignment operator /// @brief Binary addition assignment operator
/// @param rhs the right-hand operand
/// @return `*this`
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr fixedpoint & operator-=(fixedpoint rhs) noexcept constexpr fixedpoint & operator-=(fixedpoint rhs) noexcept
@ -552,6 +599,9 @@ public:
} }
/// @brief Binary multiplication operator /// @brief Binary multiplication operator
/// @tparam FractionalBits1 fractional bits1
/// @param rhs the right-hand operand
/// @return Binary multiplication operator
template <unsigned FractionalBits1> template <unsigned FractionalBits1>
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
@ -701,6 +751,8 @@ public:
// struct make_fixed_from_integral_type_tag {}; // struct make_fixed_from_integral_type_tag {};
/// @brief Determine if a floating-point is within range /// @brief Determine if a floating-point is within range
/// @param d the `d`
/// @return Determine if a floating-point is within range
HEDLEY_ALWAYS_INLINE HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW HEDLEY_NO_THROW
static constexpr bool is_in_range(double d) noexcept static constexpr bool is_in_range(double d) noexcept
@ -734,6 +786,12 @@ public:
/// @brief Bit test with the mask on the left. DPF key generation and /// @brief Bit test with the mask on the left. DPF key generation and
/// evaluation write `mask & x`. /// evaluation write `mask & x`.
/// @tparam FractionalBits number of fractional bits
/// @tparam IntegralType underlying integral type
/// @tparam Mask mask
/// @param mask the bit mask
/// @param x the `x`
/// @return Bit test with the mask on the left
template <unsigned FractionalBits, template <unsigned FractionalBits,
typename IntegralType, typename IntegralType,
typename Mask> typename Mask>
@ -797,7 +855,11 @@ static constexpr auto make_fixed(double d)
} }
/// @brief Creates a fixed-point number from a double with bounds checking. /// @brief Creates a fixed-point number from a double with bounds checking.
/// @throws std::range_error If the input double is outside the representable /// @tparam FractionalBits number of fractional bits
/// @tparam IntegralType underlying integral type
/// @param d the `d`
/// @return Creates a fixed-point number from a double with bounds checking
/// @throws std::range_error if the input double is outside the representable
/// range of the fixed-point number. /// range of the fixed-point number.
template <unsigned FractionalBits, template <unsigned FractionalBits,
typename IntegralType = GROTTO_FIXED_DEFAULT_INTEGRAL_REPRESENTATION> typename IntegralType = GROTTO_FIXED_DEFAULT_INTEGRAL_REPRESENTATION>
@ -1562,6 +1624,8 @@ struct flip_msb_for_input<grotto::fixedpoint<FractionalBits, IntegralType>>
namespace dpf::leaf_arithmetic namespace dpf::leaf_arithmetic
{ {
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
template <unsigned FractionalBits, typename IntegralType> template <unsigned FractionalBits, typename IntegralType>
struct add_t<grotto::fixedpoint<FractionalBits, IntegralType>, simde__m128i> struct add_t<grotto::fixedpoint<FractionalBits, IntegralType>, simde__m128i>
{ {
@ -1617,6 +1681,7 @@ struct multiply_t<grotto::fixedpoint<FractionalBits, IntegralType>, simde__m256i
return multiply_t<IntegralType, simde__m256i>{}(a, b.integral_representation()); return multiply_t<IntegralType, simde__m256i>{}(a, b.integral_representation());
} }
}; };
HEDLEY_PRAGMA(GCC diagnostic pop)
} // namespace dpf::leaf_arithmetic } // namespace dpf::leaf_arithmetic

View file

@ -34,6 +34,12 @@ namespace grotto
/// zero-extend. Sign-extending a narrower signed operand, and replicating the /// zero-extend. Sign-extending a narrower signed operand, and replicating the
/// product sign when `modulus_bits > multiply_bits`, are plaintext steps the /// product sign when `modulus_bits > multiply_bits`, are plaintext steps the
/// MPC protocol has to reproduce (they are not local on additive shares). /// 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 <unsigned IntegerBits, template <unsigned IntegerBits,
unsigned FractionalBits, unsigned FractionalBits,
unsigned LhsFractionalBits, unsigned LhsFractionalBits,
@ -51,12 +57,12 @@ struct fixed_mul_plan
static constexpr bool rhs_signed = std::is_signed_v<RhsIntegral>; static constexpr bool rhs_signed = std::is_signed_v<RhsIntegral>;
static constexpr bool operands_signed = lhs_signed || rhs_signed; static constexpr bool operands_signed = lhs_signed || rhs_signed;
/// Right shift applied to the raw product. Negative means a left shift. /// @brief Right shift applied to the raw product. Negative means a left shift.
static constexpr int align_shift = static_cast<int>(LhsFractionalBits) static constexpr int align_shift = static_cast<int>(LhsFractionalBits)
+ static_cast<int>(RhsFractionalBits) + static_cast<int>(RhsFractionalBits)
- static_cast<int>(FractionalBits); - static_cast<int>(FractionalBits);
/// Bits of the product that the shift reads. Zero when a left shift /// @brief Bits of the product that the shift reads. Zero when a left shift
/// moves every product bit out of the output. /// moves every product bit out of the output.
static constexpr int modulus_bits_signed = align_shift >= 0 static constexpr int modulus_bits_signed = align_shift >= 0
? align_shift + static_cast<int>(out_bits) ? align_shift + static_cast<int>(out_bits)
@ -64,14 +70,14 @@ struct fixed_mul_plan
static constexpr unsigned modulus_bits = modulus_bits_signed > 0 static constexpr unsigned modulus_bits = modulus_bits_signed > 0
? static_cast<unsigned>(modulus_bits_signed) : 0u; ? static_cast<unsigned>(modulus_bits_signed) : 0u;
/// Full two's-complement product fits in this many bits. /// @brief Full two's-complement product fits in this many bits.
static constexpr unsigned product_bits = lhs_width + rhs_width; static constexpr unsigned product_bits = lhs_width + rhs_width;
static constexpr unsigned multiply_bits = modulus_bits < product_bits static constexpr unsigned multiply_bits = modulus_bits < product_bits
? modulus_bits : product_bits; ? modulus_bits : product_bits;
static constexpr unsigned limbs = multiply_bits == 0u static constexpr unsigned limbs = multiply_bits == 0u
? 0u : (multiply_bits + 63u) / 64u; ? 0u : (multiply_bits + 63u) / 64u;
/// Signed storage exists through 128 bits. A wider window is the same /// @brief Signed storage exists through 128 bits. A wider window is the same
/// residue held in an unsigned fixed-point. /// residue held in an unsigned fixed-point.
static constexpr bool result_is_signed = operands_signed && out_bits <= 128u; static constexpr bool result_is_signed = operands_signed && out_bits <= 128u;
@ -197,7 +203,14 @@ constexpr void store_raw_limbs(const T & value, std::uint64_t out[4]) noexcept
} }
} }
/// Low `dest_bits` of `value`, sign-extended when `value` is a narrower signed integer. /// @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 <typename T> template <typename T>
HEDLEY_NO_THROW HEDLEY_NO_THROW
constexpr void reduce_operand(const T & value, unsigned src_bits, bool is_signed, constexpr void reduce_operand(const T & value, unsigned src_bits, bool is_signed,
@ -222,7 +235,11 @@ constexpr void reduce_operand(const T & value, unsigned src_bits, bool is_signed
mask_to_bits(dest, nlimbs, dest_bits); mask_to_bits(dest, nlimbs, dest_bits);
} }
/// Product modulo `2^(64*nlimbs)`, using exactly `nlimbs` limbs of each operand. /// @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 HEDLEY_NO_THROW
constexpr void mul_low_limbs(std::uint64_t * out, const std::uint64_t * lhs, constexpr void mul_low_limbs(std::uint64_t * out, const std::uint64_t * lhs,
const std::uint64_t * rhs, unsigned nlimbs) noexcept const std::uint64_t * rhs, unsigned nlimbs) noexcept
@ -328,15 +345,21 @@ constexpr T limbs_to_integral(const std::uint64_t * limbs) noexcept
} // namespace detail } // namespace detail
/// @brief Multiply two fixed-point values into a chosen integer and fraction width. /// @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 /// @tparam IntegerBits Integer bits kept in the result, including the sign bit
/// when the result is signed. Bits above this wrap. /// when the result is signed. Bits above this wrap.
/// @tparam FractionalBits Fraction bits kept in the result. Lower fraction bits /// @tparam FractionalBits Fraction bits kept in the result. Lower fraction bits
/// of the exact product are discarded (floored). /// of the exact product are discarded (floored).
/// /// @tparam LhsFractionalBits fractional bits of the left operand
/// The result is held in the smallest fixed-point word that can store /// @tparam LhsIntegral integral type of the left operand
/// `IntegerBits + FractionalBits`. A signed word is used when either operand /// @tparam RhsFractionalBits fractional bits of the right operand
/// is signed and the window is at most 128 bits; otherwise the window is the /// @tparam RhsIntegral integral type of the right operand
/// unsigned residue. /// @param lhs the left-hand operand
/// @param rhs the right-hand operand
/// @return the product at the requested width
template <unsigned IntegerBits, template <unsigned IntegerBits,
unsigned FractionalBits, unsigned FractionalBits,
unsigned LhsFractionalBits, unsigned LhsFractionalBits,

View file

@ -1,7 +1,6 @@
/// @file grotto/gadget_hints.hpp /// @file grotto/gadget_hints.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Domain, degree, and pole hints for cleartext gadget references.
/// @details
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -10,6 +9,8 @@
#ifndef LIBDPF_INCLUDE_GROTTO_GADGET_HINTS_HPP__ #ifndef LIBDPF_INCLUDE_GROTTO_GADGET_HINTS_HPP__
#define LIBDPF_INCLUDE_GROTTO_GADGET_HINTS_HPP__ #define LIBDPF_INCLUDE_GROTTO_GADGET_HINTS_HPP__
#include "hedley/hedley.h"
#include <cmath> #include <cmath>
#include <array> #include <array>
#include <limits> #include <limits>

View file

@ -1,7 +1,6 @@
/// @file grotto/gadgets.hpp /// @file grotto/gadgets.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Umbrella include for the gadget reference headers that are enabled.
/// @details
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -9,6 +8,10 @@
#ifndef LIBDPF_INCLUDE_GROTTO_GADGETS_HPP__ #ifndef LIBDPF_INCLUDE_GROTTO_GADGETS_HPP__
#define LIBDPF_INCLUDE_GROTTO_GADGETS_HPP__ #define LIBDPF_INCLUDE_GROTTO_GADGETS_HPP__
// Cleartext gadget functors are the 2019 reference layer. Fixed-point
// evaluation lives in the LUTs: `eval_reduced` (including `expm1` and
// `log1p`), `eval_window`, `make_*_lut`, and `exact_constant`. A functor
// whose map those tables implement is deprecated.
// #include "grotto/gadgets/activations.hpp" // #include "grotto/gadgets/activations.hpp"
// #include "grotto/gadgets/binary.hpp" // #include "grotto/gadgets/binary.hpp"
#include "grotto/gadgets/decimal.hpp" #include "grotto/gadgets/decimal.hpp"

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations.hpp /// @file grotto/gadgets/activations.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Includes the activations gadget references.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/celu.hpp /// @file grotto/gadgets/activations/celu.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `celu`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/elish.hpp /// @file grotto/gadgets/activations/elish.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `elish`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -21,7 +21,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct elish HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) elish
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -44,6 +48,8 @@ struct gadget_hints<elish>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/elu.hpp /// @file grotto/gadgets/activations/elu.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `elu`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/gelu.hpp /// @file grotto/gadgets/activations/gelu.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `gelu`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,7 +20,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct gelu HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) gelu
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -42,6 +46,8 @@ struct gadget_hints<gelu>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/hardelish.hpp /// @file grotto/gadgets/activations/hardelish.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `hardelish`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/hardshrink.hpp /// @file grotto/gadgets/activations/hardshrink.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `hardshrink`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,9 +20,13 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
static constexpr double hardshrink_default_lambda = 0.5; static constexpr double hardshrink_default_lambda = 0.5;
template <double const & lambda = hardshrink_default_lambda> template <double const & lambda = hardshrink_default_lambda>
struct hardshrink struct HEDLEY_DEPRECATED_FOR(2026, grotto::make_hardshrink_lut) hardshrink
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -45,6 +49,8 @@ struct gadget_hints<hardshrink<lambda>>
static constexpr std::array<double, degree+1> canonical_polys[] = { {0,1}, {0}, {0,1} }; static constexpr std::array<double, degree+1> canonical_polys[] = { {0,1}, {0}, {0,1} };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/hardsigmoid.hpp /// @file grotto/gadgets/activations/hardsigmoid.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `hardsigmoid`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,7 +20,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct hardsigmoid HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::make_hardsigmoid_lut) hardsigmoid
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -44,6 +48,8 @@ struct gadget_hints<hardsigmoid>
static constexpr std::array<double, degree+1> canonical_polys[] = { {0}, {0.5,1/6.0}, {1} }; static constexpr std::array<double, degree+1> canonical_polys[] = { {0}, {0.5,1/6.0}, {1} };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/hardswish.hpp /// @file grotto/gadgets/activations/hardswish.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `hardswish`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,7 +20,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct hardswish HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::make_hardswish_lut) hardswish
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -44,6 +48,8 @@ struct gadget_hints<hardswish>
static constexpr std::array<double, degree+1> canonical_polys[] = { {0}, {0,0.5,1/6.0}, {0,1,0} }; static constexpr std::array<double, degree+1> canonical_polys[] = { {0}, {0,0.5,1/6.0}, {0,1,0} };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/hardtanh.hpp /// @file grotto/gadgets/activations/hardtanh.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `hardtanh`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,7 +20,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct hardtanh HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::make_hardtanh_lut) hardtanh
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -44,6 +48,8 @@ struct gadget_hints<hardtanh>
static constexpr std::array<double, degree+1> canonical_polys[] = { {-1}, {0,1}, {1} }; static constexpr std::array<double, degree+1> canonical_polys[] = { {-1}, {0,1}, {1} };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/leakyrelu.hpp /// @file grotto/gadgets/activations/leakyrelu.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `leakyrelu`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,10 +20,14 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
static constexpr double leakyrelu_default_negative_slope = 0.01; static constexpr double leakyrelu_default_negative_slope = 0.01;
static constexpr double leakyrelu_zero_negative_slope = 0.0; static constexpr double leakyrelu_zero_negative_slope = 0.0;
template <double const & negative_slope = leakyrelu_default_negative_slope> template <double const & negative_slope = leakyrelu_default_negative_slope>
struct leakyrelu struct HEDLEY_DEPRECATED_FOR(2026, grotto::make_leaky_relu_lut) leakyrelu
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -46,6 +50,8 @@ struct gadget_hints<leakyrelu<negative_slope>>
inline static constexpr std::array<double, degree+1> canonical_polys[] = { {0,negative_slope}, {0,1} }; inline static constexpr std::array<double, degree+1> canonical_polys[] = { {0,negative_slope}, {0,1} };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/lecun_tanh.hpp /// @file grotto/gadgets/activations/lecun_tanh.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `lecun_tanh`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/logsigmoid.hpp /// @file grotto/gadgets/activations/logsigmoid.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `logsigmoid`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -21,7 +21,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct logsigmoid HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) logsigmoid
{ {
template <typename T> template <typename T>
T operator()(T x) { return std::log(sigmoid{}(x)); } T operator()(T x) { return std::log(sigmoid{}(x)); }
@ -40,6 +44,8 @@ struct gadget_hints<logsigmoid>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/mish.hpp /// @file grotto/gadgets/activations/mish.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `mish`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -21,7 +21,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct mish HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) mish
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -43,6 +47,8 @@ struct gadget_hints<mish>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/one_minus_relu.hpp /// @file grotto/gadgets/activations/one_minus_sigmoid.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `one_minus_sigmoid`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -21,7 +21,10 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct one_minus_sigmoid HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) one_minus_sigmoid
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -43,6 +46,8 @@ struct gadget_hints<one_minus_sigmoid>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/relu.hpp /// @file grotto/gadgets/activations/relu.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `relu`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -21,7 +21,10 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
using relu = leakyrelu<leakyrelu_zero_negative_slope>; HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
using relu HEDLEY_DEPRECATED_FOR(2026, grotto::make_relu_lut) = leakyrelu<leakyrelu_zero_negative_slope>;
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/relu6.hpp /// @file grotto/gadgets/activations/relu6.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `relu6`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,9 +20,13 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
static constexpr double relu6_default_clip = 6; static constexpr double relu6_default_clip = 6;
template <double const & clip = 6> template <double const & clip = 6>
struct relu6 struct HEDLEY_DEPRECATED_FOR(2026, grotto::make_relu6_lut) relu6
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -44,6 +48,8 @@ struct gadget_hints<relu6<clip>>
static constexpr std::array<double, degree+1> canonical_polys[] = { 0, {0,1}, clip }; static constexpr std::array<double, degree+1> canonical_polys[] = { 0, {0,1}, clip };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/selu.hpp /// @file grotto/gadgets/activations/selu.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `selu`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/serf.hpp /// @file grotto/gadgets/activations/serf.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `serf`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -21,7 +21,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct serf HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) serf
{ {
template <typename T> template <typename T>
T operator()(T x) T operator()(T x)
@ -43,6 +47,8 @@ struct gadget_hints<serf>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/sigmoid.hpp /// @file grotto/gadgets/activations/sigmoid.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `sigmoid`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,7 +20,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct sigmoid HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) sigmoid
{ {
template <typename T> template <typename T>
T operator()(T x) { return 1/(1+std::exp(-x)); } T operator()(T x) { return 1/(1+std::exp(-x)); }
@ -39,6 +43,8 @@ struct gadget_hints<sigmoid>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

View file

@ -1,6 +1,6 @@
/// @file grotto/gadgets/activations/silu.hpp /// @file grotto/gadgets/activations/silu.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca> /// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief /// @brief Cleartext reference for `silu`.
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others /// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license; /// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details. /// see [LICENSE.md](@ref GPLv2) for details.
@ -20,7 +20,11 @@ namespace grotto
namespace gadgets namespace gadgets
{ {
struct silu HEDLEY_DIAGNOSTIC_PUSH
HEDLEY_DIAGNOSTIC_DISABLE_DEPRECATED
struct HEDLEY_DEPRECATED_FOR(2026, grotto::eval_window) silu
{ {
template <typename T> template <typename T>
T operator()(T x) { return x*sigmoid{}(x); } T operator()(T x) { return x*sigmoid{}(x); }
@ -39,6 +43,8 @@ struct gadget_hints<silu>
static constexpr std::array<double, degree+1> canonical_polys[] = { }; static constexpr std::array<double, degree+1> canonical_polys[] = { };
}; };
HEDLEY_DIAGNOSTIC_POP
} // namespace gadgets } // namespace gadgets
} // namespace grotto } // namespace grotto

Some files were not shown because too many files have changed in this diff Show more