2667 lines
105 KiB
C++
2667 lines
105 KiB
C++
/// @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_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...>());
|
||
}
|
||
|
||
/// 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...>();
|
||
|
||
/// 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`.
|
||
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);
|
||
}
|
||
|
||
/// 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;
|
||
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_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;
|
||
constexpr auto bits = [] {
|
||
if constexpr (std::is_same_v<Beta, dpf::bit>)
|
||
return std::size_t{1};
|
||
else
|
||
return utils::bitlength_of_v<Beta>;
|
||
}();
|
||
spec.mask = detail::dcf_impl::default_mask_for_bits(bits);
|
||
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 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,
|
||
std::get<Is>(placed).value...);
|
||
}
|
||
|
||
template <typename ExteriorPRG, typename PlacedTuple, std::size_t... Is>
|
||
auto empty_leaves(std::index_sequence<Is...>)
|
||
{
|
||
using node = typename ExteriorPRG::block_type;
|
||
return std::make_tuple(dpf::leaf_node_t<node,
|
||
typename std::tuple_element_t<Is, PlacedTuple>::output_type>{}...);
|
||
}
|
||
|
||
template <typename ExteriorPRG, typename PlacedTuple, std::size_t... Is>
|
||
auto empty_beavers(std::index_sequence<Is...>)
|
||
{
|
||
using node = typename ExteriorPRG::block_type;
|
||
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>>{}...);
|
||
}
|
||
|
||
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
|
||
typename PlacedTuple, typename MetaHolder, std::size_t G,
|
||
typename Leaves0T, typename Beavers0T, typename Leaves1T,
|
||
typename Beavers1T>
|
||
void gen_group(InputT x, const typename InteriorPRG::block_type & s0,
|
||
const typename InteriorPRG::block_type & s1, bool sign0, PlacedTuple & placed,
|
||
Leaves0T & leaves0, Beavers0T & beavers0, Leaves1T & leaves1,
|
||
Beavers1T & beavers1)
|
||
{
|
||
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);
|
||
|
||
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>{});
|
||
}
|
||
|
||
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, bool sign0, 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, sign0, 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;
|
||
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))...);
|
||
}
|
||
|
||
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>
|
||
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
|
||
{
|
||
|
||
|
||
/// Map kind → flags. Tree keep-path stays on α; leq/gt plant δ on that path
|
||
/// via `include_eq`. Domain-edge α yields trivial always-true/false.
|
||
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;
|
||
}
|
||
|
||
/// 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.
|
||
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;
|
||
}
|
||
|
||
/// Typed overload: write party-0 / party-1 additive shares of the absorb.
|
||
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 = false>
|
||
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>;
|
||
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;
|
||
|
||
utils::flip_msb_if_signed_integral(x);
|
||
|
||
const interior_node root[2] = {dpf::unset_lo_bit(root_sampler()),
|
||
dpf::set_lo_bit(root_sampler())};
|
||
|
||
typename key_type::correction_words_array correction_words{};
|
||
typename key_type::correction_advice_array correction_advice{};
|
||
|
||
interior_node parent[2] = {root[0], root[1]};
|
||
auto mask = key_type::msb_mask;
|
||
|
||
std::array<interior_node, depth + 1> snap0{};
|
||
std::array<interior_node, depth + 1> snap1{};
|
||
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{};
|
||
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");
|
||
}
|
||
|
||
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);
|
||
}
|
||
}
|
||
};
|
||
if constexpr (CmpIdcf)
|
||
{
|
||
if (cmp.active)
|
||
snap_prefix(0);
|
||
}
|
||
|
||
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
|
||
{
|
||
bool bit = !!(mask & x);
|
||
bool advice[2];
|
||
advice[0] = dpf::get_lo_bit_and_clear_lo_2bits(parent[0]);
|
||
advice[1] = dpf::get_lo_bit_and_clear_lo_2bits(parent[1]);
|
||
|
||
auto child0 = InteriorPRG::eval01(parent[0]);
|
||
auto child1 = InteriorPRG::eval01(parent[1]);
|
||
interior_node child[2] = {child0[0] ^ child1[0], child0[1] ^ child1[1]};
|
||
|
||
bool t[2] = {static_cast<bool>(dpf::get_lo_bit(child[0]) ^ !bit),
|
||
static_cast<bool>(dpf::get_lo_bit(child[1]) ^ bit)};
|
||
auto cw = dpf::set_lo_bit(child[!bit], t[bit]);
|
||
parent[0] = dpf::xor_if(child0[bit], cw, advice[0]);
|
||
parent[1] = dpf::xor_if(child1[bit], cw, advice[1]);
|
||
|
||
correction_words[level] = child[!bit];
|
||
correction_advice[level] =
|
||
static_cast<psnip_uint8_t>(t[1] << 1) | t[0];
|
||
|
||
if constexpr (CmpBlock == 0)
|
||
{
|
||
if (cmp.active && cmp.trivial == cmp_trivial::none
|
||
&& level < cmp_nbits)
|
||
{
|
||
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(child0[0],
|
||
child0[1], child1[0], child1[1],
|
||
static_cast<uint8_t>(advice[0]),
|
||
static_cast<uint8_t>(advice[1]), 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(
|
||
child0[0], child0[1], child1[0], child1[1],
|
||
static_cast<uint8_t>(advice[0]),
|
||
static_cast<uint8_t>(advice[1]), 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(child0[0], child0[1],
|
||
child1[0], child1[1],
|
||
static_cast<uint8_t>(advice[0]),
|
||
static_cast<uint8_t>(advice[1]), 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(child0[0], child0[1],
|
||
child1[0], child1[1],
|
||
static_cast<uint8_t>(advice[0]),
|
||
static_cast<uint8_t>(advice[1]), 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)
|
||
{
|
||
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)
|
||
{
|
||
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;
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
constexpr std::size_t ngroups = [] {
|
||
std::size_t m = 0;
|
||
for (std::size_t i = 0; i < n; ++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);
|
||
|
||
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 < n; ++i)
|
||
{
|
||
if (key_type::meta[i].group_id == G)
|
||
return key_type::meta[i].tree_level;
|
||
}
|
||
return std::size_t{0};
|
||
}();
|
||
gen_group<InteriorPRG, ExteriorPRG, InputT, PlacedTuple, MetaHolder, G>(
|
||
x, dpf::unset_lo_2bits(snap0[lvl]), dpf::unset_lo_2bits(snap1[lvl]),
|
||
snap_sign[lvl], placed, leaves0, beavers0, leaves1, beavers1);
|
||
});
|
||
|
||
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;
|
||
if (cmp.active)
|
||
{
|
||
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);
|
||
}
|
||
|
||
input_type off0{}, off1{};
|
||
auto adds = extract_addends(placed, std::make_index_sequence<n>{});
|
||
return dpf::make_party_key_pair(
|
||
key_type{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},
|
||
key_type{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});
|
||
}
|
||
|
||
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 = false,
|
||
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>;
|
||
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);
|
||
|
||
const interior_node root0 =
|
||
dpf::unset_lo_bit(static_cast<interior_node>(root_sampler()));
|
||
const interior_node root1 =
|
||
dpf::set_lo_bit(static_cast<interior_node>(root_sampler()));
|
||
|
||
ds_gen_state<interior_node> st;
|
||
st.init(root0, root1);
|
||
|
||
typename key_type::correction_words_array correction_words{};
|
||
typename key_type::correction_advice_array correction_advice{};
|
||
|
||
std::array<interior_node, depth + 1> snap0{};
|
||
std::array<interior_node, depth + 1> snap1{};
|
||
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");
|
||
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;
|
||
}
|
||
}
|
||
|
||
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;
|
||
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);
|
||
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);
|
||
}
|
||
}
|
||
};
|
||
if constexpr (CmpIdcf)
|
||
{
|
||
if (cmp.active)
|
||
snap_prefix(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)
|
||
{
|
||
ds_advance_level<InteriorPRG>(st, x0, x1, mask, level, proto,
|
||
correction_words[level], correction_advice[level],
|
||
cmp_st.active ? &vcw : nullptr,
|
||
cmp_st.active ? &cmp_st : nullptr);
|
||
if (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
|
||
{
|
||
ds_advance_level<InteriorPRG>(st, x0, x1, mask, level, proto,
|
||
correction_words[level], correction_advice[level]);
|
||
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)
|
||
{
|
||
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;
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
constexpr std::size_t ngroups = [] {
|
||
std::size_t m = 0;
|
||
for (std::size_t i = 0; i < n; ++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`.
|
||
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 < n; ++i)
|
||
{
|
||
if (key_type::meta[i].group_id == G)
|
||
return key_type::meta[i].tree_level;
|
||
}
|
||
return std::size_t{0};
|
||
}();
|
||
gen_group<InteriorPRG, ExteriorPRG, InputT, PlacedTuple, MetaHolder,
|
||
G>(x, dpf::unset_lo_2bits(snap0[lvl]),
|
||
dpf::unset_lo_2bits(snap1[lvl]), snap_sign[lvl], placed,
|
||
leaves0, beavers0, leaves1, beavers1);
|
||
});
|
||
});
|
||
|
||
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;
|
||
if (cmp.active)
|
||
{
|
||
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);
|
||
}
|
||
|
||
input_type off0{}, off1{};
|
||
auto adds = extract_addends(placed, std::make_index_sequence<n>{});
|
||
return dpf::make_party_key_pair(
|
||
key_type{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},
|
||
key_type{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});
|
||
}
|
||
|
||
|
||
} // 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...>;
|
||
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>(x, std::move(placed),
|
||
dpf::uniform_sample<node>, cs);
|
||
}
|
||
else if constexpr (!args_have_cmp_v<OutputTs...> && !args_have_eq_v<OutputTs...>
|
||
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>())
|
||
{
|
||
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>(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...>;
|
||
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>(x, std::move(placed), seed, cs);
|
||
}
|
||
else if constexpr (!args_have_cmp_v<OutputTs...> && !args_have_eq_v<OutputTs...>
|
||
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>())
|
||
{
|
||
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>(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...>;
|
||
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>(x, std::move(placed), root_sampler,
|
||
cs);
|
||
}
|
||
else if constexpr (!args_have_cmp_v<OutputTs...> && !args_have_eq_v<OutputTs...>
|
||
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>())
|
||
{
|
||
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>(x, std::move(placed), root_sampler,
|
||
cs);
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// make_dpf_doerner_shelat(x0, x1, ...): classic or incremental
|
||
// ---------------------------------------------------------------------------
|
||
|
||
/// Doerner–Shelat keygen with caller-supplied roots and pad stream.
|
||
/// 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`.
|
||
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)...);
|
||
}
|
||
|
||
/// 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)...);
|
||
}
|
||
|
||
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...>;
|
||
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>(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...>
|
||
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>())
|
||
{
|
||
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>(arith, std::move(x0),
|
||
std::move(x1), rng.root, proto, std::move(placed), cs);
|
||
}
|
||
}
|
||
|
||
/// Doerner–Shelat with an injectable `CwProtocol` (local or MPC backend).
|
||
/// 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.
|
||
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...>;
|
||
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>(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_eq_v<OutputT, OutputTs...>
|
||
&& detail::incr::is_classic_placed<node, placed_tuple, bitlen>())
|
||
{
|
||
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>(arith, std::move(x0),
|
||
std::move(x1), root_sampler, proto, std::move(placed), cs);
|
||
}
|
||
}
|
||
|
||
/// Doerner–Shelat keygen. Roots and pads come from `uniform_sample`.
|
||
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;
|
||
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
|
||
dpf::uniform_sample<block>, {}};
|
||
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;
|
||
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
|
||
dpf::uniform_sample<block>, {}};
|
||
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||
arith_input, std::move(x0), std::move(x1), rng, std::forward<OutputT>(y),
|
||
std::forward<OutputTs>(ys)...);
|
||
}
|
||
|
||
/// 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
|
||
{
|
||
|
||
/// 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
|
||
|
||
/// 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.
|
||
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();
|
||
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);
|
||
}
|
||
|
||
/// Party-tagged overload: `make_dpf` returns distinct `party_key<0>` /
|
||
/// `party_key<1>` types, so the same-type pair overload cannot bind both.
|
||
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);
|
||
}
|
||
|
||
/// 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)`.
|
||
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);
|
||
}
|
||
|
||
/// Share-typed overload: subtractive shares are converted with the party
|
||
/// coefficient before the existing additive leaf absorb math.
|
||
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
|
||
{
|
||
|
||
/// Evaluate output slot `I` of an incremental key at the programmed point.
|
||
/// `N` must match the prefix of that slot (`at<N>` or full input width).
|
||
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
|
||
|
||
/// Evaluate the first deepest-prefix output (plan default for `eval_point`).
|
||
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
|
||
|
||
/// 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
|
||
{
|
||
|
||
/// 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)};
|
||
|
||
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);
|
||
}
|
||
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);
|
||
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]);
|
||
curr[i++] = kids[0];
|
||
curr[i++] = kids[1];
|
||
}
|
||
if (to_offset == true)
|
||
{
|
||
curr[i] = dpf_type::traverse_interior(prev[j], cw[0], 0);
|
||
}
|
||
}
|
||
}
|
||
|
||
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
|
||
|
||
/// Evaluate output `I` over an interval in the N-bit lane subdomain.
|
||
/// `from`/`to` are lane values in `[0, 2^N)` (not the full input domain).
|
||
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);
|
||
auto memo = basic_interval_memoizer_at<key_type, L>(segs.total);
|
||
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));
|
||
}
|
||
|
||
/// 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
|
||
|
||
/// 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
|
||
{
|
||
|
||
/// 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
|
||
|
||
/// 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::interior_prg::eval01(dpf::unset_lo_2bits(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;
|
||
}
|
||
|
||
/// Full-tree interval memoizer stopped at `StopLevel` (retains every level;
|
||
/// unlike `basic_interval_memoizer_at` which ping-pongs two buffers).
|
||
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;
|
||
}
|
||
};
|
||
|
||
/// Expand interval interior nodes for the comparison prefix (stop = nbits).
|
||
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)};
|
||
|
||
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);
|
||
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]);
|
||
curr[i++] = kids[0];
|
||
curr[i++] = kids[1];
|
||
}
|
||
if (to_offset == true)
|
||
curr[i] = KeyT::traverse_interior(prev[j], cw[0], 0);
|
||
}
|
||
}
|
||
|
||
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::interior_prg::eval01(dpf::unset_lo_2bits(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 = 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);
|
||
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);
|
||
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)");
|
||
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.
|
||
detail::incr::cmp_full_interval_memo<KeyT, stop> memo{count};
|
||
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)
|
||
{
|
||
return detail::blocked::eval_share_memo(dpf, q, a,
|
||
cmp_exclusive_end(b), memo);
|
||
}
|
||
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)");
|
||
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)
|
||
{
|
||
return detail::blocked::eval_share_memo(dpf, q, a,
|
||
cmp_exclusive_end(b), memo);
|
||
}
|
||
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__
|