libdpf/include/dpf/incremental.hpp

3643 lines
148 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

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

/// @file dpf/incremental.hpp
/// @brief Prefix-placed (incremental) DPF outputs via `dpf::at<N>`.
/// @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_INCREMENTAL_HPP__
#define LIBDPF_INCLUDE_DPF_INCREMENTAL_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <array>
#include <tuple>
#include <type_traits>
#include <utility>
#include <algorithm>
#include <optional>
#include <functional>
#include "hedley/hedley.h"
#include "portable-snippets/exact-int/exact-int.h"
#include "dpf/dpf_key.hpp"
#include "dpf/doerner_shelat.hpp"
#include "dpf/eval_target.hpp"
#include "dpf/eval_common.hpp"
#include "dpf/utils.hpp"
#include "dpf/twiddle.hpp"
#include "dpf/wildcard.hpp"
#include "dpf/leaf_node.hpp"
#include "dpf/leaf_wrapper.hpp"
#include "dpf/offset_wrapper.hpp"
#include "dpf/path_memoizer.hpp"
#include "dpf/output_buffer.hpp"
#include "dpf/interval_memoizer.hpp"
#include "dpf/subinterval_iterable.hpp"
#include "dpf/subsequence_iterable.hpp"
#include "dpf/dcf.hpp"
#include "dpf/blocked_dcf.hpp"
namespace dpf
{
// `at_pack` / `at` and the placement / slot-meta machinery now live in
// `dpf/placement.hpp` (included via `dpf/dpf_key.hpp`).
template <typename ...Args>
inline constexpr bool args_have_cmp_v =
(is_cmp_spec_v<std::decay_t<Args>> || ...);
template <typename ...Args>
inline constexpr bool args_have_verifiable_v =
(is_verifiable_tag_v<std::decay_t<Args>> || ...);
template <typename ...Args>
inline constexpr bool args_have_extractable_v =
(is_extractable_tag_v<std::decay_t<Args>> || ...);
template <typename ...Args>
inline constexpr bool args_have_eq_v =
(is_eq_spec_v<std::decay_t<Args>> || ...);
template <std::size_t BitLen>
HEDLEY_NO_THROW
constexpr std::size_t forced_cmp_depth_sum() noexcept { return 0; }
template <std::size_t BitLen, typename A0, typename ...Rest>
HEDLEY_NO_THROW
constexpr std::size_t forced_cmp_depth_sum() noexcept
{
std::size_t m = 0;
using A = std::decay_t<A0>;
if constexpr (is_cmp_spec_v<A>)
m = (A::prefix == 0) ? BitLen : A::prefix;
return std::max(m, forced_cmp_depth_sum<BitLen, Rest...>());
}
template <std::size_t BitLen, typename ...Args>
inline constexpr std::size_t forced_cmp_depth_v =
forced_cmp_depth_sum<BitLen, Args...>();
template <std::size_t BitLen>
HEDLEY_NO_THROW
constexpr std::size_t forced_cmp_out_bits_sum() noexcept { return 0; }
template <std::size_t BitLen, typename A0, typename ...Rest>
HEDLEY_NO_THROW
constexpr std::size_t forced_cmp_out_bits_sum() noexcept
{
std::size_t m = 0;
using A = std::decay_t<A0>;
if constexpr (is_cmp_spec_v<A>)
{
using Beta = typename A::beta_type;
m = std::is_same_v<Beta, dpf::bit> ? std::size_t{1}
: utils::bitlength_of_v<Beta>;
}
return std::max(m, forced_cmp_out_bits_sum<BitLen, Rest...>());
}
/// @brief Comparison output group width (bits) forced by any `lt`/`leq`/`gt`/`geq`
/// (or `_at`) spec in the pack — 0 when there is no comparison channel.
template <std::size_t BitLen, typename ...Args>
inline constexpr std::size_t forced_cmp_out_bits_v =
forced_cmp_out_bits_sum<BitLen, Args...>();
/// @brief True when the (single) comparison spec in the pack carries a wildcard
/// payload (`lt(dpf::wildcard<Beta>, ...)` etc.) — the δ is assigned after
/// keygen via `dpf::assign_cmp`.
/// @tparam A alignment of the rebound allocator
template <typename A, bool = is_cmp_spec_v<std::decay_t<A>>>
struct arg_is_wild_cmp : std::false_type {};
template <typename A>
struct arg_is_wild_cmp<A, true>
: std::bool_constant<
dpf::is_wildcard_v<typename std::decay_t<A>::beta_type>> {};
template <typename ...Args>
inline constexpr bool forced_cmp_wild_v =
(arg_is_wild_cmp<Args>::value || ...);
template <typename A, bool = is_cmp_spec_v<std::decay_t<A>>>
struct arg_cmp_block : std::integral_constant<std::size_t, 0> {};
template <typename A>
struct arg_cmp_block<A, true>
: std::integral_constant<std::size_t, std::decay_t<A>::block_width> {};
template <typename ...Args>
inline constexpr std::size_t forced_cmp_block_v =
(std::size_t{0} + ... + arg_cmp_block<Args>::value);
template <typename A, typename = void>
struct arg_cmp_idcf : std::false_type {};
template <typename A>
struct arg_cmp_idcf<A, std::enable_if_t<is_cmp_spec_v<std::decay_t<A>>>>
: spec_is_incremental<std::decay_t<A>> {};
template <typename ...Args>
inline constexpr bool forced_cmp_idcf_v =
(false || ... || arg_cmp_idcf<Args>::value);
template <typename T, typename = void>
struct spec_has_paint_fn : std::false_type {};
template <typename T>
struct spec_has_paint_fn<T, std::void_t<decltype(std::declval<T &>().fn)>>
: std::true_type {};
inline uint64_t paint_fn_adapter(std::size_t matched, uint64_t prefix, bool leaf,
const void * ctx)
{
using fn_type = std::function<uint64_t(std::size_t, uint64_t, bool)>;
return (*static_cast<const fn_type *>(ctx))(matched, prefix, leaf);
}
/// @brief Runtime description of one comparison channel peeled from `make_dpf` args.
struct dcf_runtime_spec
{
std::size_t prefix = 0; // 0 => full input bitlength
uint64_t beta = 0; // if_true - if_false
uint64_t false_value = 0;
uint64_t mask = ~0ULL;
/// @brief Payload group wider than the masked `uint64_t` ring, or a product
/// (`dpf::vec`) / XOR group. `beta_g` is δ and `false_g` is `if_false`.
bool custom = false;
detail::group_elem beta_g{};
detail::group_elem false_g{};
cmp_kind kind = cmp_kind::lt;
bool is_wildcard = false; // payload assigned after keygen (δ unknown now)
std::size_t length_bits = 0;
bool incremental = false;
std::function<uint64_t(std::size_t, uint64_t, bool)> paint;
};
namespace detail
{
namespace incr
{
// `placed`, `slot_meta`, `build_meta`, `build_group_order`, `lane_input`,
// `filter_group`, `meta_holder`, and the `out_bits_v` / `lg_opl_v` /
// `level_of_v` traits now live in `dpf/placement.hpp`.
template <std::size_t ...Levels, typename ...Betas>
auto expand_idpf(idpf_pack<std::index_sequence<Levels...>, Betas...> pack)
{
return std::apply([](auto && ...ys) {
return std::make_tuple(
placed<Levels, std::decay_t<decltype(ys)>>{
std::forward<decltype(ys)>(ys)}...);
}, std::move(pack.values));
}
template <std::size_t BitLen, typename Arg>
auto flatten_one(Arg && arg)
{
using A = std::decay_t<Arg>;
if constexpr (is_cmp_spec_v<A>)
{
return std::tuple<>{};
}
else if constexpr (is_verifiable_tag_v<A> || is_extractable_tag_v<A>)
{
return std::tuple<>{};
}
else if constexpr (is_idpf_v<A>)
{
return expand_idpf(std::forward<Arg>(arg));
}
else if constexpr (is_eq_spec_v<A>)
{
constexpr auto pref = A::prefix == 0 ? BitLen : A::prefix;
static_assert(pref <= BitLen, "eq_at<N> exceeds input bitlength");
using Beta = typename A::beta_type;
Beta delta = detail::dcf_impl::sub_beta(arg.if_true, arg.if_false);
return std::make_tuple(placed<pref, Beta>{std::move(delta), arg.if_false});
}
else if constexpr (is_at_v<A>)
{
static_assert(A::prefix <= BitLen, "at<N> exceeds input bitlength");
return std::apply([](auto && ...ys) {
return std::make_tuple(
placed<A::prefix, std::decay_t<decltype(ys)>>{
std::forward<decltype(ys)>(ys)}...);
}, std::forward<Arg>(arg).values);
}
else
{
return std::make_tuple(placed<BitLen, A>{std::forward<Arg>(arg)});
}
}
template <typename Arg>
void collect_cmp_one(dcf_runtime_spec & spec, bool & found, Arg && arg)
{
using A = std::decay_t<Arg>;
if constexpr (is_cmp_spec_v<A>)
{
if (found)
throw std::invalid_argument(
"at most one lt/leq/gt/geq (or *_at) per key");
found = true;
if constexpr (A::prefix != 0)
spec.prefix = A::prefix;
else
spec.prefix = 0;
spec.kind = A::kind;
spec.length_bits = spec_length_bits<A>::value;
spec.incremental = spec_is_incremental<A>::value;
if constexpr (spec_has_paint_fn<A>::value)
spec.paint = arg.fn;
using Beta = typename A::beta_type;
using Concrete = dpf::concrete_type_t<Beta>;
constexpr auto bits = [] {
if constexpr (std::is_same_v<Concrete, dpf::bit>)
return std::size_t{1};
else
return utils::bitlength_of_v<Concrete>;
}();
spec.mask = detail::dcf_impl::default_mask_for_bits(bits);
if constexpr (detail::cmp_group_info<Beta>::custom)
{
spec.custom = true;
const auto layout = detail::group_layout<Concrete>();
if constexpr (dpf::is_wildcard_v<Beta>)
{
spec.is_wildcard = true;
spec.beta_g = detail::group_zero(layout);
spec.false_g = detail::group_zero(layout);
}
else
{
spec.beta_g = detail::group_sub(
detail::group_from_beta(arg.if_true),
detail::group_from_beta(arg.if_false));
spec.false_g = detail::group_from_beta(arg.if_false);
}
spec.beta = 0;
spec.false_value = 0;
}
else if constexpr (dpf::is_wildcard_v<Beta>)
{
// Payload is unknown at keygen: keep δ = if_false = 0 (open the
// value CWs for the trivial payload) and record that this key's
// comparison must be `assign_cmp`ed before it can be evaluated.
spec.is_wildcard = true;
spec.beta = 0;
spec.false_value = 0;
}
else
{
spec.beta = detail::dcf_impl::beta_delta_u64(
arg.if_true, arg.if_false, spec.mask);
spec.false_value = detail::dcf_impl::beta_to_u64_simple(
arg.if_false, spec.mask);
}
}
}
template <std::size_t BitLen, typename ...Args>
auto flatten_args_and_cmp(dcf_runtime_spec & spec, bool & has_cmp, Args && ...args)
{
has_cmp = false;
(collect_cmp_one(spec, has_cmp, args), ...);
auto placed = std::tuple_cat(flatten_one<BitLen>(std::forward<Args>(args))...);
if (has_cmp && spec.prefix == 0)
spec.prefix = BitLen;
return placed;
}
template <std::size_t BitLen, typename ...Args>
auto flatten_args(Args && ...args)
{
dcf_runtime_spec spec{};
bool has_cmp = false;
return flatten_args_and_cmp<BitLen>(spec, has_cmp, std::forward<Args>(args)...);
}
template <typename NodeT, typename PlacedTuple, std::size_t BitLen, std::size_t ...Is>
constexpr bool classic_pack_impl(std::index_sequence<Is...>)
{
using first = std::tuple_element_t<0, PlacedTuple>;
constexpr auto w0 = out_bits_v<NodeT, typename first::output_type>;
return ((std::tuple_element_t<Is, PlacedTuple>::prefix == BitLen) && ...)
&& ((out_bits_v<NodeT,
typename std::tuple_element_t<Is, PlacedTuple>::output_type> == w0) && ...);
}
template <typename NodeT, typename PlacedTuple, std::size_t BitLen>
constexpr bool is_classic_placed()
{
constexpr auto n = std::tuple_size_v<PlacedTuple>;
if constexpr (n == 0) return false;
else return classic_pack_impl<NodeT, PlacedTuple, BitLen>(
std::make_index_sequence<n>{});
}
template <typename PlacedTuple, std::size_t ...Is>
auto classic_values(PlacedTuple & t, std::index_sequence<Is...>)
{
return std::make_tuple(std::move(std::get<Is>(t).value)...);
}
template <typename A, typename B>
void assign_tuple_element(A & a, B && b)
{
a = std::forward<B>(b);
}
template <std::size_t...Slots, std::size_t...Locals, typename OutL, typename OutB,
typename InL, typename InB>
void zip_install(OutL & out_leaves, OutB & out_beavers,
InL && in_leaves, InB && in_beavers,
std::index_sequence<Slots...>, std::index_sequence<Locals...>)
{
(assign_tuple_element(std::get<Slots>(out_leaves), std::get<Locals>(in_leaves)), ...);
(assign_tuple_element(std::get<Slots>(out_beavers), std::get<Locals>(in_beavers)), ...);
}
template <typename Placed>
auto concrete_placed_value(Placed & p)
{
using stored = typename Placed::stored_type;
if constexpr (dpf::is_arith_beta_v<stored>)
{
using T = typename Placed::output_type;
if constexpr (utils::has_characteristic_two_v<T>)
return static_cast<T>(p.value.y0 ^ p.value.y1);
else
return static_cast<T>(p.value.y0 + p.value.y1);
}
else
{
return p.value;
}
}
template <typename PlacedTuple, std::size_t... Is>
constexpr bool placed_has_arith_beta_impl(std::index_sequence<Is...>)
{
return (dpf::is_arith_beta_v<
typename std::tuple_element_t<Is, PlacedTuple>::stored_type>
|| ...);
}
template <typename PlacedTuple>
constexpr bool placed_has_arith_beta()
{
constexpr auto n = std::tuple_size_v<PlacedTuple>;
if constexpr (n == 0)
return false;
else
return placed_has_arith_beta_impl<PlacedTuple>(
std::make_index_sequence<n>{});
}
template <typename PlacedTuple, std::size_t... GlobalIs>
using group_outputs_t = std::tuple<
typename std::tuple_element_t<GlobalIs, PlacedTuple>::output_type...>;
template <typename ExteriorPRG, typename PlacedTuple, typename CwProtocol,
typename SeedT, typename Leaves0T, typename Leaves1T,
std::size_t LocalI, std::size_t GlobalI, std::size_t... AllGlobals>
void overwrite_arith_slot_one(CwProtocol & proto, const SeedT & s0,
const SeedT & s1, uint8_t t0, uint8_t t1, std::size_t pos_base,
std::size_t lane, PlacedTuple & placed, Leaves0T & leaves0,
Leaves1T & leaves1, std::index_sequence<AllGlobals...>)
{
using P = std::tuple_element_t<GlobalI, PlacedTuple>;
if constexpr (dpf::is_arith_beta_v<typename P::stored_type>)
{
using outs = group_outputs_t<PlacedTuple, AllGlobals...>;
auto & slot = std::get<GlobalI>(placed);
auto cw = proto.template open_arith_leaf<ExteriorPRG, LocalI, outs>(
s0, s1, t0, t1, slot.value.y0, slot.value.y1, pos_base, lane);
std::get<GlobalI>(leaves0) = cw;
std::get<GlobalI>(leaves1) = cw;
}
}
template <typename ExteriorPRG, typename PlacedTuple, typename CwProtocol,
typename SeedT, typename Leaves0T, typename Leaves1T,
std::size_t... GlobalIs, std::size_t... LocalIs>
void overwrite_arith_slots(CwProtocol & proto, const SeedT & s0,
const SeedT & s1, uint8_t t0, uint8_t t1, std::size_t pos_base,
std::size_t lane, PlacedTuple & placed, Leaves0T & leaves0,
Leaves1T & leaves1, std::index_sequence<GlobalIs...>,
std::index_sequence<LocalIs...>)
{
(overwrite_arith_slot_one<ExteriorPRG, PlacedTuple, CwProtocol, SeedT,
Leaves0T, Leaves1T, LocalIs, GlobalIs>(proto, s0, s1, t0, t1,
pos_base, lane, placed, leaves0, leaves1,
std::index_sequence<GlobalIs...>{}),
...);
}
template <typename ExteriorPRG, typename InputT, typename SeedT,
typename PlacedTuple, std::size_t... Is>
auto call_make_leaves(std::size_t pos_base, InputT lane_x, const SeedT & s0,
const SeedT & s1, bool sign, PlacedTuple & placed,
std::index_sequence<Is...>)
{
return dpf::make_leaves<ExteriorPRG>(lane_x, s0, s1, sign, pos_base,
concrete_placed_value(std::get<Is>(placed))...);
}
template <typename ExteriorPRG, typename PlacedTuple, std::size_t... Is>
auto empty_leaves(std::index_sequence<Is...>)
{
using node = typename ExteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::make_tuple(dpf::leaf_node_t<node,
typename std::tuple_element_t<Is, PlacedTuple>::output_type>{}...);
HEDLEY_PRAGMA(GCC diagnostic pop)
}
template <typename ExteriorPRG, typename PlacedTuple, std::size_t... Is>
auto empty_beavers(std::index_sequence<Is...>)
{
using node = typename ExteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::make_tuple(dpf::beaver<
dpf::is_wildcard_v<
typename std::tuple_element_t<Is, PlacedTuple>::output_type>,
node,
concrete_type_t<
typename std::tuple_element_t<Is, PlacedTuple>::output_type>>{}...);
HEDLEY_PRAGMA(GCC diagnostic pop)
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, typename MetaHolder, std::size_t G,
typename Leaves0T, typename Beavers0T, typename Leaves1T,
typename Beavers1T, typename CwProtocol = void>
void gen_group(InputT x, const typename InteriorPRG::block_type & s0,
const typename InteriorPRG::block_type & s1, uint8_t t0, uint8_t t1,
PlacedTuple & placed, Leaves0T & leaves0, Beavers0T & beavers0,
Leaves1T & leaves1, Beavers1T & beavers1, CwProtocol * proto = nullptr)
{
constexpr std::size_t n = std::tuple_size_v<PlacedTuple>;
constexpr std::size_t bitlen = utils::bitlength_of_v<InputT>;
using idxs = filter_group_t<G, MetaHolder, n>;
constexpr auto meta = MetaHolder::value;
std::size_t prefix = 0, pos_base = 0;
for (std::size_t i = 0; i < n; ++i)
{
if (meta[i].group_id == G)
{
prefix = meta[i].prefix;
pos_base = meta[i].pos_base;
break;
}
}
InputT lane_x = lane_input(x, prefix, bitlen);
const bool sign0 = static_cast<bool>(t0 & 1u);
auto built = call_make_leaves<ExteriorPRG>(pos_base, lane_x, s0, s1, sign0,
placed, idxs{});
constexpr auto nslots = idxs{}.size();
zip_install(leaves0, beavers0, built.first.first, built.first.second,
idxs{}, std::make_index_sequence<nslots>{});
zip_install(leaves1, beavers1, built.second.first, built.second.second,
idxs{}, std::make_index_sequence<nslots>{});
if constexpr (!std::is_void_v<CwProtocol>)
{
if (proto != nullptr)
{
constexpr auto to_int = utils::to_integral_type<InputT>{};
const std::size_t lane = static_cast<std::size_t>(to_int(lane_x));
overwrite_arith_slots<ExteriorPRG, PlacedTuple>(
*proto, s0, s1, t0, t1, pos_base, lane, placed, leaves0,
leaves1, idxs{}, std::make_index_sequence<nslots>{});
}
}
else
{
(void)t1;
(void)proto;
static_assert(!placed_has_arith_beta_impl<PlacedTuple>(idxs{}),
"arith_beta payloads require Doerner–Shelat gen (CwProtocol)");
}
}
template <typename F, std::size_t... Is>
void for_each_index(std::index_sequence<Is...>, F && f)
{
(static_cast<void>(f(std::integral_constant<std::size_t, Is>{})), ...);
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, typename MetaHolder, std::size_t... Gs,
typename Leaves0T, typename Beavers0T, typename Leaves1T,
typename Beavers1T>
void gen_all_groups(InputT x, const typename InteriorPRG::block_type & s0,
const typename InteriorPRG::block_type & s1, uint8_t t0, uint8_t t1,
PlacedTuple & placed, Leaves0T & leaves0, Beavers0T & beavers0,
Leaves1T & leaves1, Beavers1T & beavers1,
std::index_sequence<Gs...>)
{
(gen_group<InteriorPRG, ExteriorPRG, InputT, PlacedTuple, MetaHolder, Gs>(
x, s0, s1, t0, t1, placed, leaves0, beavers0, leaves1, beavers1),
...);
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, std::size_t... Is,
typename LeavesT, typename BeaversT>
auto wrap_leaves(LeavesT & leaves, BeaversT & beavers,
std::index_sequence<Is...>)
{
using node = typename ExteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return std::make_tuple(
dpf::leaf_wrapper<
typename std::tuple_element_t<Is, PlacedTuple>::output_type,
node>(std::get<Is>(leaves), std::get<Is>(beavers))...);
HEDLEY_PRAGMA(GCC diagnostic pop)
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, std::size_t CmpDepth = 0,
std::size_t CmpOutBits = 0, bool CmpWild = false,
std::size_t CmpBlock = 0, bool CmpIdcf = false,
bool IsVerifiable = false, bool IsExtractable = false>
auto make_incremental_impl(InputT x, PlacedTuple placed,
root_sampler_t<InteriorPRG> root_sampler,
const dcf_runtime_spec * cmp_spec = nullptr);
} // namespace incr
} // namespace detail
namespace detail
{
namespace incr
{
/// @brief Map kind → flags. Tree keep-path stays on α; leq/gt plant δ on that path
/// via `include_eq`. Domain-edge α yields trivial always-true/false.
/// @param ch the `ch`
/// @param thresh the `thresh`
/// @param nbits the width in bits
inline void adjust_cmp_threshold(detail::cmp_meta & ch,
unsigned __int128 & thresh, std::size_t nbits)
{
const unsigned __int128 domain =
(nbits >= 128) ? 0
: (((unsigned __int128)1 << nbits) - 1);
ch.include_eq = false;
ch.trivial = cmp_trivial::none;
if (is_paint_kind(ch.kind))
{
ch.eval_as_ge = false;
return;
}
switch (ch.kind)
{
case cmp_kind::lt:
ch.eval_as_ge = false;
break;
case cmp_kind::leq:
ch.eval_as_ge = false;
ch.include_eq = true;
if (nbits < 128 && thresh == domain)
ch.trivial = cmp_trivial::always_true;
break;
case cmp_kind::geq:
ch.eval_as_ge = true;
break;
case cmp_kind::gt:
ch.eval_as_ge = true;
ch.include_eq = true;
if (nbits < 128 && thresh == domain)
ch.trivial = cmp_trivial::always_false;
break;
}
if (ch.incremental)
ch.trivial = cmp_trivial::none;
}
/// @brief Split the constant absorb into party shares. `target` is `if_false` for
/// lt/leq, or `δ + if_false` when the path-sum is inverted (geq/gt), or the
/// full if_true / if_false for trivial domain edges. Clears nothing — caller
/// must not put δ on the key.
/// @param target the opened payload target
/// @param mask the bit mask
/// @param r the `r`
/// @param add0 the `add0`
/// @param add1 the `add1`
HEDLEY_NO_THROW
inline void split_cmp_addend(uint64_t target, uint64_t mask, uint64_t r,
uint64_t & add0, uint64_t & add1) noexcept
{
using namespace detail::dcf_impl;
r &= mask;
target &= mask;
add0 = r;
add1 = (target + neg_m(r, mask)) & mask;
}
/// @brief Typed overload: write party-0 / party-1 additive shares of the absorb.
/// @param target the opened payload target
/// @param mask the bit mask
/// @param r the `r`
/// @param add0 the `add0`
/// @param add1 the `add1`
HEDLEY_NO_THROW
inline void split_cmp_addend(uint64_t target, uint64_t mask, uint64_t r,
additive_share<uint64_t, 0> & add0,
additive_share<uint64_t, 1> & add1) noexcept
{
uint64_t a0 = 0, a1 = 0;
split_cmp_addend(target, mask, r, a0, a1);
add0 = additive_share<uint64_t, 0>::from_raw(a0);
add1 = additive_share<uint64_t, 1>::from_raw(a1);
}
template <typename PlacedTuple, std::size_t... Is>
auto extract_addends(const PlacedTuple & placed, std::index_sequence<Is...>)
{
return std::make_tuple(std::get<Is>(placed).addend...);
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, std::size_t CmpDepth, std::size_t CmpOutBits,
bool CmpWild, std::size_t CmpBlock, bool CmpIdcf,
bool IsVerifiable, bool IsExtractable>
auto make_incremental_impl(InputT x, PlacedTuple placed,
root_sampler_t<InteriorPRG> root_sampler,
const dcf_runtime_spec * cmp_spec)
{
static_assert(!(CmpIdcf && CmpBlock > 0),
"idcf uses the per-level path, not blocked checkpoints");
using key_type =
incr_dpf_key_of_t<InteriorPRG, ExteriorPRG, InputT, PlacedTuple, CmpDepth,
CmpOutBits, CmpWild, CmpBlock, CmpIdcf, IsVerifiable, IsExtractable>;
using interior_node = typename key_type::interior_node;
using input_type = typename key_type::input_type;
constexpr auto depth = key_type::depth;
constexpr auto n = key_type::num_outputs;
using MetaHolder = meta_holder<key_type>;
using namespace detail::dcf_impl;
if constexpr (IsExtractable)
{
static_assert(n > 0, "extractable requires at least one output");
[]<std::size_t... Is>(std::index_sequence<Is...>) {
static_assert((extractable_codomain_ok_v<
typename key_type::template concrete_output_type<Is>> && ...),
"extractable: each output must be fp61 or >= 128 bits");
}(std::make_index_sequence<n>{});
}
utils::flip_msb_if_signed_integral(x);
using tree = dpf::tree_traits<InteriorPRG>;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
interior_node root[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(root, root_sampler);
typename key_type::correction_words_array correction_words{};
typename key_type::correction_advice_array correction_advice{};
typename key_type::correction_seeds_array correction_seeds{};
interior_node parent[2] = {root[0], root[1]};
auto mask = key_type::msb_mask;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::array<interior_node, depth + 1> snap0{};
std::array<interior_node, depth + 1> snap1{};
HEDLEY_PRAGMA(GCC diagnostic pop)
std::array<bool, depth + 1> snap_sign{};
std::array<bool, depth + 1> need_snap{};
for (std::size_t i = 0; i <= depth; ++i)
need_snap[i] = false;
for (std::size_t i = 0; i < n; ++i)
need_snap[key_type::meta[i].tree_level] = true;
if (need_snap[0])
{
snap0[0] = parent[0];
snap1[0] = parent[1];
snap_sign[0] = dpf::get_lo_bit(parent[0]);
}
detail::cmp_meta cmp{};
// ... rest continues from existing body — patched via smaller edits below
typename key_type::value_cw_array value_cws{};
typename key_type::tail_array tail{};
typename key_type::tail_array tail_coeff{};
uint64_t cw_last = 0;
// Wildcard payload: per-level δ-coefficients (`value_cw(1) − value_cw(0)`),
// computed with a parallel β = 1 accumulator `Va1`. Unused when !CmpWild.
// Blocked keys store one coefficient per checkpoint and tail slot instead.
typename key_type::value_cw_array value_cw_coeff{};
uint64_t cw_last_coeff = 0;
uint64_t Va1 = 0;
unsigned __int128 thresh = 0;
std::size_t cmp_nbits = 0;
uint64_t Va = 0;
uint64_t delta = 0;
uint64_t false_value = 0;
if (cmp_spec != nullptr)
{
cmp_nbits = cmp_spec->prefix ? cmp_spec->prefix
: utils::bitlength_of_v<input_type>;
cmp.nbits = static_cast<int>(cmp_nbits);
cmp.mask = cmp_spec->mask;
delta = cmp_spec->beta & cmp_spec->mask;
false_value = cmp_spec->false_value & cmp_spec->mask;
cmp.kind = cmp_spec->kind;
cmp.active = true;
cmp.incremental = cmp_spec->incremental;
cmp.eval_as_ge = false;
cmp.trivial = cmp_trivial::none;
cmp.block_width = static_cast<int>(CmpBlock);
cmp.tail_bits = static_cast<int>(key_type::cmp_q);
auto lane = lane_input(x, cmp_nbits, utils::bitlength_of_v<input_type>);
thresh = static_cast<unsigned __int128>(
utils::to_integral_type<input_type>{}(lane));
adjust_cmp_threshold(cmp, thresh, cmp_nbits);
if (CmpBlock > 0 && is_paint_kind(cmp.kind))
throw std::invalid_argument(
"path recipes use the per-level comparison channel");
if (cmp_spec->custom && (CmpBlock > 0 || is_paint_kind(cmp.kind)))
throw std::invalid_argument(
"this comparison payload uses the per-level channel");
}
const bool custom_cmp = cmp_spec != nullptr && cmp_spec->custom;
detail::group_elem g_delta = custom_cmp
? cmp_spec->beta_g : detail::group_elem{};
detail::group_elem g_false = custom_cmp
? cmp_spec->false_g : detail::group_elem{};
detail::group_elem g_va = detail::group_zero(g_delta);
detail::group_elem g_va1 = detail::group_zero(g_delta);
detail::group_elem g_one = detail::group_one(g_delta);
detail::group_elem g_on = detail::group_zero(g_delta);
detail::group_elem g_on_unit = detail::group_zero(g_delta);
if (custom_cmp && cmp.include_eq && !is_paint_kind(cmp.kind))
{
g_on = g_delta;
g_on_unit = g_one;
}
typename key_type::value_cw_word g_last{};
typename key_type::value_cw_word g_last_coeff{};
const paint_callback paint_cb =
(cmp_spec != nullptr && cmp_spec->paint)
? &paint_fn_adapter : nullptr;
const void * paint_ctx =
(cmp_spec != nullptr && cmp_spec->paint) ? &cmp_spec->paint : nullptr;
const std::size_t paint_length_bits =
cmp_spec != nullptr ? cmp_spec->length_bits : 0;
typename key_type::prefix_cw_array prefix_cw{};
typename key_type::prefix_cw_array prefix_coeff{};
const auto on_path_scaled = [&](uint64_t scale) -> uint64_t {
if (!is_paint_kind(cmp.kind))
return (cmp.include_eq ? scale : 0ULL) & cmp.mask;
const uint64_t unit = detail::dcf_impl::paint_unit(cmp.kind, cmp_nbits,
thresh, cmp_nbits, paint_length_bits, true, paint_cb, paint_ctx);
return detail::dcf_impl::scale_plant(unit, scale, cmp.mask);
};
const uint64_t on_path = on_path_scaled(delta);
const uint64_t on_path_unit = on_path_scaled(1ULL);
auto snap_prefix = [&](std::size_t at) {
if constexpr (CmpIdcf)
{
if (!cmp.incremental || cmp.trivial != cmp_trivial::none)
return;
const uint64_t word = detail::dcf_impl::make_final_cw(parent[0],
parent[1], static_cast<uint8_t>(dpf::get_lo_bit(parent[1])),
Va, cmp.mask, on_path);
prefix_cw[at] = static_cast<typename key_type::value_cw_word>(word);
if constexpr (CmpWild)
{
const uint64_t w1 = detail::dcf_impl::make_final_cw(parent[0],
parent[1], static_cast<uint8_t>(dpf::get_lo_bit(parent[1])),
Va1, cmp.mask, on_path_unit);
prefix_coeff[at] = static_cast<typename key_type::value_cw_word>(
(w1 + detail::dcf_impl::neg_m(word, cmp.mask)) & cmp.mask);
}
}
};
auto snap_prefix_group = [&](std::size_t at) {
if constexpr (CmpIdcf)
{
if (!cmp.incremental || cmp.trivial != cmp_trivial::none)
return;
const auto word = detail::group_final_cw<InteriorPRG>(parent[0],
parent[1], static_cast<uint8_t>(dpf::get_lo_bit(parent[1])),
g_va, g_on);
prefix_cw[at] = detail::group_to_word<typename key_type::value_cw_word>(word);
if constexpr (CmpWild)
{
const auto w1 = detail::group_final_cw<InteriorPRG>(parent[0],
parent[1], static_cast<uint8_t>(dpf::get_lo_bit(parent[1])),
g_va1, g_on_unit);
prefix_coeff[at] = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(w1, word));
}
}
};
if constexpr (CmpIdcf)
{
if (cmp.active && !custom_cmp)
snap_prefix(0);
if (cmp.active && custom_cmp)
snap_prefix_group(0);
}
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{
bool bit = !!(mask & x);
const bool is_last = tree::is_last_level(level, depth);
const bool advice0 = static_cast<bool>(dpf::get_lo_bit(parent[0]));
const bool advice1 = static_cast<bool>(dpf::get_lo_bit(parent[1]));
auto child0 = tree::expand(parent[0], is_last);
auto child1 = tree::expand(parent[1], is_last);
// Value Convert stretch (HT mid: two-tweak, independent of seed expand).
const auto val0 = tree::expand_value(parent[0]);
const auto val1 = tree::expand_value(parent[1]);
interior_node cw{};
psnip_uint8_t tpack = 0;
tree::make_cw(cw, tpack, child0, child1, parent[0], parent[1], bit,
is_last);
parent[0] = tree::advance(parent[0], child0, cw, tpack, bit, advice0,
is_last);
parent[1] = tree::advance(parent[1], child1, cw, tpack, bit, advice1,
is_last);
correction_words[level] = cw;
correction_advice[level] = tpack;
if constexpr (IsVerifiable)
{
// Prefix bits of α through this level (same labelling as eval fold).
const auto prefix = static_cast<psnip_uint64_t>(
utils::to_integral_type<input_type>{}(x)
>> (utils::bitlength_of_v<input_type> - (level + 1)));
correction_seeds[level] = detail::vdpf::make_cs(level, prefix,
parent[0], parent[1]);
}
if constexpr (CmpBlock == 0)
{
if (cmp.active && cmp.trivial == cmp_trivial::none
&& level < cmp_nbits && !custom_cmp)
{
const int ai = static_cast<int>(
(thresh >> (cmp_nbits - 1 - level)) & 1);
uint64_t base = 0;
if (is_paint_kind(cmp.kind))
{
const uint64_t unit = detail::dcf_impl::paint_unit(cmp.kind,
level, thresh, cmp_nbits, paint_length_bits, false,
paint_cb, paint_ctx);
base = detail::dcf_impl::make_value_cw_planted(val0[0],
val0[1], val1[0], val1[1],
static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va,
detail::dcf_impl::scale_plant(unit, delta, cmp.mask),
cmp.mask);
if constexpr (CmpWild)
{
const uint64_t v1 = detail::dcf_impl::make_value_cw_planted(
val0[0], val0[1], val1[0], val1[1],
static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va1,
detail::dcf_impl::scale_plant(unit, 1ULL, cmp.mask),
cmp.mask);
value_cw_coeff[level] =
static_cast<typename key_type::value_cw_word>(
(v1 + detail::dcf_impl::neg_m(base, cmp.mask))
& cmp.mask);
}
}
else
{
base = make_value_cw(val0[0], val0[1],
val1[0], val1[1],
static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va, delta, cmp.mask);
if constexpr (CmpWild)
{
// Same recurrence with β = 1 on a parallel accumulator; the
// value CW is affine in β so `coeff = value_cw(1) − base`.
const uint64_t v1 = make_value_cw(val0[0], val0[1],
val1[0], val1[1],
static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va1, 1ULL, cmp.mask);
value_cw_coeff[level] =
static_cast<typename key_type::value_cw_word>(
(v1 + detail::dcf_impl::neg_m(base, cmp.mask))
& cmp.mask);
}
}
value_cws[level] =
static_cast<typename key_type::value_cw_word>(base);
snap_prefix(level + 1);
}
else if (cmp.active && cmp.trivial == cmp_trivial::none
&& level < cmp_nbits && custom_cmp)
{
const int ai = static_cast<int>(
(thresh >> (cmp_nbits - 1 - level)) & 1);
const auto base = detail::group_value_cw<InteriorPRG>(val0[0],
val0[1], val1[0], val1[1],
static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, g_va, g_delta);
value_cws[level] =
detail::group_to_word<typename key_type::value_cw_word>(base);
if constexpr (CmpWild)
{
const auto v1 = detail::group_value_cw<InteriorPRG>(val0[0],
val0[1], val1[0], val1[1],
static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, g_va1, g_one);
value_cw_coeff[level] =
detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(v1, base));
}
snap_prefix_group(level + 1);
}
}
else if (cmp.active && cmp.trivial == cmp_trivial::none)
{
using sched = detail::blocked::schedule<key_type::cmp_h, CmpBlock>;
const std::size_t c = level + 1;
if (c <= key_type::cmp_h && sched::contains(c))
{
const auto wi = sched::index(c);
const uint64_t word = detail::blocked::checkpoint_word<InteriorPRG>(
parent[0], parent[1], delta, cmp.mask);
value_cws[wi] =
static_cast<typename key_type::value_cw_word>(word);
if constexpr (CmpWild)
{
value_cw_coeff[wi] =
static_cast<typename key_type::value_cw_word>(
detail::blocked::checkpoint_coeff(
parent[0], parent[1], cmp.mask));
}
}
if (c == key_type::cmp_h && key_type::cmp_q > 0)
{
uint64_t words[4]{};
uint64_t coeffs[4]{};
const uint64_t suffix = static_cast<uint64_t>(thresh)
& ((1ULL << key_type::cmp_q) - 1ULL);
detail::blocked::tail_words<InteriorPRG>(parent[0], parent[1],
delta, cmp.mask, cmp.include_eq, suffix, key_type::cmp_q,
words, CmpWild ? coeffs : nullptr);
for (std::size_t z = 0; z < key_type::cmp_tail; ++z)
{
tail[z] = static_cast<typename key_type::value_cw_word>(words[z]);
if constexpr (CmpWild)
{
tail_coeff[z] =
static_cast<typename key_type::value_cw_word>(coeffs[z]);
}
}
}
}
if (need_snap[level + 1])
{
snap0[level + 1] = parent[0];
snap1[level + 1] = parent[1];
snap_sign[level + 1] = dpf::get_lo_bit(parent[0]);
}
if constexpr (CmpBlock == 0)
{
if (cmp.active && cmp.trivial == cmp_trivial::none
&& level + 1 == cmp_nbits && !custom_cmp)
{
cw_last = make_final_cw(parent[0], parent[1],
static_cast<uint8_t>(dpf::get_lo_bit(parent[1])), Va, cmp.mask,
on_path);
if constexpr (CmpWild)
{
const uint64_t l1 = make_final_cw(parent[0], parent[1],
static_cast<uint8_t>(dpf::get_lo_bit(parent[1])), Va1,
cmp.mask, on_path_unit);
cw_last_coeff =
(l1 + detail::dcf_impl::neg_m(cw_last, cmp.mask)) & cmp.mask;
}
}
else if (cmp.active && cmp.trivial == cmp_trivial::none
&& level + 1 == cmp_nbits && custom_cmp)
{
g_last = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_final_cw<InteriorPRG>(parent[0], parent[1],
static_cast<uint8_t>(dpf::get_lo_bit(parent[1])), g_va, g_on));
if constexpr (CmpWild)
{
const auto l1 = detail::group_final_cw<InteriorPRG>(parent[0],
parent[1], static_cast<uint8_t>(dpf::get_lo_bit(parent[1])),
g_va1, g_on_unit);
g_last_coeff = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(l1,
detail::group_from_word(g_last, g_delta)));
}
}
}
}
// `key_type::num_outputs` is a static constexpr, so the loop bound is a
// constant expression. A local `n` is not usable here: capturing it is not
// a constant, and the loop would lower to a goto.
constexpr std::size_t ngroups = [] {
std::size_t m = 0;
for (std::size_t i = 0; i < key_type::num_outputs; ++i)
m = std::max(m, key_type::meta[i].group_id + 1);
return m;
}();
auto leaves0 = empty_leaves<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
auto beavers0 = empty_beavers<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
auto leaves1 = empty_leaves<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
auto beavers1 = empty_beavers<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
constexpr auto order =
build_group_order<decltype(key_type::meta), ngroups>(key_type::meta, n);
using leaf_prg = std::conditional_t<IsExtractable,
detail::vdpf::extractable_leaf_prg<ExteriorPRG>, ExteriorPRG>;
for_each_index(std::make_index_sequence<ngroups>{}, [&](auto oi) {
constexpr std::size_t G = order[decltype(oi)::value];
constexpr std::size_t lvl = [] {
for (std::size_t i = 0; i < key_type::num_outputs; ++i)
{
if (key_type::meta[i].group_id == G)
return key_type::meta[i].tree_level;
}
return std::size_t{0};
}();
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
gen_group<InteriorPRG, leaf_prg, InputT, PlacedTuple, MetaHolder, G>(
x, dpf::unset_lo_2bits(snap0[lvl]), dpf::unset_lo_2bits(snap1[lvl]),
static_cast<uint8_t>(snap_sign[lvl]),
static_cast<uint8_t>(dpf::get_lo_bit(snap1[lvl])), placed, leaves0,
beavers0, leaves1, beavers1);
HEDLEY_PRAGMA(GCC diagnostic pop)
});
auto wrap0 = wrap_leaves<InteriorPRG, ExteriorPRG, InputT, PlacedTuple>(
leaves0, beavers0, std::make_index_sequence<n>{});
auto wrap1 = wrap_leaves<InteriorPRG, ExteriorPRG, InputT, PlacedTuple>(
leaves1, beavers1, std::make_index_sequence<n>{});
uint64_t cmp_add0 = 0, cmp_add1 = 0;
typename key_type::value_cw_word g_add0{};
typename key_type::value_cw_word g_add1{};
if (cmp.active && !custom_cmp)
{
uint64_t target = false_value;
if (cmp.trivial == cmp_trivial::always_true)
target = (delta + false_value) & cmp.mask;
else if (cmp.trivial == cmp_trivial::always_false)
target = false_value;
else if (cmp.eval_as_ge)
target = (delta + false_value) & cmp.mask;
// Group-width addend blind (same helper the DS path uses). The blind
// keeps only `popcount(mask)` live bits; the addend share is stored at
// `value_cw_word` width on the key.
const uint64_t r = detail::dcf_impl::sample_addend_blind(cmp.mask,
[&]() -> interior_node { return root_sampler(); });
split_cmp_addend(target, cmp.mask, r, cmp_add0, cmp_add1);
}
else if (cmp.active && custom_cmp)
{
detail::group_elem target = g_false;
if (cmp.trivial == cmp_trivial::always_true || cmp.eval_as_ge)
target = detail::group_add(g_delta, g_false);
const auto blind = detail::group_from_node<InteriorPRG>(root_sampler(), g_delta);
g_add0 = detail::group_to_word<typename key_type::value_cw_word>(blind);
g_add1 = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(target, blind));
}
input_type off0{}, off1{};
auto adds = extract_addends(placed, std::make_index_sequence<n>{});
key_type key0{root[0], correction_words, correction_advice, std::move(wrap0),
off0, cmp, value_cws, cw_last, cmp_add0, adds, value_cw_coeff,
cw_last_coeff, tail, tail_coeff, prefix_cw, prefix_coeff,
correction_seeds};
key_type key1{root[1], correction_words, correction_advice, std::move(wrap1),
off1, cmp, value_cws, cw_last, cmp_add1, adds, value_cw_coeff,
cw_last_coeff, tail, tail_coeff, prefix_cw, prefix_coeff,
correction_seeds};
if (custom_cmp)
{
key0.set_cmp_scalars(g_last, g_add0, g_last_coeff);
key1.set_cmp_scalars(g_last, g_add1, g_last_coeff);
}
return dpf::make_party_key_pair(std::move(key0), std::move(key1));
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename PlacedTuple, std::size_t CmpDepth, std::size_t CmpOutBits,
bool CmpWild, std::size_t CmpBlock, bool CmpIdcf,
bool IsVerifiable, bool IsExtractable,
typename RootSampler,
typename CwProtocol>
auto make_incremental_ds_impl(bool arith, InputT x0, InputT x1,
RootSampler & root_sampler, CwProtocol & proto, PlacedTuple placed,
const dcf_runtime_spec * cmp_spec)
{
static_assert(!dpf::is_wildcard_v<InputT>,
"Doerner–Shelat gen takes shares of a concrete point");
static_assert(sizeof(typename InteriorPRG::block_type) == sizeof(simde__m128i),
"Doerner–Shelat gen uses the AES-block interior node");
static_assert(!(CmpIdcf && CmpBlock > 0),
"idcf uses the per-level path, not blocked checkpoints");
using key_type =
incr_dpf_key_of_t<InteriorPRG, ExteriorPRG, InputT, PlacedTuple, CmpDepth,
CmpOutBits, CmpWild, CmpBlock, CmpIdcf, IsVerifiable, IsExtractable>;
using interior_node = typename key_type::interior_node;
using input_type = typename key_type::input_type;
constexpr auto depth = key_type::depth;
constexpr auto n = key_type::num_outputs;
using MetaHolder = meta_holder<key_type>;
using namespace detail::dcf_impl;
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")
interior_node roots[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(roots, [&]() -> interior_node {
return static_cast<interior_node>(root_sampler());
});
const interior_node root0 = roots[0];
const interior_node root1 = roots[1];
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_gen_state<interior_node> st;
HEDLEY_PRAGMA(GCC diagnostic pop)
st.init(root0, root1);
typename key_type::correction_words_array correction_words{};
typename key_type::correction_advice_array correction_advice{};
typename key_type::correction_seeds_array correction_seeds{};
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
std::array<interior_node, depth + 1> snap0{};
std::array<interior_node, depth + 1> snap1{};
HEDLEY_PRAGMA(GCC diagnostic pop)
std::array<bool, depth + 1> snap_sign{};
std::array<bool, depth + 1> need_snap{};
for (std::size_t i = 0; i <= depth; ++i)
need_snap[i] = false;
for (std::size_t i = 0; i < n; ++i)
need_snap[key_type::meta[i].tree_level] = true;
if (need_snap[0])
{
snap0[0] = st.seed0();
snap1[0] = st.seed1();
snap_sign[0] = dpf::get_lo_bit(st.seed0());
}
detail::cmp_meta cmp{};
typename key_type::value_cw_array value_cws{};
typename key_type::value_cw_array value_cw_coeff{};
typename key_type::tail_array tail{};
typename key_type::tail_array tail_coeff{};
uint64_t cw_last = 0;
uint64_t cw_last_coeff = 0;
ds_cmp_gen_state cmp_st{};
uint64_t delta = 0;
uint64_t false_value = 0;
if (cmp_spec != nullptr)
{
// The comparison channel is a dealer-equivalent computation: in this
// local joint simulation the threshold lane comes from the two shares.
// (The *leaf* phase does not reconstruct `x` at this call site — it is
// handed to `proto.open_leaf_group` below.)
const input_type cmp_x = utils::xor_input_shares(x0, x1);
const std::size_t cmp_nbits = cmp_spec->prefix ? cmp_spec->prefix
: utils::bitlength_of_v<input_type>;
cmp.nbits = static_cast<int>(cmp_nbits);
cmp.mask = cmp_spec->mask;
delta = cmp_spec->beta & cmp_spec->mask;
false_value = cmp_spec->false_value & cmp_spec->mask;
cmp.kind = cmp_spec->kind;
cmp.active = true;
cmp.incremental = cmp_spec->incremental;
cmp.eval_as_ge = false;
cmp.trivial = cmp_trivial::none;
cmp.block_width = static_cast<int>(CmpBlock);
cmp.tail_bits = static_cast<int>(key_type::cmp_q);
auto lane = lane_input(cmp_x, cmp_nbits,
utils::bitlength_of_v<input_type>);
unsigned __int128 thresh = static_cast<unsigned __int128>(
utils::to_integral_type<input_type>{}(lane));
adjust_cmp_threshold(cmp, thresh, cmp_nbits);
if (CmpBlock > 0 && is_paint_kind(cmp.kind))
throw std::invalid_argument(
"path recipes use the per-level comparison channel");
if (cmp_spec->custom && (CmpBlock > 0 || is_paint_kind(cmp.kind)))
throw std::invalid_argument(
"this comparison payload uses the per-level channel");
cmp_st.active = cmp.active;
cmp_st.nbits = cmp_nbits;
cmp_st.mask = cmp.mask;
cmp_st.beta = delta;
cmp_st.include_eq = cmp.include_eq;
cmp_st.trivial = cmp.trivial;
cmp_st.thresh = thresh;
cmp_st.kind = cmp.kind;
cmp_st.paint = is_paint_kind(cmp.kind);
cmp_st.length_bits = cmp_spec->length_bits;
cmp_st.paint_cb = cmp_spec->paint ? &paint_fn_adapter : nullptr;
cmp_st.paint_ctx = cmp_spec->paint ? &cmp_spec->paint : nullptr;
cmp_st.Va = 0;
// Wildcard payload: open CWs for δ = 0 and stash β = 1 coefficients
// (same dual-accumulator scheme as dealer `make_incremental_impl`).
if constexpr (CmpWild)
{
cmp_st.track_coeff = true;
cmp_st.Va1 = 0;
}
}
const bool custom_cmp = cmp_spec != nullptr && cmp_spec->custom;
detail::group_elem g_delta = custom_cmp
? cmp_spec->beta_g : detail::group_elem{};
detail::group_elem g_false = custom_cmp
? cmp_spec->false_g : detail::group_elem{};
detail::group_elem g_va = detail::group_zero(g_delta);
detail::group_elem g_va1 = detail::group_zero(g_delta);
detail::group_elem g_one = detail::group_one(g_delta);
detail::group_elem g_on = detail::group_zero(g_delta);
detail::group_elem g_on_unit = detail::group_zero(g_delta);
if (custom_cmp && cmp.include_eq && !is_paint_kind(cmp.kind))
{
g_on = g_delta;
g_on_unit = g_one;
}
typename key_type::value_cw_word g_last{};
typename key_type::value_cw_word g_last_coeff{};
typename key_type::prefix_cw_array prefix_cw{};
typename key_type::prefix_cw_array prefix_coeff{};
const uint64_t on_path = [&]() -> uint64_t {
if (!is_paint_kind(cmp.kind))
return (cmp.include_eq ? delta : 0ULL) & cmp.mask;
const uint64_t unit = paint_unit(cmp.kind, cmp_st.nbits, cmp_st.thresh,
cmp_st.nbits, cmp_st.length_bits, true, cmp_st.paint_cb,
cmp_st.paint_ctx);
return scale_plant(unit, delta, cmp.mask);
}();
const uint64_t on_path_unit = [&]() -> uint64_t {
if (!is_paint_kind(cmp.kind))
return (cmp.include_eq ? 1ULL : 0ULL) & cmp.mask;
const uint64_t unit = paint_unit(cmp.kind, cmp_st.nbits, cmp_st.thresh,
cmp_st.nbits, cmp_st.length_bits, true, cmp_st.paint_cb,
cmp_st.paint_ctx);
return scale_plant(unit, 1ULL, cmp.mask);
}();
auto snap_prefix = [&](std::size_t at) {
if constexpr (CmpIdcf)
{
if (!cmp.incremental || cmp.trivial != cmp_trivial::none)
return;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
const uint64_t word = proto.open_final_cw(st.seed0(), st.seed1(),
static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())), cmp_st.Va,
cmp.mask, on_path);
HEDLEY_PRAGMA(GCC diagnostic pop)
prefix_cw[at] = static_cast<typename key_type::value_cw_word>(word);
if constexpr (CmpWild)
{
const uint64_t w1 = proto.open_final_cw(st.seed0(), st.seed1(),
static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())),
cmp_st.Va1, cmp.mask, on_path_unit);
prefix_coeff[at] = static_cast<typename key_type::value_cw_word>(
(w1 + neg_m(word, cmp.mask)) & cmp.mask);
}
}
};
auto snap_prefix_group = [&](std::size_t at) {
if constexpr (CmpIdcf)
{
if (!cmp.incremental || cmp.trivial != cmp_trivial::none)
return;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
const auto word = detail::group_final_cw<InteriorPRG>(st.seed0(),
st.seed1(), static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())),
g_va, g_on);
HEDLEY_PRAGMA(GCC diagnostic pop)
prefix_cw[at] = detail::group_to_word<typename key_type::value_cw_word>(word);
if constexpr (CmpWild)
{
const auto w1 = detail::group_final_cw<InteriorPRG>(st.seed0(),
st.seed1(), static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())),
g_va1, g_on_unit);
prefix_coeff[at] = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(w1, word));
}
}
};
if constexpr (CmpIdcf)
{
if (cmp.active && !custom_cmp)
snap_prefix(0);
if (cmp.active && custom_cmp)
snap_prefix_group(0);
}
auto mask = key_type::msb_mask;
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{
uint64_t vcw = 0;
if constexpr (CmpBlock == 0)
{
if (custom_cmp && cmp_st.active && cmp_st.trivial == cmp_trivial::none
&& level < cmp_st.nbits)
{
using tree = dpf::tree_traits<InteriorPRG>;
auto s0 = st.seed0();
auto s1 = st.seed1();
const auto a0 = static_cast<uint8_t>(dpf::get_lo_bit(s0));
const auto a1 = static_cast<uint8_t>(dpf::get_lo_bit(s1));
auto v0 = tree::expand_value(s0);
auto v1 = tree::expand_value(s1);
const int ai = static_cast<int>(
(cmp_st.thresh >> (cmp_st.nbits - 1 - level)) & 1);
const auto base = detail::group_value_cw<InteriorPRG>(v0[0],
v0[1], v1[0], v1[1], a0, a1, ai, g_va, g_delta);
value_cws[level] =
detail::group_to_word<typename key_type::value_cw_word>(base);
if constexpr (CmpWild)
{
const auto v1w = detail::group_value_cw<InteriorPRG>(v0[0],
v0[1], v1[0], v1[1], a0, a1, ai, g_va1, g_one);
value_cw_coeff[level] =
detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(v1w, base));
}
}
ds_advance_level<InteriorPRG>(st, x0, x1, mask, level, depth, proto,
correction_words[level], correction_advice[level],
(!custom_cmp && cmp_st.active) ? &vcw : nullptr,
(!custom_cmp && cmp_st.active) ? &cmp_st : nullptr);
if constexpr (IsVerifiable)
{
const auto prefix = static_cast<psnip_uint64_t>(
utils::to_integral_type<input_type>{}(
utils::xor_input_shares(x0, x1))
>> (utils::bitlength_of_v<input_type> - (level + 1)));
correction_seeds[level] = detail::vdpf::make_cs(level, prefix,
st.seed0(), st.seed1());
}
if (!custom_cmp && cmp_st.active && cmp_st.trivial == cmp_trivial::none
&& level < cmp_st.nbits)
{
value_cws[level] =
static_cast<typename key_type::value_cw_word>(vcw);
if constexpr (CmpWild)
{
value_cw_coeff[level] =
static_cast<typename key_type::value_cw_word>(
cmp_st.last_vcw_coeff);
}
snap_prefix(level + 1);
}
else if (custom_cmp && cmp_st.active
&& cmp_st.trivial == cmp_trivial::none && level < cmp_st.nbits)
{
snap_prefix_group(level + 1);
}
}
else
{
ds_advance_level<InteriorPRG>(st, x0, x1, mask, level, depth, proto,
correction_words[level], correction_advice[level]);
if constexpr (IsVerifiable)
{
const auto prefix = static_cast<psnip_uint64_t>(
utils::to_integral_type<input_type>{}(
utils::xor_input_shares(x0, x1))
>> (utils::bitlength_of_v<input_type> - (level + 1)));
correction_seeds[level] = detail::vdpf::make_cs(level, prefix,
st.seed0(), st.seed1());
}
if (cmp_st.active && cmp_st.trivial == cmp_trivial::none)
{
using sched = detail::blocked::schedule<key_type::cmp_h, CmpBlock>;
const std::size_t c = level + 1;
if (c <= key_type::cmp_h && sched::contains(c))
{
const auto wi = sched::index(c);
const uint64_t word =
detail::blocked::checkpoint_word<InteriorPRG>(
st.seed0(), st.seed1(), delta, cmp.mask);
value_cws[wi] =
static_cast<typename key_type::value_cw_word>(word);
if constexpr (CmpWild)
{
value_cw_coeff[wi] =
static_cast<typename key_type::value_cw_word>(
detail::blocked::checkpoint_coeff(
st.seed0(), st.seed1(), cmp.mask));
}
}
if (c == key_type::cmp_h && key_type::cmp_q > 0)
{
uint64_t words[4]{};
uint64_t coeffs[4]{};
const uint64_t suffix = static_cast<uint64_t>(cmp_st.thresh)
& ((1ULL << key_type::cmp_q) - 1ULL);
detail::blocked::tail_words<InteriorPRG>(st.seed0(),
st.seed1(), delta, cmp.mask, cmp.include_eq, suffix,
key_type::cmp_q, words, CmpWild ? coeffs : nullptr);
for (std::size_t z = 0; z < key_type::cmp_tail; ++z)
{
tail[z] =
static_cast<typename key_type::value_cw_word>(words[z]);
if constexpr (CmpWild)
{
tail_coeff[z] =
static_cast<typename key_type::value_cw_word>(
coeffs[z]);
}
}
}
}
}
if (need_snap[level + 1])
{
snap0[level + 1] = st.seed0();
snap1[level + 1] = st.seed1();
snap_sign[level + 1] = dpf::get_lo_bit(st.seed0());
}
if constexpr (CmpBlock == 0)
{
if (cmp_st.active && cmp_st.trivial == cmp_trivial::none
&& level + 1 == cmp_st.nbits && !custom_cmp)
{
cw_last = proto.open_final_cw(st.seed0(), st.seed1(),
static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())), cmp_st.Va,
cmp_st.mask, on_path);
if constexpr (CmpWild)
{
const uint64_t l1 = proto.open_final_cw(st.seed0(), st.seed1(),
static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())),
cmp_st.Va1, cmp_st.mask, on_path_unit);
cw_last_coeff =
(l1 + neg_m(cw_last, cmp_st.mask)) & cmp_st.mask;
}
}
else if (cmp_st.active && cmp_st.trivial == cmp_trivial::none
&& level + 1 == cmp_st.nbits && custom_cmp)
{
g_last = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_final_cw<InteriorPRG>(st.seed0(), st.seed1(),
static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())), g_va, g_on));
if constexpr (CmpWild)
{
const auto l1 = detail::group_final_cw<InteriorPRG>(st.seed0(),
st.seed1(), static_cast<uint8_t>(dpf::get_lo_bit(st.seed1())),
g_va1, g_on_unit);
g_last_coeff = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(l1, detail::group_from_word(g_last, g_delta)));
}
}
}
}
constexpr std::size_t ngroups = [] {
std::size_t m = 0;
for (std::size_t i = 0; i < key_type::num_outputs; ++i)
m = std::max(m, key_type::meta[i].group_id + 1);
return m;
}();
auto leaves0 = empty_leaves<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
auto beavers0 = empty_beavers<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
auto leaves1 = empty_leaves<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
auto beavers1 = empty_beavers<ExteriorPRG, PlacedTuple>(
std::make_index_sequence<n>{});
constexpr auto order =
build_group_order<decltype(key_type::meta), ngroups>(key_type::meta, n);
// Leaf phase: hand both point shares to the protocol. In the local joint
// simulation `open_leaf_group` reconstructs `x` internally and runs
// `make_leaves` per group via this closure; the gen body never forms `x`.
using leaf_prg = std::conditional_t<IsExtractable,
detail::vdpf::extractable_leaf_prg<ExteriorPRG>, ExteriorPRG>;
proto.open_leaf_group(x0, x1, [&](input_type x) {
for_each_index(std::make_index_sequence<ngroups>{}, [&](auto oi) {
constexpr std::size_t G = order[decltype(oi)::value];
constexpr std::size_t lvl = [] {
for (std::size_t i = 0; i < key_type::num_outputs; ++i)
{
if (key_type::meta[i].group_id == G)
return key_type::meta[i].tree_level;
}
return std::size_t{0};
}();
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
gen_group<InteriorPRG, leaf_prg, InputT, PlacedTuple, MetaHolder,
G>(x, dpf::unset_lo_2bits(snap0[lvl]),
dpf::unset_lo_2bits(snap1[lvl]),
static_cast<uint8_t>(snap_sign[lvl]),
static_cast<uint8_t>(dpf::get_lo_bit(snap1[lvl])), placed,
leaves0, beavers0, leaves1, beavers1, &proto);
HEDLEY_PRAGMA(GCC diagnostic pop)
});
});
auto wrap0 = wrap_leaves<InteriorPRG, ExteriorPRG, InputT, PlacedTuple>(
leaves0, beavers0, std::make_index_sequence<n>{});
auto wrap1 = wrap_leaves<InteriorPRG, ExteriorPRG, InputT, PlacedTuple>(
leaves1, beavers1, std::make_index_sequence<n>{});
uint64_t cmp_add0 = 0, cmp_add1 = 0;
typename key_type::value_cw_word g_add0{};
typename key_type::value_cw_word g_add1{};
if (cmp.active && !custom_cmp)
{
uint64_t target = false_value;
if (cmp.trivial == cmp_trivial::always_true)
target = (delta + false_value) & cmp.mask;
else if (cmp.trivial == cmp_trivial::always_false)
target = false_value;
else if (cmp.eval_as_ge)
target = (delta + false_value) & cmp.mask;
// Group-width addend blind from the protocol. The local backend reuses
// the shared root sampler (matched tapes with the dealer); the addend
// share itself is stored at `value_cw_word` width on the key, so no
// padded uint64 rides the wire when the payload is narrow.
// For wildcard payloads, target is 0 here; `assign_cmp` later replaces
// these shares with a fresh split of the real absorb target.
const uint64_t r = proto.sample_addend_blind(cmp.mask,
[&]() -> interior_node {
return static_cast<interior_node>(root_sampler());
});
split_cmp_addend(target, cmp.mask, r, cmp_add0, cmp_add1);
}
else if (cmp.active && custom_cmp)
{
detail::group_elem target = g_false;
if (cmp.trivial == cmp_trivial::always_true || cmp.eval_as_ge)
target = detail::group_add(g_delta, g_false);
const auto blind = detail::group_from_node<InteriorPRG>(
static_cast<interior_node>(root_sampler()), g_delta);
g_add0 = detail::group_to_word<typename key_type::value_cw_word>(blind);
g_add1 = detail::group_to_word<typename key_type::value_cw_word>(
detail::group_sub(target, blind));
}
input_type off0{}, off1{};
auto adds = extract_addends(placed, std::make_index_sequence<n>{});
key_type key0{root0, correction_words, correction_advice, std::move(wrap0),
off0, cmp, value_cws, cw_last, cmp_add0, adds, value_cw_coeff,
cw_last_coeff, tail, tail_coeff, prefix_cw, prefix_coeff,
correction_seeds};
key_type key1{root1, correction_words, correction_advice, std::move(wrap1),
off1, cmp, value_cws, cw_last, cmp_add1, adds, value_cw_coeff,
cw_last_coeff, tail, tail_coeff, prefix_cw, prefix_coeff,
correction_seeds};
if (custom_cmp)
{
key0.set_cmp_scalars(g_last, g_add0, g_last_coeff);
key1.set_cmp_scalars(g_last, g_add1, g_last_coeff);
}
return dpf::make_party_key_pair(std::move(key0), std::move(key1));
}
} // namespace incr
} // namespace detail
// ---------------------------------------------------------------------------
// Convenience make_dpf(x, y...): classic or incremental (+ optional cmp)
// ---------------------------------------------------------------------------
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename ...OutputTs,
typename = std::enable_if_t<no_ic_pack_v<OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf(InputT && x, OutputTs && ...ys)
{
using input_type = std::decay_t<InputT>;
static_assert(!is_secret_share_v<input_type>,
"make_dpf: domain point must be plaintext (use raw() or reconstruct)");
static_assert((!is_secret_share_v<std::decay_t<OutputTs>> && ...),
"make_dpf: payloads must be plaintext (use raw() or reconstruct)");
using node = typename ExteriorPRG::block_type;
constexpr auto bitlen = utils::bitlength_of_v<input_type>;
constexpr std::size_t CD = forced_cmp_depth_v<bitlen, OutputTs...>;
constexpr std::size_t CB = forced_cmp_out_bits_v<bitlen, OutputTs...>;
constexpr bool CW = forced_cmp_wild_v<OutputTs...>;
constexpr std::size_t BK = forced_cmp_block_v<OutputTs...>;
constexpr bool ID = forced_cmp_idcf_v<OutputTs...>;
constexpr bool V = args_have_verifiable_v<OutputTs...>;
constexpr bool E = args_have_extractable_v<OutputTs...>;
dcf_runtime_spec dcf_spec{};
bool has_cmp = false;
auto placed = detail::incr::flatten_args_and_cmp<bitlen>(dcf_spec, has_cmp,
std::forward<OutputTs>(ys)...);
using placed_tuple = decltype(placed);
const dcf_runtime_spec * cs = has_cmp ? &dcf_spec : nullptr;
if constexpr (std::tuple_size_v<placed_tuple> == 0)
{
if (!has_cmp)
throw std::invalid_argument("make_dpf: no outputs");
return detail::incr::make_incremental_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(x, std::move(placed),
dpf::uniform_sample<node>, cs);
}
else if constexpr (!args_have_cmp_v<OutputTs...> && !args_have_eq_v<OutputTs...>
&& !args_have_verifiable_v<OutputTs...> && !args_have_extractable_v<OutputTs...>
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>()
&& !detail::incr::placed_has_arith_beta<placed_tuple>())
{
auto vals = detail::incr::classic_values(placed,
std::make_index_sequence<std::tuple_size_v<placed_tuple>>{});
return std::apply(
[&](auto && ...zs) {
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(
dpf::make_dpfargs(x, std::forward<decltype(zs)>(zs)...));
},
std::move(vals));
}
else
{
return detail::incr::make_incremental_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(x, std::move(placed),
dpf::uniform_sample<node>, cs);
}
}
template <typename InteriorPRG,
typename ExteriorPRG,
auto RootSampler,
typename InputT,
typename ...OutputTs,
typename = std::enable_if_t<no_ic_pack_v<OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf(InputT && x, OutputTs && ...ys)
{
using input_type = std::decay_t<InputT>;
static_assert(!is_secret_share_v<input_type>,
"make_dpf: domain point must be plaintext (use raw() or reconstruct)");
static_assert((!is_secret_share_v<std::decay_t<OutputTs>> && ...),
"make_dpf: payloads must be plaintext (use raw() or reconstruct)");
using node = typename ExteriorPRG::block_type;
constexpr auto bitlen = utils::bitlength_of_v<input_type>;
constexpr std::size_t CD = forced_cmp_depth_v<bitlen, OutputTs...>;
constexpr std::size_t CB = forced_cmp_out_bits_v<bitlen, OutputTs...>;
constexpr bool CW = forced_cmp_wild_v<OutputTs...>;
constexpr std::size_t BK = forced_cmp_block_v<OutputTs...>;
constexpr bool ID = forced_cmp_idcf_v<OutputTs...>;
constexpr bool V = args_have_verifiable_v<OutputTs...>;
constexpr bool E = args_have_extractable_v<OutputTs...>;
dcf_runtime_spec dcf_spec{};
bool has_cmp = false;
auto placed = detail::incr::flatten_args_and_cmp<bitlen>(dcf_spec, has_cmp,
std::forward<OutputTs>(ys)...);
using placed_tuple = decltype(placed);
const dcf_runtime_spec * cs = has_cmp ? &dcf_spec : nullptr;
auto seed = root_sampler_t<InteriorPRG>{RootSampler};
if constexpr (std::tuple_size_v<placed_tuple> == 0)
{
if (!has_cmp)
throw std::invalid_argument("make_dpf: no outputs");
return detail::incr::make_incremental_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(x, std::move(placed), seed, cs);
}
else if constexpr (!args_have_cmp_v<OutputTs...> && !args_have_eq_v<OutputTs...>
&& !args_have_verifiable_v<OutputTs...> && !args_have_extractable_v<OutputTs...>
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>()
&& !detail::incr::placed_has_arith_beta<placed_tuple>())
{
auto vals = detail::incr::classic_values(placed,
std::make_index_sequence<std::tuple_size_v<placed_tuple>>{});
return std::apply(
[&](auto && ...zs) {
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(
dpf::make_dpfargs(x, std::forward<decltype(zs)>(zs)...),
RootSampler);
},
std::move(vals));
}
else
{
return detail::incr::make_incremental_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(x, std::move(placed), seed, cs);
}
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename ...OutputTs,
typename = std::enable_if_t<no_ic_pack_v<OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf(InputT && x, root_sampler_t<InteriorPRG> root_sampler,
OutputTs && ...ys)
{
using input_type = std::decay_t<InputT>;
static_assert(!is_secret_share_v<input_type>,
"make_dpf: domain point must be plaintext (use raw() or reconstruct)");
static_assert((!is_secret_share_v<std::decay_t<OutputTs>> && ...),
"make_dpf: payloads must be plaintext (use raw() or reconstruct)");
constexpr auto bitlen = utils::bitlength_of_v<input_type>;
constexpr std::size_t CD = forced_cmp_depth_v<bitlen, OutputTs...>;
constexpr std::size_t CB = forced_cmp_out_bits_v<bitlen, OutputTs...>;
constexpr bool CW = forced_cmp_wild_v<OutputTs...>;
constexpr std::size_t BK = forced_cmp_block_v<OutputTs...>;
constexpr bool ID = forced_cmp_idcf_v<OutputTs...>;
constexpr bool V = args_have_verifiable_v<OutputTs...>;
constexpr bool E = args_have_extractable_v<OutputTs...>;
dcf_runtime_spec dcf_spec{};
bool has_cmp = false;
auto placed = detail::incr::flatten_args_and_cmp<bitlen>(dcf_spec, has_cmp,
std::forward<OutputTs>(ys)...);
using placed_tuple = decltype(placed);
using node = typename ExteriorPRG::block_type;
const dcf_runtime_spec * cs = has_cmp ? &dcf_spec : nullptr;
if constexpr (std::tuple_size_v<placed_tuple> == 0)
{
if (!has_cmp)
throw std::invalid_argument("make_dpf: no outputs");
return detail::incr::make_incremental_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(x, std::move(placed), root_sampler,
cs);
}
else if constexpr (!args_have_cmp_v<OutputTs...> && !args_have_eq_v<OutputTs...>
&& !args_have_verifiable_v<OutputTs...> && !args_have_extractable_v<OutputTs...>
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>()
&& !detail::incr::placed_has_arith_beta<placed_tuple>())
{
auto vals = detail::incr::classic_values(placed,
std::make_index_sequence<std::tuple_size_v<placed_tuple>>{});
return std::apply(
[&](auto && ...zs) {
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(
dpf::make_dpfargs(x, std::forward<decltype(zs)>(zs)...),
std::move(root_sampler));
},
std::move(vals));
}
else
{
return detail::incr::make_incremental_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(x, std::move(placed), root_sampler,
cs);
}
}
// ---------------------------------------------------------------------------
// make_dpf_doerner_shelat(x0, x1, ...): classic or incremental
// ---------------------------------------------------------------------------
/// @brief Doerner–Shelat keygen with caller-supplied roots and pad stream.
/// @details The point is `x0 XOR x1` (before the signed-MSB flip `make_dpf` applies).
/// `rng.pad` is the Beaver randomness; it cancels and must not draw from
/// `uniform_fill` when beaver coins are being matched to `make_dpf`.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam InputT input domain type
/// @tparam OutputTs output ts
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param rng the Doerner–Shelat randomness tapes
/// @param ys the `ys`
/// @return Doerner–Shelat keygen with caller-supplied roots and pad stream
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename PadRng,
typename InputT,
typename ...OutputTs,
typename = std::enable_if_t<no_ic_pack_v<OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputTs && ...ys)
{
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
false, std::move(x0), std::move(x1), std::move(rng),
std::forward<OutputTs>(ys)...);
}
/// @brief Doerner–Shelat with additive shares: the point is `x0 + x1` in the input
/// ring (unsigned wrap; signed MSB flipped after the carry chain).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam InputT input domain type
/// @tparam OutputTs output ts
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param rng the Doerner–Shelat randomness tapes
/// @param ys the `ys`
/// @return Doerner–Shelat with additive shares: the point is `x0 + x1` in the input ring (unsigned
/// wrap; signed MSB flipped after the carry chain)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename PadRng,
typename InputT,
typename ...OutputTs,
typename = std::enable_if_t<no_ic_pack_v<OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_input_t, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputTs && ...ys)
{
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
true, std::move(x0), std::move(x1), std::move(rng),
std::forward<OutputTs>(ys)...);
}
/// @brief XOR-index shares, additively shared payload `y0 + y1 = β` (single concrete).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputT output type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param rng the Doerner–Shelat randomness tapes
/// @param y0 the `y0`
/// @param y1 the `y1`
/// @return XOR-index shares, additively shared payload `y0 + y1 = β` (single concrete)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename PadRng,
typename InputT,
typename OutputT,
typename = std::enable_if_t<!is_arith_beta_v<OutputT> && !is_at_v<OutputT>
&& !is_wildcard_v<OutputT> && no_ic_pack_v<OutputT>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_output_t, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputT y0, OutputT y1)
{
local_cw_protocol<PadRng> proto{rng.pad};
return detail::make_dpf_doerner_shelat_impl<InteriorPRG, ExteriorPRG>(
false, true, std::move(x0), std::move(x1), rng.root, proto,
std::move(y0), std::move(y1));
}
/// @brief Additive index and additive payload shares (single concrete).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputT output type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param rng the Doerner–Shelat randomness tapes
/// @param y0 the `y0`
/// @param y1 the `y1`
/// @return Additive index and additive payload shares (single concrete)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename PadRng,
typename InputT,
typename OutputT,
typename = std::enable_if_t<!is_arith_beta_v<OutputT> && !is_at_v<OutputT>
&& !is_wildcard_v<OutputT> && no_ic_pack_v<OutputT>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_input_t, arith_output_t, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputT y0, OutputT y1)
{
local_cw_protocol<PadRng> proto{rng.pad};
return detail::make_dpf_doerner_shelat_impl<InteriorPRG, ExteriorPRG>(
true, true, std::move(x0), std::move(x1), rng.root, proto,
std::move(y0), std::move(y1));
}
/// @brief Shared payloads via `arith_beta` / packs (XOR index).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputTs output ts
/// @tparam OutputTs output ts
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param rng the Doerner–Shelat randomness tapes
/// @param y the `y`
/// @param ys the `ys`
/// @return Shared payloads via `arith_beta` / packs (XOR index)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename PadRng,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename = std::enable_if_t<
(is_arith_beta_v<OutputT> || is_at_v<OutputT> || (sizeof...(OutputTs) > 0))
&& no_ic_pack_v<OutputT, OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_output_t, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputT && y, OutputTs && ...ys)
{
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
false, std::move(x0), std::move(x1), std::move(rng),
std::forward<OutputT>(y), std::forward<OutputTs>(ys)...);
}
/// @brief Shared payloads via `arith_beta` / packs (additive index).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputTs output ts
/// @tparam OutputTs output ts
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param rng the Doerner–Shelat randomness tapes
/// @param y the `y`
/// @param ys the `ys`
/// @return Shared payloads via `arith_beta` / packs (additive index)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename PadRng,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename = std::enable_if_t<
(is_arith_beta_v<OutputT> || is_at_v<OutputT> || (sizeof...(OutputTs) > 0))
&& no_ic_pack_v<OutputT, OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_input_t, arith_output_t, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputT && y, OutputTs && ...ys)
{
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
true, std::move(x0), std::move(x1), std::move(rng),
std::forward<OutputT>(y), std::forward<OutputTs>(ys)...);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename PadRng,
typename InputT,
typename ...OutputTs,
typename = std::enable_if_t<no_ic_pack_v<OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(bool arith, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputTs && ...ys)
{
using input_type = std::decay_t<InputT>;
static_assert(!is_secret_share_v<input_type>,
"Doerner–Shelat: use additive_share of xor_wrapper, or raw shares");
using node = typename ExteriorPRG::block_type;
constexpr auto bitlen = utils::bitlength_of_v<input_type>;
constexpr std::size_t CD = forced_cmp_depth_v<bitlen, OutputTs...>;
constexpr std::size_t CB = forced_cmp_out_bits_v<bitlen, OutputTs...>;
constexpr bool CW = forced_cmp_wild_v<OutputTs...>;
constexpr std::size_t BK = forced_cmp_block_v<OutputTs...>;
constexpr bool ID = forced_cmp_idcf_v<OutputTs...>;
constexpr bool V = args_have_verifiable_v<OutputTs...>;
constexpr bool E = args_have_extractable_v<OutputTs...>;
dcf_runtime_spec dcf_spec{};
bool has_cmp = false;
auto placed = detail::incr::flatten_args_and_cmp<bitlen>(dcf_spec, has_cmp,
std::forward<OutputTs>(ys)...);
using placed_tuple = decltype(placed);
const dcf_runtime_spec * cs = has_cmp ? &dcf_spec : nullptr;
if constexpr (std::tuple_size_v<placed_tuple> == 0)
{
if (!has_cmp)
throw std::invalid_argument("make_dpf_doerner_shelat: no outputs");
local_cw_protocol<PadRng> proto{rng.pad};
return detail::incr::make_incremental_ds_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(arith, std::move(x0),
std::move(x1), rng.root, proto, std::move(placed), cs);
}
else if constexpr (!args_have_cmp_v<OutputTs...> && !args_have_eq_v<OutputTs...>
&& !args_have_verifiable_v<OutputTs...> && !args_have_extractable_v<OutputTs...>
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>()
&& !detail::incr::placed_has_arith_beta<placed_tuple>())
{
auto vals = detail::incr::classic_values(placed,
std::make_index_sequence<std::tuple_size_v<placed_tuple>>{});
local_cw_protocol<PadRng> proto{rng.pad};
return std::apply(
[&](auto && ...zs) {
return detail::make_dpf_doerner_shelat_impl<InteriorPRG,
ExteriorPRG>(arith, std::move(x0), std::move(x1), rng.root,
proto, std::forward<decltype(zs)>(zs)...);
},
std::move(vals));
}
else
{
local_cw_protocol<PadRng> proto{rng.pad};
return detail::incr::make_incremental_ds_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(arith, std::move(x0),
std::move(x1), rng.root, proto, std::move(placed), cs);
}
}
/// @brief Doerner–Shelat with an injectable `CwProtocol` (local or MPC backend).
/// @details Signature is `make_dpf_doerner_shelat(x0, x1, root_sampler, proto, y...)`
/// so it does not collide with the `ds_randomness` or bare-output overloads.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam CwProtocol correction-word protocol
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputTs output ts
/// @tparam OutputTs output ts
/// @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 Doerner–Shelat with an injectable `CwProtocol` (local or MPC backend)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename CwProtocol,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename = std::enable_if_t<
detail::is_cw_protocol<std::decay_t<CwProtocol>>::value
&& no_ic_pack_v<OutputT, OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(InputT x0, InputT x1, RootSampler root_sampler,
CwProtocol & proto, OutputT && y, OutputTs && ...ys)
{
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
false, std::move(x0), std::move(x1), std::move(root_sampler), proto,
std::forward<OutputT>(y), std::forward<OutputTs>(ys)...);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename CwProtocol,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename = std::enable_if_t<
detail::is_cw_protocol<std::decay_t<CwProtocol>>::value
&& no_ic_pack_v<OutputT, OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_input_t, InputT x0, InputT x1,
RootSampler root_sampler, CwProtocol & proto, OutputT && y,
OutputTs && ...ys)
{
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
true, std::move(x0), std::move(x1), std::move(root_sampler), proto,
std::forward<OutputT>(y), std::forward<OutputTs>(ys)...);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename RootSampler,
typename CwProtocol,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename = std::enable_if_t<
detail::is_cw_protocol<std::decay_t<CwProtocol>>::value
&& no_ic_pack_v<OutputT, OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(bool arith, InputT x0, InputT x1,
RootSampler root_sampler, CwProtocol & proto, OutputT && y,
OutputTs && ...ys)
{
using input_type = std::decay_t<InputT>;
using node = typename ExteriorPRG::block_type;
constexpr auto bitlen = utils::bitlength_of_v<input_type>;
constexpr std::size_t CD = forced_cmp_depth_v<bitlen, OutputT, OutputTs...>;
constexpr std::size_t CB = forced_cmp_out_bits_v<bitlen, OutputT, OutputTs...>;
constexpr bool CW = forced_cmp_wild_v<OutputT, OutputTs...>;
constexpr std::size_t BK = forced_cmp_block_v<OutputT, OutputTs...>;
constexpr bool ID = forced_cmp_idcf_v<OutputT, OutputTs...>;
constexpr bool V = args_have_verifiable_v<OutputT, OutputTs...>;
constexpr bool E = args_have_extractable_v<OutputT, OutputTs...>;
dcf_runtime_spec dcf_spec{};
bool has_cmp = false;
auto placed = detail::incr::flatten_args_and_cmp<bitlen>(dcf_spec, has_cmp,
std::forward<OutputT>(y), std::forward<OutputTs>(ys)...);
using placed_tuple = decltype(placed);
const dcf_runtime_spec * cs = has_cmp ? &dcf_spec : nullptr;
if constexpr (std::tuple_size_v<placed_tuple> == 0)
{
if (!has_cmp)
throw std::invalid_argument("make_dpf_doerner_shelat: no outputs");
return detail::incr::make_incremental_ds_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(arith, std::move(x0),
std::move(x1), root_sampler, proto, std::move(placed), cs);
}
else if constexpr (!args_have_cmp_v<OutputT, OutputTs...>
&& !args_have_verifiable_v<OutputT, OutputTs...>
&& !args_have_extractable_v<OutputT, OutputTs...>
&& !args_have_eq_v<OutputT, OutputTs...>
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>()
&& !detail::incr::placed_has_arith_beta<placed_tuple>())
{
auto vals = detail::incr::classic_values(placed,
std::make_index_sequence<std::tuple_size_v<placed_tuple>>{});
return std::apply(
[&](auto && ...zs) {
return detail::make_dpf_doerner_shelat_impl<InteriorPRG,
ExteriorPRG>(arith, std::move(x0), std::move(x1),
root_sampler, proto, std::forward<decltype(zs)>(zs)...);
},
std::move(vals));
}
else
{
return detail::incr::make_incremental_ds_impl<InteriorPRG, ExteriorPRG,
input_type, placed_tuple, CD, CB, CW, BK, ID, V, E>(arith, std::move(x0),
std::move(x1), root_sampler, proto, std::move(placed), cs);
}
}
/// @brief Doerner–Shelat keygen. Roots and pads come from `uniform_sample`.
/// @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 OutputTs output ts
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param y the `y`
/// @param ys the `ys`
/// @return Doerner–Shelat keygen
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename = std::enable_if_t<
!detail::is_ds_randomness<std::decay_t<OutputT>>::value
&& !detail::is_cw_protocol<std::decay_t<OutputT>>::value
&& !detail::first_is_cw_protocol<OutputTs...>::value
&& no_ic_pack_v<OutputT, OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(InputT x0, InputT x1, OutputT && y, OutputTs && ...ys)
{
using block = typename InteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
dpf::uniform_sample<block>, {}};
HEDLEY_PRAGMA(GCC diagnostic pop)
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
std::move(x0), std::move(x1), rng, std::forward<OutputT>(y),
std::forward<OutputTs>(ys)...);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename = std::enable_if_t<
!detail::is_ds_randomness<std::decay_t<OutputT>>::value
&& !detail::is_cw_protocol<std::decay_t<OutputT>>::value
&& !detail::first_is_cw_protocol<OutputTs...>::value
&& no_ic_pack_v<OutputT, OutputTs...>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_input_t, InputT x0, InputT x1, OutputT && y,
OutputTs && ...ys)
{
using block = typename InteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
dpf::uniform_sample<block>, {}};
HEDLEY_PRAGMA(GCC diagnostic pop)
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
arith_input, std::move(x0), std::move(x1), rng, std::forward<OutputT>(y),
std::forward<OutputTs>(ys)...);
}
/// @brief XOR-index, additive payload shares (urandom roots/pads).
/// @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 OutputT output type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param y0 the `y0`
/// @param y1 the `y1`
/// @return XOR-index, additive payload shares (urandom roots/pads)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename = std::enable_if_t<!is_arith_beta_v<OutputT> && !is_at_v<OutputT>
&& !is_wildcard_v<OutputT>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_output_t, InputT x0, InputT x1, OutputT y0,
OutputT y1)
{
using block = typename InteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
dpf::uniform_sample<block>, {}};
HEDLEY_PRAGMA(GCC diagnostic pop)
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
arith_output, std::move(x0), std::move(x1), rng, std::move(y0),
std::move(y1));
}
/// @brief Additive index and additive payload (urandom roots/pads).
/// @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 OutputT output type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param y0 the `y0`
/// @param y1 the `y1`
/// @return Additive index and additive payload (urandom roots/pads)
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename = std::enable_if_t<!is_arith_beta_v<OutputT> && !is_at_v<OutputT>
&& !is_wildcard_v<OutputT>>>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(arith_input_t, arith_output_t, InputT x0, InputT x1,
OutputT y0, OutputT y1)
{
using block = typename InteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
dpf::uniform_sample<block>, {}};
HEDLEY_PRAGMA(GCC diagnostic pop)
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
arith_input, arith_output, std::move(x0), std::move(x1), rng,
std::move(y0), std::move(y1));
}
/// @brief Doerner–Shelat from party-tagged additive XOR shares of the point.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam U rebound value type
/// @tparam Args args
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param args the arguments forwarded to the constructor
/// @return Doerner–Shelat from party-tagged additive XOR shares of the point
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename U,
typename ...Args>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_doerner_shelat(
const additive_share<xor_wrapper<U>, 0> & x0,
const additive_share<xor_wrapper<U>, 1> & x1,
Args && ...args)
{
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
x0.raw(), x1.raw(), std::forward<Args>(args)...);
}
// ---------------------------------------------------------------------------
// Wildcard comparison payload assignment (cw-protocol item 4)
// ---------------------------------------------------------------------------
namespace detail
{
namespace incr
{
/// @brief Resolve the constant absorb target for a comparison payload δ = if_true −
/// if_false, matching the branch logic used at keygen.
/// @param ch the `ch`
/// @param delta the payload difference `if_true - if_false`
/// @param false_value the `false_value`
/// @return Resolve the constant absorb target for a comparison payload δ = if_true − if_false,
/// matching the branch logic used at keygen
HEDLEY_NO_THROW
inline uint64_t cmp_assign_target(const detail::cmp_meta & ch, uint64_t delta,
uint64_t false_value) noexcept
{
const uint64_t mask = ch.mask;
if (ch.trivial == cmp_trivial::always_true)
return (delta + false_value) & mask;
if (ch.trivial == cmp_trivial::always_false)
return false_value & mask;
if (ch.eval_as_ge)
return (delta + false_value) & mask;
return false_value & mask;
}
} // namespace incr
} // namespace detail
/// @brief Assign a previously-wildcard comparison payload on a *pair* of keys.
///
/// The key was generated with `dpf::lt(dpf::wildcard<Beta>{}, ...)` (or any
/// `leq`/`gt`/`geq`), which opened the value CWs / `cw_last` for δ = 0 and
/// stashed the per-level δ-coefficients. This patches both keys' (public,
/// identical) value CWs in place — `value_cw[i] += coeff[i]·δ`,
/// `cw_last += coeff_last·δ` — and splits the constant absorb into fresh
/// additive `cmp_addend` shares. No tree re-walk and no re-PRG.
/// @tparam KeyT key type
/// @tparam Beta payload type
/// @param key0 the `key0`
/// @param key1 the `key1`
/// @param if_true the payload on a true comparison
/// @param if_false the payload on a false comparison
template <typename KeyT, typename Beta>
void assign_cmp(KeyT & key0, KeyT & key1, const Beta & if_true,
const Beta & if_false = Beta{})
{
static_assert(KeyT::cmp_is_wildcard,
"assign_cmp: key's comparison payload is not a wildcard");
if (!key0.has_cmp() || !key1.has_cmp())
throw std::invalid_argument("assign_cmp: key has no comparison channel");
const auto & ch = key0.cmp();
using Concrete = dpf::concrete_type_t<Beta>;
if constexpr (detail::cmp_group_info<Concrete>::custom)
{
const auto layout = detail::group_layout<Concrete>();
const auto delta = detail::group_sub(
detail::group_from_beta(if_true), detail::group_from_beta(if_false));
const auto false_value = detail::group_from_beta(if_false);
detail::group_elem target = false_value;
if (ch.trivial == cmp_trivial::always_true || ch.eval_as_ge)
target = detail::group_add(delta, false_value);
using prg = typename KeyT::interior_prg;
const auto blind = detail::group_from_node<prg>(
dpf::uniform_sample<typename KeyT::interior_node>(), layout);
const auto add0 = detail::group_to_word<typename KeyT::value_cw_word>(blind);
const auto add1 = detail::group_to_word<typename KeyT::value_cw_word>(
detail::group_sub(target, blind));
key0.assign_cmp_group(delta, add0);
key1.assign_cmp_group(delta, add1);
return;
}
const uint64_t mask = ch.mask;
const uint64_t delta =
detail::dcf_impl::beta_delta_u64(if_true, if_false, mask);
const uint64_t false_value =
detail::dcf_impl::beta_to_u64_simple(if_false, mask);
const uint64_t target =
detail::incr::cmp_assign_target(ch, delta, false_value);
// Fresh additive split of the absorb target. Any blind works for
// reconstruction; draw one so neither share reveals the payload.
uint64_t add0 = 0, add1 = 0;
const uint64_t r = detail::dcf_impl::sample_addend_blind(mask,
[] { return dpf::uniform_sample<typename KeyT::interior_node>(); });
detail::incr::split_cmp_addend(target, mask, r, add0, add1);
key0.assign_cmp_delta(delta, add0);
key1.assign_cmp_delta(delta, add1);
}
/// @brief Party-tagged overload: `make_dpf` returns distinct `party_key<0>` /
/// `party_key<1>` types, so the same-type pair overload cannot bind both.
/// @tparam Key key type
/// @tparam Beta payload type
/// @param key0 the `key0`
/// @param key1 the `key1`
/// @param if_true the payload on a true comparison
/// @param if_false the payload on a false comparison
template <typename Key, typename Beta>
void assign_cmp(party_key<0, Key> & key0, party_key<1, Key> & key1,
const Beta & if_true, const Beta & if_false = Beta{})
{
assign_cmp(key0.key(), key1.key(), if_true, if_false);
}
/// @brief Party-local variant: patch one key with a *public* (already-opened) δ and a
/// caller-supplied `cmp_addend` share. Both parties must call this with the
/// same δ (so the public value CWs stay identical) and additive shares of the
/// absorb target that reconstruct to `cmp_assign_target(cmp(), δ, if_false)`.
/// @tparam KeyT key type
/// @param key the `key`
/// @param delta the payload difference `if_true - if_false`
/// @param addend_share the `addend_share`
/// @throws std::invalid_argument if `key has no comparison channel`
template <typename KeyT>
void assign_cmp_local(KeyT & key, uint64_t delta, uint64_t addend_share)
{
static_assert(KeyT::cmp_is_wildcard,
"assign_cmp_local: key's comparison payload is not a wildcard");
if (!key.has_cmp())
throw std::invalid_argument(
"assign_cmp_local: key has no comparison channel");
key.assign_cmp_delta(delta, addend_share);
}
/// @brief Share-typed overload: subtractive shares are converted with the party
/// coefficient before the existing additive leaf absorb math.
/// @tparam KeyT key type
/// @tparam T value type
/// @tparam Party party index, `0` or `1`
/// @tparam Scheme scheme
/// @param key the `key`
/// @param delta the payload difference `if_true - if_false`
/// @param addend_share the `addend_share`
template <typename KeyT, typename T, std::size_t Party, sharing Scheme>
void assign_cmp_local(KeyT & key, uint64_t delta,
const secret_share<T, Party, Scheme> & addend_share)
{
if constexpr (is_party_key_v<KeyT>)
{
static_assert(party_of_v<KeyT> == Party,
"assign_cmp_local: share party must match party_key");
}
const auto additive = addend_share.as_additive();
assign_cmp_local(key, delta,
static_cast<uint64_t>(additive.raw()));
}
// ---------------------------------------------------------------------------
// Point / interval / sequence / cmp eval (detail + public sugar)
// ---------------------------------------------------------------------------
namespace internal
{
template <typename DpfKey, typename InputT, typename PathMemoizer>
inline void eval_to_level(const DpfKey & dpf, InputT && x, PathMemoizer && path,
std::size_t to_level)
{
detail::ensure_level(dpf, x, path, to_level);
}
} // namespace internal
namespace detail
{
namespace incr
{
/// @brief Evaluate output slot `I` of an incremental key at the programmed point.
/// @details `N` must match the prefix of that slot (`at<N>` or full input width).
/// @tparam N width in bits
/// @tparam I output index
/// @tparam KeyT key type
/// @tparam QueryT query type
/// @tparam PathMemoizer path memoizer type
/// @param dpf the DPF key
/// @param x the `x`
/// @param path the root-to-leaf path
/// @return the evaluation result
template <std::size_t N, std::size_t I = 0, typename KeyT, typename QueryT,
typename PathMemoizer = dpf::nonmemoizing_path_memoizer<KeyT>>
auto eval_out_point_impl(
const KeyT & dpf,
QueryT && x, PathMemoizer && path = PathMemoizer{})
{
using key_type = KeyT;
static_assert(key_type::meta[I].prefix == N,
"out point eval: N does not match the prefix of output I");
using output_type = typename key_type::template concrete_output_type<I>;
// `traverse_exterior<I>` is `noexcept`; an unassigned wildcard leaf would
// otherwise `std::terminate` there. Fail loudly (matching the classic
// `eval_point` path) before touching the leaf.
if constexpr (dpf::is_wildcard_v<typename key_type::template output_type_t<I>>)
dpf::assert_not_wildcard_output<I>(dpf);
auto tx = dpf.offset_x(x);
utils::flip_msb_if_signed_integral(tx);
constexpr auto level = key_type::meta[I].tree_level;
internal::eval_to_level(dpf, tx, path, level);
auto node = dpf.template traverse_exterior<I>(path[level]);
auto lane_x = detail::incr::lane_input(tx, N, key_type::input_bits);
// Public eq if_false: absorb so subtractive reconstruction picks it up.
// Leaf shares open as y0 − y1 (XOR groups: y0 ⊕ y1); party 0 absorbs.
// Wildcard slots carry no public addend and cannot be materialized until
// assigned; `traverse_exterior<I>` above already throws for an unassigned
// wildcard leaf, so the addend fold is simply skipped for wildcard slots
// (its `public_addends` element is wildcard-typed and not comparable).
if constexpr (!dpf::is_wildcard_v<typename key_type::template output_type_t<I>>)
{
const auto & add = std::get<I>(dpf.public_addends);
if (add != output_type{})
{
const bool absorb = !dpf::get_lo_bit(dpf.root()); // party 0
if (absorb)
{
using exterior_node = typename key_type::exterior_node;
auto addon = dpf::make_naked_leaf<exterior_node>(lane_x, add);
node = dpf::add_leaf<output_type>(node, addon);
}
}
}
return make_eval_dpf_output<key_type, output_type>(node, lane_x);
}
} // namespace incr
} // namespace detail
/// @brief Evaluate the first deepest-prefix output (plan default for `eval_point`).
/// @tparam KeyT key type
/// @tparam QueryT query type
/// @tparam PathMemoizer path memoizer type
/// @tparam KeyT key type
/// @param dpf the DPF key
/// @param x the `x`
/// @param path the root-to-leaf path
/// @return the evaluation result
template <typename KeyT, typename QueryT,
typename PathMemoizer = dpf::nonmemoizing_path_memoizer<KeyT>,
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
auto eval_point(
const KeyT & dpf,
QueryT && x, PathMemoizer && path = PathMemoizer{})
{
using key_type = KeyT;
constexpr auto I = key_type::deepest_output;
return detail::incr::eval_out_point_impl<key_type::meta[I].prefix, I>(dpf, std::forward<QueryT>(x),
std::forward<PathMemoizer>(path));
}
namespace detail
{
namespace incr
{
// ---------------------------------------------------------------------------
// Incremental interval / full / sequence (per-slot stop level + pos_base)
// ---------------------------------------------------------------------------
} // namespace incr
} // namespace detail
/// @brief Per-slot buffer: `num_leaf_nodes * outputs_per_leaf_of<I>`.
/// @tparam I output index
/// @tparam KeyT key type
/// @param num_leaf_nodes the `num_leaf_nodes`
/// @return Per-slot buffer: `num_leaf_nodes * outputs_per_leaf_of<I>`
template <std::size_t I = 0, typename KeyT>
auto make_output_buffer_for(const KeyT &, std::size_t num_leaf_nodes)
{
using output_type = typename KeyT::template concrete_output_type<I>;
using buffer_elem = leaf_buffer_elem_t<KeyT, output_type>;
constexpr auto opl = KeyT::template outputs_per_leaf_of<I>;
return dpf::output_buffer<buffer_elem>(num_leaf_nodes * opl);
}
namespace detail
{
namespace incr
{
/// @brief Buffer sized for lane-domain interval `[from, to]` of output `I` at prefix `N`.
/// @tparam N width in bits
/// @tparam I output index
/// @tparam KeyT key type
/// @tparam LaneT lane type
/// @param key the `key`
/// @param from the inclusive start of the range
/// @param to the `to`
/// @return Buffer sized for lane-domain interval `[from, to]` of output `I` at prefix `N`
template <std::size_t N, std::size_t I = 0, typename KeyT, typename LaneT>
auto make_output_buffer_for_out_interval_impl(const KeyT & key, LaneT from, LaneT to)
{
static_assert(KeyT::meta[I].prefix == N, "interval buffer: prefix mismatch");
using output_type = typename KeyT::template concrete_output_type<I>;
using buffer_elem = leaf_buffer_elem_t<KeyT, output_type>;
constexpr auto opl = KeyT::template outputs_per_leaf_of<I>;
using integral_type = typename KeyT::integral_type;
constexpr auto to_int = utils::to_integral_type<LaneT>{};
utils::flip_msb_if_signed_integral(from);
utils::flip_msb_if_signed_integral(to);
constexpr auto lg = KeyT::template lg_outputs_per_leaf_of<I>;
const auto from_i = static_cast<integral_type>(to_int(from));
const auto to_i = static_cast<integral_type>(to_int(to));
const integral_type from_node = utils::leaf_node_floor(from_i, lg);
const integral_type to_node = utils::leaf_node_ceil_exclusive(to_i, lg);
const bool wraps = utils::interval_wraps(from_i, to_i, N);
const auto segs = utils::split_leaf_nodes(from_node, to_node,
KeyT::meta[I].tree_level, wraps);
return dpf::output_buffer<buffer_elem>(segs.total * opl);
}
namespace internal
{
template <std::size_t N, std::size_t I, typename DpfKey, typename IntegralT,
typename IntervalMemoizer>
void eval_out_interval_interior(const DpfKey & dpf, IntegralT from_node,
IntegralT to_node, IntervalMemoizer & memoizer)
{
using dpf_type = DpfKey;
using node_type = typename DpfKey::interior_node;
constexpr auto lg_opl = dpf_type::template lg_outputs_per_leaf_of<I>;
constexpr auto to_level = dpf_type::meta[I].tree_level;
static_assert(dpf_type::meta[I].prefix == N, "out interval eval: prefix");
// Lane MSB for the N-bit subdomain (same role as key.msb_mask for classic).
using input_type = typename dpf_type::input_type;
const input_type lane_msb =
static_cast<input_type>(input_type{1} << (N - 1));
std::size_t level_index = memoizer.assign_interval(dpf, from_node, to_node);
std::size_t nodes_at_level = memoizer.get_nodes_at_level();
IntegralT mask = static_cast<IntegralT>(
utils::to_integral_type<input_type>{}(lane_msb)
>> (level_index - 1 + lg_opl));
for (; level_index <= to_level;
level_index = memoizer.advance_level(),
nodes_at_level = memoizer.get_nodes_at_level(), mask >>= 1)
{
std::size_t i = 0, j = 0;
bool from_offset = mask & from_node,
to_offset = from_offset ^ (nodes_at_level & 1);
const node_type cw[2] = {
dpf.correction_word(level_index - 1, 0),
dpf.correction_word(level_index - 1, 1)};
const bool is_last = dpf_type::tree::is_last_level(level_index - 1,
dpf_type::depth);
auto *prev = memoizer[level_index - 1];
auto *curr = memoizer[level_index];
if (from_offset == true)
{
curr[i++] = dpf_type::traverse_interior(prev[j++], cw[1], 1, is_last);
}
const std::size_t both_end = nodes_at_level - to_offset;
while (i + 8 <= both_end)
{
alignas(node_type) node_type parents[4];
alignas(node_type) node_type left[4];
alignas(node_type) node_type right[4];
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
parents[t] = prev[j + t];
dpf_type::traverse_interior01_x4(parents, cw[0], cw[1], left, right,
is_last);
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
{
curr[i + 2 * t] = left[t];
curr[i + 2 * t + 1] = right[t];
}
i += 8;
j += 4;
}
DPF_UNROLL_LOOP
for (; i < both_end;)
{
auto cur_node = prev[j++];
auto kids = dpf_type::traverse_interior01(cur_node, cw[0], cw[1],
is_last);
curr[i++] = kids[0];
curr[i++] = kids[1];
}
if (to_offset == true)
{
curr[i] = dpf_type::traverse_interior(prev[j], cw[0], 0, is_last);
}
}
}
template <std::size_t N, std::size_t I, typename DpfKey, typename OutputBuffer,
typename IntervalMemoizer, typename IntegralT>
void eval_out_interval_exterior(const DpfKey & dpf, IntegralT from_node,
IntegralT to_node, OutputBuffer && outbuf, IntervalMemoizer && memoizer,
std::size_t start = 0)
{
using dpf_type = DpfKey;
using output_type = typename dpf_type::template concrete_output_type<I>;
constexpr auto opl = dpf_type::template outputs_per_leaf_of<I>;
constexpr auto to_level = dpf_type::meta[I].tree_level;
if (HEDLEY_UNLIKELY(to_node < from_node && to_node != IntegralT{0}))
throw std::runtime_error("to_node<from_node");
std::size_t nodes_in_interval =
static_cast<std::size_t>(to_node - from_node);
auto *nodes = memoizer[to_level];
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
DPF_UNROLL_LOOP
for (std::size_t j = 0, k = start; j < nodes_in_interval; ++j, ++k)
{
auto leaf = dpf.template traverse_exterior<I>(nodes[j]);
if constexpr (utils::is_packed_subbyte_v<output_type>)
{
store_leaf_bytes(outbuf, k, leaf);
}
else
{
std::memcpy(&outbuf[k * opl], &leaf, sizeof(output_type) * opl);
}
}
HEDLEY_PRAGMA(GCC diagnostic pop)
}
} // namespace internal
/// @brief Evaluate output `I` over an interval in the N-bit lane subdomain.
/// @details `from`/`to` are lane values in `[0, 2^N)` (not the full input domain).
/// @tparam N width in bits
/// @tparam I output index
/// @tparam KeyT key type
/// @tparam LaneT lane type
/// @tparam OutputBuffer output buffer type
/// @tparam IntervalMemoizer interval memoizer type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the `to`
/// @param outbuf the `outbuf`
/// @param memoizer the memoizer built for this key
/// @return the evaluation result
template <std::size_t N, std::size_t I = 0, typename KeyT, typename LaneT,
typename OutputBuffer, typename IntervalMemoizer>
auto eval_out_interval_impl(
const KeyT & dpf,
LaneT from, LaneT to, OutputBuffer && outbuf, IntervalMemoizer && memoizer)
{
using key_type = KeyT;
static_assert(key_type::meta[I].prefix == N,
"out interval eval: N does not match output I");
using integral_type = typename key_type::integral_type;
constexpr auto opl = key_type::template outputs_per_leaf_of<I>;
constexpr auto lg_opl = key_type::template lg_outputs_per_leaf_of<I>;
constexpr auto to_int = utils::to_integral_type<LaneT>{};
utils::flip_msb_if_signed_integral(from);
utils::flip_msb_if_signed_integral(to);
const auto from_i = static_cast<integral_type>(to_int(from));
const auto to_i = static_cast<integral_type>(to_int(to));
integral_type from_node = utils::leaf_node_floor(from_i, lg_opl);
integral_type to_node = utils::leaf_node_ceil_exclusive(to_i, lg_opl);
const bool wraps = utils::interval_wraps(from_i, to_i, N);
const auto segs = utils::split_leaf_nodes(from_node, to_node,
key_type::meta[I].tree_level, wraps);
std::size_t start = 0;
for (std::size_t s = 0; s < segs.n; ++s)
{
const auto & seg = segs.seg[s];
internal::eval_out_interval_interior<N, I>(dpf, seg.from_node,
seg.to_node, memoizer);
internal::eval_out_interval_exterior<N, I>(dpf, seg.from_node,
seg.to_node, outbuf, memoizer, start);
start += seg.count;
}
constexpr auto mod_pow_2 = utils::mod_pow_2<LaneT>{};
auto from_bits = to_int(from);
auto span = to_int(to) - from_bits;
if constexpr (N < utils::bitlength_of_v<decltype(span)>)
span &= (decltype(span){1} << N) - 1;
auto from_sz = static_cast<std::size_t>(from_bits);
auto to_sz = from_sz + static_cast<std::size_t>(span);
return subinterval_iterable(std::begin(outbuf), utils::size(outbuf),
from_sz, to_sz, mod_pow_2(from, lg_opl), opl);
}
template <std::size_t N, std::size_t I = 0, typename KeyT, typename LaneT,
typename OutputBuffer>
auto eval_out_interval_impl(
const KeyT & dpf,
LaneT from, LaneT to, OutputBuffer && outbuf)
{
using key_type = KeyT;
constexpr auto L = key_type::meta[I].tree_level;
using integral_type = typename key_type::integral_type;
constexpr auto to_int = utils::to_integral_type<LaneT>{};
LaneT f = from, t = to;
utils::flip_msb_if_signed_integral(f);
utils::flip_msb_if_signed_integral(t);
constexpr auto lg = key_type::template lg_outputs_per_leaf_of<I>;
const auto from_i = static_cast<integral_type>(to_int(f));
const auto to_i = static_cast<integral_type>(to_int(t));
const integral_type from_node = utils::leaf_node_floor(from_i, lg);
const integral_type to_node = utils::leaf_node_ceil_exclusive(to_i, lg);
const bool wraps = utils::interval_wraps(from_i, to_i, N);
const auto segs = utils::split_leaf_nodes(from_node, to_node, L, wraps);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto memo = basic_interval_memoizer_at<key_type, L>(segs.total);
HEDLEY_PRAGMA(GCC diagnostic pop)
return eval_out_interval_impl<N, I>(dpf, from, to, outbuf, memo);
}
template <std::size_t N, std::size_t I = 0, typename KeyT, typename LaneT>
auto eval_out_interval_impl(
const KeyT & dpf,
LaneT from, LaneT to)
{
auto buf = make_output_buffer_for_out_interval_impl<N, I>(dpf, from, to);
auto it = eval_out_interval_impl<N, I>(dpf, from, to, buf);
return std::make_pair(std::move(buf), std::move(it));
}
/// @brief Full N-bit lane-domain eval of output `I`.
/// @tparam N width in bits
/// @tparam I output index
/// @tparam KeyT key type
/// @tparam OutputBuffer output buffer type
/// @tparam IntervalMemoizer interval memoizer type
/// @param dpf the DPF key
/// @param outbuf the `outbuf`
/// @param memoizer the memoizer built for this key
/// @return Full N-bit lane-domain eval of output `I`
template <std::size_t N, std::size_t I = 0, typename KeyT, typename OutputBuffer,
typename IntervalMemoizer>
auto eval_out_full_impl(
const KeyT & dpf,
OutputBuffer && outbuf, IntervalMemoizer && memoizer)
{
using lane_t = typename KeyT::input_type;
constexpr lane_t lo = 0;
constexpr lane_t hi = (N >= utils::bitlength_of_v<lane_t>)
? static_cast<lane_t>(~lane_t{0})
: static_cast<lane_t>((lane_t{1} << N) - 1);
return eval_out_interval_impl<N, I>(dpf, lo, hi, outbuf, memoizer);
}
template <std::size_t N, std::size_t I = 0, typename KeyT>
auto eval_out_full_impl(
const KeyT & dpf)
{
using lane_t = typename KeyT::input_type;
constexpr lane_t lo = 0;
constexpr lane_t hi = (N >= utils::bitlength_of_v<lane_t>)
? static_cast<lane_t>(~lane_t{0})
: static_cast<lane_t>((lane_t{1} << N) - 1);
return eval_out_interval_impl<N, I>(dpf, lo, hi);
}
} // namespace incr
} // namespace detail
/// @brief Deepest-group full-domain eval (first cut of incremental `eval_full`).
/// @tparam KeyT key type
/// @tparam KeyT key type
/// @param dpf the DPF key
/// @return Deepest-group full-domain eval (first cut of incremental `eval_full`)
template <typename KeyT,
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
auto eval_full(
const KeyT & dpf)
{
using key_type = KeyT;
constexpr auto I = key_type::deepest_output;
return detail::incr::eval_out_full_impl<key_type::meta[I].prefix, I>(dpf);
}
namespace detail
{
namespace incr
{
/// @brief Sequence eval over lane points for output `I` at prefix `N`.
/// @tparam N width in bits
/// @tparam I output index
/// @tparam KeyT key type
/// @tparam ForwardIterator forward iterator type
/// @tparam OutputBuffer output buffer type
/// @tparam PathMemoizer path memoizer type
/// @param dpf the DPF key
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param outbuf the `outbuf`
/// @param path the root-to-leaf path
/// @return Sequence eval over lane points for output `I` at prefix `N`
template <std::size_t N, std::size_t I = 0, typename KeyT,
typename ForwardIterator, typename OutputBuffer,
typename PathMemoizer = dpf::basic_path_memoizer<KeyT>>
auto eval_out_sequence_impl(
const KeyT & dpf,
ForwardIterator begin, ForwardIterator end, OutputBuffer && outbuf,
PathMemoizer && path = PathMemoizer{})
{
using key_type = KeyT;
static_assert(key_type::meta[I].prefix == N, "out sequence eval: prefix");
constexpr auto opl = key_type::template outputs_per_leaf_of<I>;
using output_type = typename key_type::template concrete_output_type<I>;
std::size_t i = 0;
for (auto it = begin; it != end; ++it, ++i)
{
// Build a full-domain query whose top-N bits equal the lane point.
auto lane = static_cast<typename key_type::input_type>(*it);
typename key_type::input_type q = lane;
if constexpr (N < key_type::input_bits)
q = static_cast<typename key_type::input_type>(
lane << (key_type::input_bits - N));
auto out = eval_out_point_impl<N, I>(dpf, q, path);
if constexpr (utils::is_packed_subbyte_v<output_type>)
{
store_leaf_bytes(outbuf, i, out.node);
}
else
{
std::memcpy(&outbuf[i * opl], &out.node, sizeof(output_type) * opl);
}
}
return subsequence_iterable<key_type, decltype(std::begin(outbuf)),
ForwardIterator>(std::begin(outbuf), begin, end);
}
} // namespace incr
} // namespace detail
/// @brief Deepest-group sequence eval (first cut of incremental `eval_sequence`).
/// @tparam KeyT key type
/// @tparam ForwardIterator forward iterator type
/// @tparam OutputBuffer output buffer type
/// @tparam KeyT key type
/// @param dpf the DPF key
/// @param begin the iterator to the first query
/// @param end the iterator past the last query
/// @param outbuf the `outbuf`
/// @return Deepest-group sequence eval (first cut of incremental `eval_sequence`)
template <typename KeyT, typename ForwardIterator, typename OutputBuffer,
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
auto eval_sequence(
const KeyT & dpf,
ForwardIterator begin, ForwardIterator end, OutputBuffer && outbuf)
{
using key_type = KeyT;
constexpr auto I = key_type::deepest_output;
return detail::incr::eval_out_sequence_impl<key_type::meta[I].prefix, I>(dpf, begin, end,
outbuf);
}
// ---------------------------------------------------------------------------
// cmp eval: tree path-sum comparison (native key value CWs)
// ---------------------------------------------------------------------------
namespace detail
{
namespace incr
{
template <typename KeyT, typename PathMemoizer>
uint64_t eval_cmp_path_sum(const KeyT & dpf, typename KeyT::input_type tx,
PathMemoizer & path, bool as_prefix = false, std::size_t prefix_len = 0)
{
using namespace detail::dcf_impl;
const auto & ch = dpf.cmp();
const uint64_t mask = ch.mask;
const uint64_t add = [&]() -> uint64_t {
if constexpr (is_party_key_v<KeyT>)
return dpf.cmp_addend().raw();
else
return dpf.cmp_addend();
}();
if constexpr (unwrap_party_key_t<KeyT>::cmp_block > 0)
{
if (as_prefix)
throw std::invalid_argument(
"cmp_prefix is not defined for a blocked comparison");
return detail::blocked::eval_share(dpf, tx, path);
}
const std::size_t full_bits = static_cast<std::size_t>(ch.nbits);
const std::size_t nbits = as_prefix ? prefix_len : full_bits;
if ((ch.trivial == cmp_trivial::always_true
|| ch.trivial == cmp_trivial::always_false)
&& (!as_prefix || prefix_len == full_bits))
return add;
dpf::detail::ensure_level(dpf, tx, path, nbits);
const int party = dpf::get_lo_bit(dpf.root()) ? 1 : 0;
uint64_t V = 0;
auto bit_mask = KeyT::msb_mask;
for (std::size_t i = 0; i < nbits; ++i, bit_mask >>= 1)
{
const bool xi = !!(bit_mask & tx);
const auto & parent = path[i];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(parent));
auto kids = KeyT::tree::expand_value(parent);
const uint64_t v = convert_node(kids[xi ? 1 : 0], mask);
const uint64_t contrib =
(v + (t ? dpf.value_cw(i) : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
}
const auto & leaf = path[nbits];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(leaf));
const uint64_t c = convert_node(leaf, mask);
const uint64_t last = as_prefix ? dpf.prefix_cw(nbits) : dpf.cw_last();
const uint64_t contrib = (c + (t ? last : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
if (ch.eval_as_ge)
V = neg_m(V, mask);
return (V + add) & mask;
}
/// @brief Full-tree interval memoizer stopped at `StopLevel` (retains every level;
/// unlike `basic_interval_memoizer_at` which ping-pongs two buffers).
/// @tparam DpfKey DPF key type
/// @tparam StopLevel stop level
/// @tparam interior_node interior node
template <typename DpfKey, std::size_t StopLevel,
typename Allocator = aligned_allocator<
typename unwrap_party_key_t<DpfKey>::interior_node>>
struct cmp_full_interval_memo
{
using dpf_type = unwrap_party_key_t<DpfKey>;
using integral_type = typename dpf_type::integral_type;
using node_type = typename dpf_type::interior_node;
using return_type = node_type *;
using unique_ptr = typename Allocator::unique_ptr;
static constexpr std::size_t depth = StopLevel;
static constexpr bool retains_all_levels = true;
explicit cmp_full_interval_memo(std::size_t output_len,
Allocator alloc = Allocator{})
: output_length{output_len},
level_index{0},
level_endpoints{initialize_endpoints(output_len)},
buf{alloc.allocate_unique_ptr(level_endpoints[depth] + output_len)},
from_{std::nullopt},
to_{std::nullopt}
{
if (HEDLEY_UNLIKELY(buf == nullptr)) throw std::bad_alloc{};
}
std::size_t assign_interval(const dpf_type & dpf, integral_type new_from,
integral_type new_to)
{
static constexpr auto complement_of = std::bit_not{};
if (from_.has_value() == false
|| std::memcmp(&dpf_root_, &dpf.root(), sizeof(node_type)) != 0
|| std::memcmp(&dpf_common_part_hash_, &dpf.common_part_hash(),
sizeof(digest_type)) != 0
|| from_.value_or(complement_of(new_from)) != new_from
|| to_.value_or(complement_of(new_to)) != new_to)
{
if (new_to - new_from > output_length)
throw std::length_error("size of new interval is too large for memoizer");
(*this)[0][0] = dpf.root();
dpf_root_ = dpf.root();
dpf_common_part_hash_ = dpf.common_part_hash();
from_ = new_from;
to_ = new_to;
level_index = 1;
}
return level_index;
}
std::size_t advance_level() { return ++level_index; }
std::size_t get_nodes_at_level() const
{
return get_nodes_at_level(level_index, from_.value_or(0), to_.value_or(0));
}
std::size_t get_nodes_at_level(std::size_t level) const
{
return get_nodes_at_level(level, from_.value_or(0), to_.value_or(0));
}
static std::size_t get_nodes_at_level(std::size_t level, integral_type from_node,
integral_type to_node)
{
std::size_t offset = depth - level;
return utils::shift_right(to_node - integral_type{1}, offset)
- utils::shift_right(from_node, offset) + 1;
}
HEDLEY_NO_THROW
return_type operator[](std::size_t level) const noexcept
{
return Allocator::assume_aligned(&buf[level_endpoints[level]]);
}
private:
std::size_t output_length;
std::size_t level_index;
const std::array<std::size_t, depth + 1> level_endpoints;
unique_ptr buf;
node_type dpf_root_;
digest_type dpf_common_part_hash_;
std::optional<integral_type> from_;
std::optional<integral_type> to_;
static constexpr auto initialize_endpoints(std::size_t len_in)
{
std::array<std::size_t, depth + 1> eps{};
auto len = len_in;
for (std::size_t level = depth; level > 0; --level)
{
len = std::min((len + 2) >> 1, std::size_t{1} << (level - 1));
eps[level] = len;
}
for (std::size_t level = 0; level < depth; ++level)
eps[level + 1] = eps[level] + eps[level + 1];
return eps;
}
};
/// @brief Expand interval interior nodes for the comparison prefix (stop = nbits).
/// @tparam KeyT key type
/// @tparam IntegralT integral type
/// @tparam IntervalMemoizer interval memoizer type
/// @param dpf the DPF key
/// @param from_node the `from_node`
/// @param to_node the `to_node`
/// @param nbits the width in bits
/// @param memoizer the memoizer built for this key
/// @param tree_levels the `tree_levels`
template <typename KeyT, typename IntegralT, typename IntervalMemoizer>
void eval_cmp_interval_impl_interior(const KeyT & dpf, IntegralT from_node,
IntegralT to_node, std::size_t nbits, IntervalMemoizer & memoizer,
std::size_t tree_levels = static_cast<std::size_t>(-1))
{
if (tree_levels == static_cast<std::size_t>(-1))
tree_levels = nbits;
using node_type = typename KeyT::interior_node;
using input_type = typename KeyT::input_type;
const input_type lane_msb =
static_cast<input_type>(input_type{1} << (nbits - 1));
std::size_t level_index = memoizer.assign_interval(dpf, from_node, to_node);
std::size_t nodes_at_level = memoizer.get_nodes_at_level();
IntegralT mask = static_cast<IntegralT>(
utils::to_integral_type<input_type>{}(lane_msb) >> (level_index - 1));
for (; level_index <= tree_levels;
level_index = memoizer.advance_level(),
nodes_at_level = memoizer.get_nodes_at_level(), mask >>= 1)
{
std::size_t i = 0, j = 0;
bool from_offset = mask & from_node,
to_offset = from_offset ^ (nodes_at_level & 1);
const node_type cw[2] = {
dpf.correction_word(level_index - 1, 0),
dpf.correction_word(level_index - 1, 1)};
const bool is_last = KeyT::tree::is_last_level(level_index - 1,
KeyT::depth);
auto *prev = memoizer[level_index - 1];
auto *curr = memoizer[level_index];
if (from_offset == true)
curr[i++] = KeyT::traverse_interior(prev[j++], cw[1], 1, is_last);
const std::size_t both_end = nodes_at_level - to_offset;
for (; i < both_end;)
{
auto cur_node = prev[j++];
auto kids = KeyT::traverse_interior01(cur_node, cw[0], cw[1],
is_last);
curr[i++] = kids[0];
curr[i++] = kids[1];
}
if (to_offset == true)
curr[i] = KeyT::traverse_interior(prev[j], cw[0], 0, is_last);
}
}
template <typename KeyT, typename IntervalMemoizer>
uint64_t eval_cmp_from_interval_memo(const KeyT & dpf,
typename KeyT::integral_type lane, typename KeyT::integral_type from_lane,
std::size_t nbits, IntervalMemoizer & memo)
{
using namespace detail::dcf_impl;
const auto & ch = dpf.cmp();
const uint64_t mask = ch.mask;
const uint64_t add = [&]() -> uint64_t {
if constexpr (is_party_key_v<KeyT>)
return dpf.cmp_addend().raw();
else
return dpf.cmp_addend();
}();
if (ch.trivial == cmp_trivial::always_true
|| ch.trivial == cmp_trivial::always_false)
return add;
const int party = dpf::get_lo_bit(dpf.root()) ? 1 : 0;
uint64_t V = 0;
for (std::size_t i = 0; i < nbits; ++i)
{
const std::size_t shift = nbits - i;
const auto from_i = from_lane >> shift;
const auto q_i = lane >> shift;
const std::size_t idx = static_cast<std::size_t>(q_i - from_i);
const auto & parent = memo[i][idx];
const bool xi = !!(lane & (typename KeyT::integral_type{1}
<< (nbits - 1 - i)));
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(parent));
auto kids = KeyT::tree::expand_value(parent);
const uint64_t v = convert_node(kids[xi ? 1 : 0], mask);
const uint64_t contrib =
(v + (t ? dpf.value_cw(i) : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
}
{
const auto & leaf = memo[nbits][static_cast<std::size_t>(lane - from_lane)];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(leaf));
const uint64_t c = convert_node(leaf, mask);
const uint64_t contrib = (c + (t ? dpf.cw_last() : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
}
if (ch.eval_as_ge)
V = neg_m(V, mask);
return (V + add) & mask;
}
} // namespace incr
} // namespace detail
namespace detail
{
namespace incr
{
template <typename Beta, typename KeyT, typename PathMemoizer>
auto eval_group_path_sum(const KeyT & dpf, typename KeyT::input_type tx,
PathMemoizer & path, bool as_prefix = false, std::size_t prefix_len = 0)
{
using prg = typename unwrap_party_key_t<KeyT>::interior_prg;
using concrete = dpf::concrete_type_t<Beta>;
const auto layout = detail::group_layout<concrete>();
const auto & ch = dpf.cmp();
const auto add = detail::group_from_word(dpf.cmp_addend_word(), layout);
const std::size_t full_bits = static_cast<std::size_t>(ch.nbits);
const std::size_t nbits = as_prefix ? prefix_len : full_bits;
if ((ch.trivial == cmp_trivial::always_true
|| ch.trivial == cmp_trivial::always_false)
&& (!as_prefix || prefix_len == full_bits))
return detail::group_to_beta<concrete>(add);
dpf::detail::ensure_level(dpf, tx, path, nbits);
const int party = dpf::get_lo_bit(dpf.root()) ? 1 : 0;
auto V = detail::group_zero(layout);
const auto zero = detail::group_zero(layout);
auto bit_mask = KeyT::msb_mask;
for (std::size_t i = 0; i < nbits; ++i, bit_mask >>= 1)
{
const bool xi = !!(bit_mask & tx);
const auto & parent = path[i];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(parent));
auto kids = KeyT::tree::expand_value(parent);
const auto v = detail::group_from_node<prg>(kids[xi ? 1 : 0], layout);
const auto cw = detail::group_from_word(dpf.value_cw()[i], layout);
const auto contrib = detail::group_add(v, t ? cw : zero);
V = detail::group_add(V, party ? detail::group_neg(contrib) : contrib);
}
const auto & leaf = path[nbits];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(leaf));
const auto c = detail::group_from_node<prg>(leaf, layout);
const auto last = detail::group_from_word(
as_prefix ? dpf.prefix_cws()[nbits] : dpf.cw_last_word(), layout);
const auto contrib = detail::group_add(c, t ? last : zero);
V = detail::group_add(V, party ? detail::group_neg(contrib) : contrib);
if (ch.eval_as_ge)
V = detail::group_neg(V);
return detail::group_to_beta<concrete>(detail::group_add(V, add));
}
template <typename Beta = uint64_t, typename KeyT, typename QueryT,
typename PathMemoizer = dpf::basic_path_memoizer<KeyT>>
auto eval_cmp_point_impl(const KeyT & dpf, QueryT && x,
PathMemoizer && path = PathMemoizer{})
{
if (!dpf.has_cmp())
throw std::invalid_argument("cmp eval: key has no comparison channel");
if (!dpf.cmp_assigned())
throw std::invalid_argument(
"cmp eval: wildcard comparison payload not assigned (call assign_cmp)");
auto tx = dpf.offset_x(std::forward<QueryT>(x));
utils::flip_msb_if_signed_integral(tx);
if constexpr (detail::cmp_group_info<dpf::concrete_type_t<Beta>>::custom)
{
return make_eval_cmp_result<KeyT>(
detail::incr::eval_group_path_sum<Beta>(dpf, tx, path));
}
else
{
const uint64_t raw =
detail::incr::eval_cmp_path_sum(dpf, tx, path);
return make_eval_cmp_result<KeyT>(
detail::dcf_impl::u64_to_beta<Beta>(raw));
}
}
template <std::size_t L, typename Beta = uint64_t, typename KeyT, typename QueryT,
typename PathMemoizer = dpf::basic_path_memoizer<KeyT>>
auto eval_cmp_prefix_point_impl(const KeyT & dpf, QueryT && x,
PathMemoizer && path = PathMemoizer{})
{
static_assert(unwrap_party_key_t<KeyT>::cmp_idcf,
"cmp_prefix requires an idcf key");
if (!dpf.has_cmp())
throw std::invalid_argument("cmp_prefix: key has no comparison channel");
if (!dpf.cmp().incremental)
throw std::invalid_argument("cmp_prefix: key was not built with idcf");
if (!dpf.cmp_assigned())
throw std::invalid_argument(
"cmp_prefix: wildcard comparison payload not assigned (call assign_cmp)");
if (L > static_cast<std::size_t>(dpf.cmp().nbits))
throw std::invalid_argument("cmp_prefix: prefix is longer than the comparison");
auto tx = dpf.offset_x(std::forward<QueryT>(x));
utils::flip_msb_if_signed_integral(tx);
if constexpr (detail::cmp_group_info<dpf::concrete_type_t<Beta>>::custom)
{
return make_eval_cmp_result<KeyT>(
detail::incr::eval_group_path_sum<Beta>(dpf, tx, path, true, L));
}
else
{
const uint64_t raw =
detail::incr::eval_cmp_path_sum(dpf, tx, path, true, L);
return make_eval_cmp_result<KeyT>(
detail::dcf_impl::u64_to_beta<Beta>(raw));
}
}
// ---------------------------------------------------------------------------
// Buffered / interval / sequence comparison evals
// ---------------------------------------------------------------------------
template <typename Beta = uint64_t, typename KeyT>
auto make_output_buffer_for_cmp_impl(const KeyT &, std::size_t n)
{
return dpf::output_buffer<cmp_buffer_elem_t<KeyT, Beta>>(n);
}
template <typename Integral>
std::size_t cmp_inclusive_count(Integral a, Integral b)
{
if (b < a)
throw std::invalid_argument("cmp interval: to < from");
const Integral one{1};
if (a == Integral{0} && static_cast<Integral>(b + one) < b)
throw std::length_error("cmp interval does not fit in size_t");
return static_cast<std::size_t>(b - a + one);
}
template <typename Integral>
Integral cmp_exclusive_end(Integral b)
{
const Integral next = static_cast<Integral>(b + Integral{1});
// Overflow: the exclusive end is 2^width, encoded as 0 for the walk.
if (next < b)
return Integral{0};
return next;
}
template <typename Memoizer, typename = void>
struct interval_memoizer_retains_levels : std::false_type {};
template <typename Memoizer>
struct interval_memoizer_retains_levels<Memoizer,
std::void_t<decltype(Memoizer::retains_all_levels)>>
: std::bool_constant<Memoizer::retains_all_levels> {};
template <typename Beta = uint64_t, typename KeyT, typename LaneT>
auto make_output_buffer_for_cmp_interval_impl(const KeyT & dpf, LaneT from, LaneT to)
{
if (!dpf.has_cmp())
throw std::invalid_argument("make_output_buffer(cmp): no cmp");
using integral = typename KeyT::integral_type;
constexpr auto to_int = utils::to_integral_type<LaneT>{};
utils::flip_msb_if_signed_integral(from);
utils::flip_msb_if_signed_integral(to);
const auto a = static_cast<integral>(to_int(from));
const auto b = static_cast<integral>(to_int(to));
return dpf::output_buffer<cmp_buffer_elem_t<KeyT, Beta>>(
cmp_inclusive_count(a, b));
}
template <typename Beta = uint64_t, typename KeyT, typename ForwardIterator,
typename OutputBuffer,
typename PathMemoizer = dpf::basic_path_memoizer<KeyT>>
void eval_cmp_sequence_impl(const KeyT & dpf, ForwardIterator begin,
ForwardIterator end, OutputBuffer && outbuf,
PathMemoizer && path = PathMemoizer{})
{
if (!dpf.has_cmp())
throw std::invalid_argument("cmp sequence eval: no comparison channel");
std::size_t i = 0;
for (auto it = begin; it != end; ++it, ++i)
outbuf[i] = eval_cmp_point_impl<Beta>(dpf, *it, path);
}
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
typename OutputBuffer>
void eval_cmp_interval_impl(const KeyT & dpf, LaneT from, LaneT to,
OutputBuffer && outbuf)
{
if (!dpf.has_cmp())
throw std::invalid_argument("cmp interval eval: no comparison channel");
if (!dpf.cmp_assigned())
throw std::invalid_argument(
"cmp interval eval: wildcard payload not assigned (call assign_cmp)");
if constexpr (detail::cmp_group_info<dpf::concrete_type_t<Beta>>::custom)
{
std::size_t i = 0;
dpf::basic_path_memoizer<KeyT> path;
for (LaneT q = from; ; ++q, ++i)
{
outbuf[i] = eval_cmp_point_impl<Beta>(dpf, q, path);
if (q == to)
break;
}
return;
}
constexpr auto to_int = utils::to_integral_type<LaneT>{};
utils::flip_msb_if_signed_integral(from);
utils::flip_msb_if_signed_integral(to);
const auto nbits = static_cast<std::size_t>(dpf.cmp().nbits);
using integral = typename KeyT::integral_type;
const auto a = static_cast<integral>(to_int(from));
const auto b = static_cast<integral>(to_int(to));
const auto count = cmp_inclusive_count(a, b);
// Full-tree interval memoizer retains every level (ping-pong basic_* does not),
// so path-sum can read ancestors directly.
constexpr std::size_t stop =
KeyT::cmp_depth == 0 ? KeyT::depth : KeyT::cmp_depth;
// Rebind key depth for the memoizer by using a stop-level wrapper built on
// the same assign/advance API as basic_interval_memoizer_at, but allocate
// per-level storage. Stop stays the logical comparison width so prefix
// indexes match `nbits`, even when a blocked key's seed spine is shorter.
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
detail::incr::cmp_full_interval_memo<KeyT, stop> memo{count};
HEDLEY_PRAGMA(GCC diagnostic pop)
const std::size_t levels = unwrap_party_key_t<KeyT>::cmp_block > 0
? unwrap_party_key_t<KeyT>::cmp_h : nbits;
detail::incr::eval_cmp_interval_impl_interior(dpf, a, cmp_exclusive_end(b),
nbits, memo, levels);
for (std::size_t i = 0; i < count; ++i)
{
const auto q = static_cast<integral>(a + static_cast<integral>(i));
const uint64_t raw = [&] {
if constexpr (unwrap_party_key_t<KeyT>::cmp_block > 0)
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return detail::blocked::eval_share_memo(dpf, q, a,
cmp_exclusive_end(b), memo);
HEDLEY_PRAGMA(GCC diagnostic pop)
}
else
{
return detail::incr::eval_cmp_from_interval_memo(
dpf, q, a, nbits, memo);
}
}();
outbuf[i] = make_eval_cmp_result<KeyT>(
detail::dcf_impl::u64_to_beta<Beta>(raw));
}
}
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
typename OutputBuffer, typename IntervalMemoizer>
void eval_cmp_interval_impl(const KeyT & dpf, LaneT from, LaneT to,
OutputBuffer && outbuf, IntervalMemoizer && memo)
{
if (!dpf.has_cmp())
throw std::invalid_argument("cmp interval eval: no comparison channel");
if (!dpf.cmp_assigned())
throw std::invalid_argument(
"cmp interval eval: wildcard payload not assigned (call assign_cmp)");
if constexpr (detail::cmp_group_info<dpf::concrete_type_t<Beta>>::custom)
{
(void)memo;
std::size_t i = 0;
dpf::basic_path_memoizer<KeyT> path;
for (LaneT q = from; ; ++q, ++i)
{
outbuf[i] = eval_cmp_point_impl<Beta>(dpf, q, path);
if (q == to)
break;
}
return;
}
constexpr auto to_int = utils::to_integral_type<LaneT>{};
utils::flip_msb_if_signed_integral(from);
utils::flip_msb_if_signed_integral(to);
const auto nbits = static_cast<std::size_t>(dpf.cmp().nbits);
using integral = typename KeyT::integral_type;
static_assert(interval_memoizer_retains_levels<std::decay_t<IntervalMemoizer>>::value,
"eval_cmp_interval_impl memoizer must set retains_all_levels; a ping-pong memoizer drops ancestor levels");
const auto a = static_cast<integral>(to_int(from));
const auto b = static_cast<integral>(to_int(to));
const auto count = cmp_inclusive_count(a, b);
const std::size_t levels = unwrap_party_key_t<KeyT>::cmp_block > 0
? unwrap_party_key_t<KeyT>::cmp_h : nbits;
detail::incr::eval_cmp_interval_impl_interior(dpf, a, cmp_exclusive_end(b),
nbits, memo, levels);
for (std::size_t i = 0; i < count; ++i)
{
const auto q = static_cast<integral>(a + static_cast<integral>(i));
const uint64_t raw = [&] {
if constexpr (unwrap_party_key_t<KeyT>::cmp_block > 0)
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
return detail::blocked::eval_share_memo(dpf, q, a,
cmp_exclusive_end(b), memo);
HEDLEY_PRAGMA(GCC diagnostic pop)
}
else
{
return detail::incr::eval_cmp_from_interval_memo(
dpf, q, a, nbits, memo);
}
}();
outbuf[i] = make_eval_cmp_result<KeyT>(
detail::dcf_impl::u64_to_beta<Beta>(raw));
}
}
template <typename Beta = uint64_t, typename KeyT, typename LaneT>
auto eval_cmp_interval_impl(const KeyT & dpf, LaneT from, LaneT to)
{
auto buf = make_output_buffer_for_cmp_interval_impl<Beta>(dpf, from, to);
eval_cmp_interval_impl<Beta>(dpf, from, to, buf);
return buf;
}
} // namespace incr
} // namespace detail
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_INCREMENTAL_HPP__