/// @file dpf/eval_unified.hpp /// @brief Target-first eval surface for DPF / iDPF / DCF channels. /// @details `eval_*(out, …)` selects point-output slot `I`. /// `eval_*(cmp, …)` selects the comparison channel. Memoizer and /// buffer arguments match the classic overloads: a path memoizer /// on `eval_point`, an output buffer then an interval memoizer on /// `eval_interval`. `make_output_buffer(out, key, from, to)` and /// `make_output_buffer(cmp, key, n)` size the buffer for that channel. /// @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_EVAL_UNIFIED_HPP__ #define LIBDPF_INCLUDE_DPF_EVAL_UNIFIED_HPP__ #include "hedley/hedley.h" #include #include #include #include #include #include #include #include #include #include "dpf/eval_target.hpp" #include "dpf/eval_point.hpp" #include "dpf/eval_interval.hpp" #include "dpf/eval_full.hpp" #include "dpf/eval_sequence.hpp" #include "dpf/sequence_recipe.hpp" #include "dpf/incremental.hpp" #include "dpf/path_memoizer.hpp" #include "dpf/interval_memoizer.hpp" #include "dpf/aligned_allocator.hpp" #include "dpf/leaf_node.hpp" namespace dpf { namespace detail { template HEDLEY_NO_THROW constexpr std::size_t resolved_out_prefix() noexcept { if constexpr (is_multilevel_key_v) { if constexpr (N != prefix_deduce) { static_assert(KeyT::meta[I].prefix == N, "out: N does not match key::meta[I].prefix"); return N; } else return KeyT::meta[I].prefix; } else { (void)N; return utils::bitlength_of_v; } } } // namespace detail // --------------------------------------------------------------------------- // eval_point(target, key, x [, path]) // --------------------------------------------------------------------------- template > auto eval_point(out_t, const KeyT & key, QueryT && x, PathMemoizer && path = PathMemoizer{}) { if constexpr (is_multilevel_key_v) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_point_impl(key, std::forward(x), std::forward(path)); } else { return eval_point(key, std::forward(x), std::forward(path)); } } template > auto eval_point(cmp_t, const KeyT & key, QueryT && x, PathMemoizer && path = PathMemoizer{}) { return detail::incr::eval_cmp_point_impl(key, std::forward(x), std::forward(path)); } template > auto eval_point(cmp_prefix_t, const KeyT & key, QueryT && x, PathMemoizer && path = PathMemoizer{}) { return detail::incr::eval_cmp_prefix_point_impl(key, std::forward(x), std::forward(path)); } // --------------------------------------------------------------------------- // eval_interval(target, key, from, to [, buf [, memo]]) // --------------------------------------------------------------------------- template auto eval_interval(out_t, const KeyT & key, LaneT from, LaneT to, OutputBuffer && outbuf, IntervalMemoizer && memo) { if constexpr (is_multilevel_key_v) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_interval_impl(key, from, to, std::forward(outbuf), std::forward(memo)); } else { return eval_interval(key, from, to, std::forward(outbuf), std::forward(memo)); } } template auto eval_interval(out_t, const KeyT & key, LaneT from, LaneT to, OutputBuffer && outbuf) { if constexpr (is_multilevel_key_v) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_interval_impl(key, from, to, std::forward(outbuf)); } else { return eval_interval(key, from, to, std::forward(outbuf)); } } template auto eval_interval(out_t, const KeyT & key, LaneT from, LaneT to) { if constexpr (is_multilevel_key_v) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_interval_impl(key, from, to); } else { return eval_interval(key, from, to); } } template void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to, OutputBuffer && outbuf) { detail::incr::eval_cmp_interval_impl(key, from, to, std::forward(outbuf)); } template void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to, OutputBuffer && outbuf, IntervalMemoizer && memo) { detail::incr::eval_cmp_interval_impl(key, from, to, std::forward(outbuf), std::forward(memo)); } template auto eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to) { return detail::incr::eval_cmp_interval_impl(key, from, to); } // --------------------------------------------------------------------------- // eval_full(target, key [, …]) // --------------------------------------------------------------------------- template auto eval_full(out_t, const KeyT & key, OutputBuffer && outbuf, IntervalMemoizer && memo) { if constexpr (is_multilevel_key_v) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_full_impl(key, std::forward(outbuf), std::forward(memo)); } else { return eval_full(key, std::forward(outbuf), std::forward(memo)); } } template auto eval_full(out_t, const KeyT & key) { if constexpr (is_multilevel_key_v) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_full_impl(key); } else { return eval_full(key); } } template auto eval_full(cmp_t, const KeyT & key) { if (!key.has_cmp()) throw std::invalid_argument("eval_full(cmp): no comparison channel"); using lane_t = typename KeyT::integral_type; const auto nbits = static_cast(key.cmp().nbits); const lane_t lo = 0; const lane_t hi = (nbits >= 8 * sizeof(lane_t)) ? static_cast(~lane_t{0}) : static_cast((lane_t{1} << nbits) - 1); return detail::incr::eval_cmp_interval_impl(key, lo, hi); } // --------------------------------------------------------------------------- // eval_sequence(target, key, begin, end, buf [, path]) // --------------------------------------------------------------------------- template > auto eval_sequence(out_t, const KeyT & key, ForwardIterator begin, ForwardIterator end, OutputBuffer && outbuf, PathMemoizer && path = PathMemoizer{}) { if constexpr (is_multilevel_key_v) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_sequence_impl(key, begin, end, std::forward(outbuf), std::forward(path)); } else { return eval_sequence(key, begin, end, std::forward(outbuf)); } } template > void eval_sequence(cmp_t, const KeyT & key, ForwardIterator begin, ForwardIterator end, OutputBuffer && outbuf, PathMemoizer && path = PathMemoizer{}) { detail::incr::eval_cmp_sequence_impl(key, begin, end, std::forward(outbuf), std::forward(path)); } // --------------------------------------------------------------------------- // make_output_buffer(target, …) // --------------------------------------------------------------------------- template auto make_output_buffer(cmp_t, const KeyT & key, std::size_t n) { return detail::incr::make_output_buffer_for_cmp_impl(key, n); } template auto make_output_buffer(cmp_t, const KeyT & key, LaneT from, LaneT to) { return detail::incr::make_output_buffer_for_cmp_interval_impl( key, from, to); } template auto make_output_buffer(out_t, const KeyT & key, LaneT from, LaneT to) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::make_output_buffer_for_out_interval_impl( key, from, to); } // --------------------------------------------------------------------------- // eval_inner_product(target, key, from, to, weights [, memo]) // // Point-slot inner product: same interior walk as `eval_interval(out, …)` // but each packed leaf is multiply-accumulated against a public weight vector // instead of being materialized. Additive outputs sum `DPF_I(x)·w[x]`; XOR // outputs (`bit` / `xor_wrapper`) xor `DPF_I(x) & w[x]`. Weights are indexed in // the slot's lane domain, matching `eval_interval`'s destination layout. // // Cmp inner product: dot of the per-point comparison path-sum shares with the // weights (no leaf MAC); the two parties' results reconstruct to the true dot. // --------------------------------------------------------------------------- namespace detail { namespace incr { template struct ml_ip_accum { static constexpr bool xor_mode = std::is_same_v || utils::is_xor_wrapper_v; psnip_uint64_t acc = 0; template void mac(const LeafT & leaf, std::size_t base, std::size_t opl, W && w) { for (std::size_t p = 0; p < opl; ++p) { psnip_uint64_t val; if constexpr (utils::is_packed_subbyte_v) { val = static_cast( dpf::extract_leaf(leaf, p)); } else { OutputT v; std::memcpy(&v, reinterpret_cast(std::addressof(leaf)) + p * sizeof(OutputT), sizeof(v)); if constexpr (utils::is_xor_wrapper_v) { // `static_cast(v)` is ambiguous for // `xor_wrapper` (both `operator bool` and `operator T` // are viable). Go through the concrete underlying bits. val = static_cast(v.data()); } else { val = static_cast(v); } } const auto wt = static_cast(w[base + p]); if constexpr (xor_mode) acc ^= (val & wt); else if constexpr (utils::is_packed_subbyte_v) { constexpr auto mask = (static_cast(1) << utils::packed_lane_bits_v) - 1; acc = (acc + (val & mask) * (wt & mask)) & mask; } else acc += val * wt; } } OutputT finish() const { if constexpr (std::is_same_v) return OutputT{static_cast(acc & 1)}; else if constexpr (utils::is_packed_subbyte_v) return static_cast(acc); else if constexpr (utils::is_xor_wrapper_v) return OutputT{static_cast(acc)}; else return static_cast(acc); } }; template auto eval_out_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to, Weights && weights, IntervalMemoizer && memoizer) { using key_type = KeyT; static_assert(key_type::meta[I].prefix == N, "out inner product: N does not match output I"); using output_type = typename key_type::template concrete_output_type; using exterior_node = typename key_type::exterior_node; using integral_type = typename key_type::integral_type; constexpr auto opl = key_type::template outputs_per_leaf_of; constexpr auto lg_opl = key_type::template lg_outputs_per_leaf_of; constexpr auto to_level = key_type::meta[I].tree_level; constexpr auto to_int = utils::to_integral_type{}; utils::flip_msb_if_signed_integral(from); utils::flip_msb_if_signed_integral(to); const auto from_i = static_cast(to_int(from)); const auto to_i = static_cast(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, to_level, wraps); ml_ip_accum acc{}; 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(dpf, seg.from_node, seg.to_node, memoizer); auto * nodes = memoizer[to_level]; const std::size_t count = static_cast(seg.to_node - seg.from_node); for (std::size_t j = 0; j < count; ++j) { auto leaf = dpf.template traverse_exterior(nodes[j]); acc.mac(leaf, (start + j) * opl, opl, weights); } start += seg.count; } return acc.finish(); } template Beta eval_cmp_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to, Weights && weights) { if (!dpf.has_cmp()) throw std::invalid_argument("cmp inner product: no comparison channel"); if (!dpf.cmp_assigned()) throw std::invalid_argument( "cmp inner product: wildcard payload not assigned (call assign_cmp)"); constexpr auto to_int = utils::to_integral_type{}; utils::flip_msb_if_signed_integral(from); utils::flip_msb_if_signed_integral(to); const auto nbits = static_cast(dpf.cmp().nbits); const uint64_t mask = dpf.cmp().mask; using integral = typename KeyT::integral_type; const auto a = static_cast(to_int(from)); const auto b = static_cast(to_int(to)); const auto count = cmp_inclusive_count(a, b); constexpr std::size_t stop = KeyT::cmp_depth == 0 ? KeyT::depth : KeyT::cmp_depth; detail::incr::cmp_full_interval_memo memo{count}; const std::size_t levels = unwrap_party_key_t::cmp_block > 0 ? unwrap_party_key_t::cmp_h : nbits; detail::incr::eval_cmp_interval_impl_interior(dpf, a, cmp_exclusive_end(b), nbits, memo, levels); uint64_t dot = 0; for (std::size_t i = 0; i < count; ++i) { const auto q = static_cast(a + static_cast(i)); const uint64_t raw = [&] { if constexpr (unwrap_party_key_t::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); } }(); const uint64_t wt = static_cast(weights[i]) & mask; dot = (dot + ((raw & mask) * wt)) & mask; } return detail::dcf_impl::u64_to_beta(dot); } } // namespace incr } // namespace detail template , bool> = true> auto eval_inner_product(out_t, const KeyT & key, LaneT from, LaneT to, Weights && weights, IntervalMemoizer && memo) { constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_inner_product_impl(key, from, to, std::forward(weights), std::forward(memo)); } template , bool> = true> auto eval_inner_product(out_t, const KeyT & key, LaneT from, LaneT to, Weights && weights) { auto memo = make_basic_interval_memoizer(from, to); constexpr auto pref = detail::resolved_out_prefix(); return detail::incr::eval_out_inner_product_impl(key, from, to, std::forward(weights), memo); } template Beta eval_inner_product(cmp_t, const KeyT & key, LaneT from, LaneT to, Weights && weights) { return detail::incr::eval_cmp_inner_product_impl(key, from, to, std::forward(weights)); } // --------------------------------------------------------------------------- // eval_sequence_breadth_first(out, key, begin, end [, outbuf]) // // Breadth-first sequence eval that stops the interior walk at `meta[I] // .tree_level` (the leaf level of slot `I`) instead of the full key depth. // `begin`/`end` are a *sorted* range of lane points in `[0, 2^N)` (top-N-bit // prefixes); the result is written output-only, one value per query point in // query order (`outbuf[i]` is the output for the `i`-th query). // --------------------------------------------------------------------------- namespace detail { namespace incr { template void eval_out_sequence_breadth_first_impl(const KeyT & dpf, ForwardIterator begin, ForwardIterator end, OutputBuffer && outbuf) { using key_type = KeyT; static_assert(key_type::meta[I].prefix == N, "breadth-first out sequence: N does not match output I"); using input_type = typename key_type::input_type; using node_type = typename key_type::interior_node; using exterior_node = typename key_type::exterior_node; using output_type = typename key_type::template concrete_output_type; constexpr std::size_t stop = key_type::meta[I].tree_level; constexpr std::size_t lg_opl = key_type::template lg_outputs_per_leaf_of; constexpr std::size_t opl = std::size_t{1} << lg_opl; if (HEDLEY_UNLIKELY(!std::is_sorted(begin, end))) throw std::runtime_error("breadth-first sequence: list must be sorted"); if (begin == end) return; using allocator = aligned_allocator; allocator alloc{}; const std::size_t nseq = static_cast(std::distance(begin, end)); auto memo = alloc.allocate_unique_ptr(nseq * 2); if (HEDLEY_UNLIKELY(memo == nullptr)) throw std::bad_alloc{}; input_type mask = static_cast(input_type{1} << (N - 1)); bool curhalf = (stop ^ 1) & 1; memo[static_cast(!curhalf) * nseq + 0] = dpf.root(); std::list splits{begin, end}; std::size_t level_index = 1; auto step = [&]() { std::size_t i = 0, j = 0; const node_type cw[2] = { dpf.correction_word(level_index - 1, 0), dpf.correction_word(level_index - 1, 1)}; const std::size_t cur = static_cast(curhalf) * nseq; const std::size_t prv = static_cast(!curhalf) * nseq; for (auto upper = std::begin(splits), lower = upper++; upper != std::end(splits); lower = upper++) { auto it = std::upper_bound(*lower, *upper, mask, [](auto a, auto b) { return static_cast(a & b); }); if (it == *lower) { memo[cur + i++] = key_type::traverse_interior( memo[prv + j++], cw[1], 1); } else if (it == *upper) { memo[cur + i++] = key_type::traverse_interior( memo[prv + j++], cw[0], 0); } else { auto kids = key_type::traverse_interior01(memo[prv + j++], cw[0], cw[1]); memo[cur + i++] = kids[0]; memo[cur + i++] = kids[1]; splits.insert(upper, it); } } }; for (; level_index <= stop; ++level_index, mask >>= 1, curhalf = !curhalf) step(); auto * buf = memo.get(); // deepest built level (stop) lands in half 0 auto curr = begin, prev = begin; std::size_t j = 0; for (std::size_t i = 0; i < nseq; ++i) { if (i > 0 && (static_cast(*curr) >> lg_opl) != (static_cast(*prev) >> lg_opl)) ++j; auto leaf = dpf.template traverse_exterior(buf[j]); const std::size_t off = static_cast(static_cast(*curr) & (opl - 1)); auto v = dpf::extract_leaf(leaf, off); if constexpr (is_party_key_v) outbuf[i] = subtractive_share>::from_raw(v); else outbuf[i] = v; prev = curr++; } } } // namespace incr } // namespace detail template , bool> = true> void eval_sequence_breadth_first(out_t, const KeyT & key, ForwardIterator begin, ForwardIterator end, OutputBuffer && outbuf) { constexpr auto pref = detail::resolved_out_prefix(); detail::incr::eval_out_sequence_breadth_first_impl(key, begin, end, std::forward(outbuf)); } template , bool> = true> auto eval_sequence_breadth_first(out_t, const KeyT & key, ForwardIterator begin, ForwardIterator end) { using output_type = typename KeyT::template concrete_output_type; const std::size_t n = static_cast(std::distance(begin, end)); dpf::output_buffer> buf(n); eval_sequence_breadth_first(out_t{}, key, begin, end, buf); return buf; } /// Build a sequence recipe stopped at slot `I`'s tree level (prefix domain). template , bool> = true> auto make_sequence_recipe(out_t, const KeyT & key, ForwardIterator begin, ForwardIterator end) { constexpr auto pref = detail::resolved_out_prefix(); using input_type = typename KeyT::input_type; constexpr auto stop = KeyT::meta[I].tree_level; constexpr auto lg = KeyT::template lg_outputs_per_leaf_of; const input_type lane_msb = static_cast(input_type{1} << (pref - 1)); (void)key; return make_sequence_recipe_at(lane_msb, begin, end); } } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_EVAL_UNIFIED_HPP__