libdpf/include/dpf/eval_until.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

323 lines
11 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

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

/// @file dpf/eval_until.hpp
/// @brief Prefix-resuming evaluation for multilevel / incremental DPF keys.
/// @details Google's incremental-DPF `EvaluateUntil(level, prefixes, ctx)`
/// (Poplar / private heavy hitters, ePrint 2021/017) walks only the
/// live prefixes. `eval_prefixes` always restarts at the root and
/// materializes all `2^N` nodes; a path memoizer resumes one path.
/// `idpf_eval_ctx` keeps the interior node under each live prefix.
/// `eval_until(ctx, level, prefixes)` returns the output shares at
/// that hierarchy level for those prefixes only, then updates the
/// context. Empty `prefixes` on a fresh context (level 0) leaves the
/// root seed in place. Each later prefix must extend a prefix from
/// the previous call.
/// @see dpf/eval_walk.hpp (`eval_prefixes`), dpf/idpf_agg.hpp
/// @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_UNTIL_HPP__
#define LIBDPF_INCLUDE_DPF_EVAL_UNTIL_HPP__
#include <cstddef>
#include <cstdint>
#include <stdexcept>
#include <type_traits>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/dpf_key.hpp"
#include "dpf/eval_common.hpp"
#include "dpf/incremental.hpp"
#include "dpf/placement.hpp"
#include "dpf/secret_share.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// @brief Live-prefix evaluation context for a multilevel / idpf key.
/// @details Hierarchy level 0 holds only the root seed. After
/// `eval_until(..., L, prefixes)`, `level()` is `L` and
/// `node_count()` equals `prefixes.size()`.
/// \complexity O(|prefixes|) stored nodes (not O(2^level)).
template <typename KeyT>
class idpf_eval_ctx
{
public:
using key_type = unwrap_party_key_t<KeyT>;
using input_type = typename key_type::input_type;
using node_type = typename key_type::interior_node;
explicit idpf_eval_ctx(const KeyT & key)
: key_{&key}, level_{0}, prefixes_{}, nodes_{}
{
nodes_.push_back(static_cast<const key_type &>(key).root());
// Compact parent of every length-1 prefix is the empty prefix `0`.
prefixes_.push_back(input_type{0});
}
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
const KeyT & key() const noexcept { return *key_; }
/// @brief Last hierarchy level that `eval_until` wrote (0 = root only).
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
std::size_t level() const noexcept { return level_; }
/// @brief Number of saved interior nodes (one per live prefix).
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
std::size_t node_count() const noexcept { return nodes_.size(); }
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
const std::vector<input_type> & prefixes() const noexcept
{
return prefixes_;
}
/// @brief Drop every saved prefix except `prefix` (must be live).
void retain(input_type prefix)
{
for (std::size_t i = 0; i < prefixes_.size(); ++i)
{
if (prefixes_[i] == prefix)
{
prefixes_ = {prefix};
nodes_ = {nodes_[i]};
return;
}
}
throw std::invalid_argument(
"idpf_eval_ctx::retain: prefix is not live in this context");
}
private:
template <typename K, typename PrefRange>
friend auto eval_until(idpf_eval_ctx<K> & ctx, std::size_t level,
PrefRange && prefixes);
const KeyT * key_;
std::size_t level_;
std::vector<input_type> prefixes_;
std::vector<node_type> nodes_;
};
namespace detail
{
namespace eval_until_detail
{
template <typename KeyT>
HEDLEY_NO_THROW
constexpr std::size_t slot_for_prefix(std::size_t prefix_len) noexcept
{
for (std::size_t i = 0; i < KeyT::num_outputs; ++i)
{
if (KeyT::meta[i].prefix == prefix_len)
return i;
}
return static_cast<std::size_t>(-1);
}
template <typename KeyT>
HEDLEY_NO_THROW
constexpr std::size_t tree_level_for_prefix(std::size_t prefix_len) noexcept
{
const auto i = slot_for_prefix<KeyT>(prefix_len);
return (i == static_cast<std::size_t>(-1)) ? static_cast<std::size_t>(-1)
: KeyT::meta[i].tree_level;
}
/// @brief Walk `from_tree_level` → `to_tree_level` along the high bits of `x`.
/// \complexity O(to − from) interior traversals.
template <typename KeyT, typename Node>
HEDLEY_ALWAYS_INLINE
Node walk_interior(const KeyT & key, Node node, typename KeyT::input_type x,
std::size_t from_tree_level, std::size_t to_tree_level)
{
using key_type = KeyT;
if (to_tree_level <= from_tree_level)
return node;
auto level_index = from_tree_level + 1;
auto mask = key.msb_mask >> (level_index - 1);
for (; level_index <= to_tree_level; ++level_index, mask >>= 1)
{
const bool bit = !!(mask & x);
auto cw = key.correction_word(level_index - 1, bit);
const bool is_last = key_type::tree::is_last_level(level_index - 1,
key.depth);
node = key_type::traverse_interior(node, cw, bit, is_last);
}
return node;
}
template <std::size_t I, typename KeyT, typename Node, typename InputT>
HEDLEY_ALWAYS_INLINE
auto exterior_at(const KeyT & key, const Node & node, InputT domain_x)
{
using key_type = unwrap_party_key_t<KeyT>;
using output_type = typename key_type::template concrete_output_type<I>;
constexpr auto N = key_type::meta[I].prefix;
// `traverse_exterior` lives on the underlying key; party_key inherits it.
auto leaf = key.template traverse_exterior<I>(node);
auto lane_x = detail::incr::lane_input(domain_x, N, key_type::input_bits);
detail::incr::absorb_public_addend_lane<I>(
static_cast<const key_type &>(key), leaf, lane_x);
// Pass the party-tagged key type so the share carries the party id.
return *make_eval_dpf_output<KeyT, output_type>(leaf, lane_x);
}
template <typename KeyT, typename Node, typename InputT, typename Out,
std::size_t... Is>
HEDLEY_ALWAYS_INLINE
bool exterior_dispatch(std::size_t slot, const KeyT & key, const Node & node,
InputT domain_x, Out & out, std::index_sequence<Is...>)
{
return ((Is == slot
? (out = exterior_at<Is>(key, node, domain_x), true)
: false)
|| ...);
}
template <typename KeyT>
HEDLEY_ALWAYS_INLINE
auto share_type_tag()
{
using key_type = unwrap_party_key_t<KeyT>;
using output_type = typename key_type::template concrete_output_type<0>;
return eval_leaf_result_t<KeyT, output_type>{};
}
} // namespace eval_until_detail
} // namespace detail
/// @brief Evaluate output shares at hierarchy `level` for `prefixes` only.
/// @details Each prefix is a compact integer in `[0, 2^level)`. The first call
/// may pass an empty range at level 0 to keep the root seed. Later
/// calls require `level > ctx.level()` and every prefix to extend one
/// saved parent. Returns one share per prefix, in the same order.
/// Opened values match `eval_point(out<I,N>, key, domain)` for the
/// same prefix length (see `eval_until_test`).
/// \complexity O(|prefixes| · (level − ctx.level())) interior traversals.
/// Saved nodes stay O(|prefixes|), not O(2^level).
/// @see ePrint 2021/017 (Poplar EvaluateUntil), ePrint 2024/1190 (I-DPF agg)
template <typename KeyT, typename PrefRange>
HEDLEY_WARN_UNUSED_RESULT
auto eval_until(idpf_eval_ctx<KeyT> & ctx, std::size_t level,
PrefRange && prefixes)
{
using key_type = typename idpf_eval_ctx<KeyT>::key_type;
using input_type = typename idpf_eval_ctx<KeyT>::input_type;
using node_type = typename idpf_eval_ctx<KeyT>::node_type;
using share_type = decltype(detail::eval_until_detail::share_type_tag<KeyT>());
static_assert(is_multilevel_key_v<key_type>,
"eval_until requires a multilevel / incremental DPF key");
const KeyT & key_ref = ctx.key();
const key_type & key = static_cast<const key_type &>(key_ref);
std::vector<input_type> pref_list;
for (auto && p : prefixes)
pref_list.push_back(static_cast<input_type>(p));
if (pref_list.empty())
{
if (ctx.level_ != 0 || level != 0)
{
throw std::invalid_argument(
"eval_until: empty prefixes only valid on a fresh context at level 0");
}
return std::vector<share_type>{};
}
if (level == 0)
{
throw std::invalid_argument(
"eval_until: level 0 has only the root; pass a positive hierarchy level");
}
if (level <= ctx.level_)
{
throw std::invalid_argument(
"eval_until: level must strictly advance past the context level");
}
if (level > key_type::input_bits)
{
throw std::invalid_argument(
"eval_until: level exceeds the input bit length");
}
const auto slot = detail::eval_until_detail::slot_for_prefix<key_type>(level);
if (slot == static_cast<std::size_t>(-1))
{
throw std::invalid_argument(
"eval_until: key has no output planted at that prefix length");
}
const auto to_tree =
detail::eval_until_detail::tree_level_for_prefix<key_type>(level);
const std::size_t from_tree = (ctx.level_ == 0)
? 0
: detail::eval_until_detail::tree_level_for_prefix<key_type>(ctx.level_);
constexpr auto bits = key_type::input_bits;
const auto parent_shift = level - ctx.level_;
std::vector<node_type> new_nodes;
std::vector<share_type> shares;
new_nodes.reserve(pref_list.size());
shares.reserve(pref_list.size());
for (auto p : pref_list)
{
if (level < bits && (static_cast<std::uint64_t>(p) >> level) != 0)
{
throw std::invalid_argument(
"eval_until: prefix does not fit in the requested bit length");
}
const input_type parent = static_cast<input_type>(p >> parent_shift);
std::size_t parent_idx = static_cast<std::size_t>(-1);
for (std::size_t i = 0; i < ctx.prefixes_.size(); ++i)
{
if (ctx.prefixes_[i] == parent)
{
parent_idx = i;
break;
}
}
if (parent_idx == static_cast<std::size_t>(-1))
{
throw std::invalid_argument(
"eval_until: prefix does not extend a live parent");
}
auto domain_x = static_cast<input_type>(
static_cast<std::uint64_t>(p) << (bits - level));
domain_x = key.offset_x(domain_x);
utils::flip_msb_if_signed_integral(domain_x);
node_type node = detail::eval_until_detail::walk_interior(key,
ctx.nodes_[parent_idx], domain_x, from_tree, to_tree);
share_type share{};
const bool ok = detail::eval_until_detail::exterior_dispatch(slot,
key_ref, node, domain_x, share,
std::make_index_sequence<key_type::num_outputs>{});
if (!ok)
{
throw std::logic_error("eval_until: exterior dispatch missed slot");
}
new_nodes.push_back(node);
shares.push_back(share);
}
ctx.level_ = level;
ctx.prefixes_ = std::move(pref_list);
ctx.nodes_ = std::move(new_nodes);
return shares;
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_EVAL_UNTIL_HPP__