libdpf/include/dpf/leaf_later.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

232 lines
8.4 KiB
C++
Raw Permalink 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/leaf_later.hpp
/// @brief Defer the leaf correction until a public group element is known.
/// @details `eval_full` / `eval_full_add_into` with `leaf_later` skip the leaf
/// correction word, write the uncorrected group share, and fill a
/// parallel control-bit buffer the caller owns. After a rotate of
/// value and control together, `apply_leaf_correction` does
/// `buf[i] += F * control[i]` so Duoram's `v[i−s] + F·t[i−s]` never
/// reapplies a walk.
/// @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_LEAF_LATER_HPP__
#define LIBDPF_INCLUDE_DPF_LEAF_LATER_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <limits>
#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/eval_interval.hpp"
#include "dpf/interval_memoizer.hpp"
#include "dpf/leaf_arithmetic.hpp"
#include "dpf/leaf_node.hpp"
#include "dpf/twiddle.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// @brief Tag: expand without applying the leaf correction word.
struct leaf_later
{
static constexpr bool is_leaf_later_tag = true;
};
template <typename T>
struct is_leaf_later : std::false_type {};
template <>
struct is_leaf_later<leaf_later> : std::true_type {};
template <typename T>
inline constexpr bool is_leaf_later_v = is_leaf_later<std::decay_t<T>>::value;
// Forward declaration: defined in eval_walk.hpp (same offset convention).
struct rotate;
namespace detail_leaf_later
{
template <typename KeyT>
HEDLEY_NO_THROW
constexpr std::size_t domain_size() noexcept
{
constexpr std::size_t bits =
utils::bitlength_of_v<typename KeyT::input_type>;
static_assert(bits < 8 * sizeof(std::size_t),
"leaf_later: input domain does not fit a std::size_t index");
return std::size_t{1} << bits;
}
/// @brief Uncorrected exterior leaf: `0 − mask` (same group as traverse_exterior).
template <std::size_t I = 0, typename DpfKey>
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
auto traverse_exterior_uncorrected(const typename DpfKey::interior_node & node)
{
using output_type = typename DpfKey::concrete_output_type<I>;
using exterior_prg = typename DpfKey::exterior_prg;
using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
using leaf_type = dpf::leaf_node_t<typename DpfKey::exterior_node, output_type>;
leaf_type zero{};
return dpf::subtract_leaf<output_type>(
zero,
make_leaf_mask_inner<exterior_prg, I, outputs_tuple>(
unset_lo_2bits(node)));
}
template <typename DpfKey, typename = void>
struct has_opl_of : std::false_type {};
template <typename DpfKey>
struct has_opl_of<DpfKey,
std::void_t<decltype(DpfKey::template outputs_per_leaf_of<0>)>>
: std::true_type {};
template <typename DpfKey>
HEDLEY_NO_THROW
constexpr std::size_t opl_of() noexcept
{
if constexpr (has_opl_of<DpfKey>::value)
return DpfKey::template outputs_per_leaf_of<0>;
else
return DpfKey::outputs_per_leaf;
}
} // namespace detail_leaf_later
/// @brief Full-domain expansion without the leaf correction; also fills `control`.
/// @details `buf[i]` receives the uncorrected share at domain point `i`, and
/// `control[i]` is the leaf node's control bit. Both must be sized to
/// `2^n`.
/// \complexity One full-domain expansion, `Θ(2^n)`.
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
&& !is_multilevel_key_v<DpfKey>, bool> = true>
void eval_full(Buffer & buf, Control & control, const DpfKey & dpf,
leaf_later) // NOLINT(runtime/references)
{
using input_type = typename DpfKey::input_type;
using output_type = typename DpfKey::concrete_output_type<I>;
using exterior_node = typename DpfKey::exterior_node;
constexpr std::size_t opl = DpfKey::outputs_per_leaf;
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
if (buf.size() < n || control.size() < n)
throw std::invalid_argument("eval_full(leaf_later): buffer too small");
auto memo = make_basic_full_memoizer(dpf);
const input_type from = std::numeric_limits<input_type>::min();
const input_type to = std::numeric_limits<input_type>::max();
const auto from_node = utils::get_from_node<DpfKey>(from);
const auto to_node = utils::get_to_node<DpfKey>(to);
internal::eval_interval_interior(dpf, from_node, to_node, memo);
const std::size_t nodes_in_interval =
static_cast<std::size_t>(to_node - from_node);
auto * nodes = memo[DpfKey::depth];
for (std::size_t j = 0; j < nodes_in_interval; ++j)
{
const auto & node = nodes[j];
const bool tbit = static_cast<bool>(get_lo_bit(node));
auto leaf = detail_leaf_later::traverse_exterior_uncorrected<I, DpfKey>(
node);
for (std::size_t o = 0; o < opl; ++o)
{
const std::size_t idx = j * opl + o;
if (idx >= n)
break;
output_type y = extract_leaf<exterior_node, output_type>(leaf, o);
buf[idx] = y;
control[idx] = tbit;
}
}
}
/// @brief Add an uncorrected full-domain expansion into `buf` and fill `control`.
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
&& !is_multilevel_key_v<DpfKey>, bool> = true>
void eval_full_add_into(Buffer & buf, Control & control, const DpfKey & dpf,
leaf_later) // NOLINT(runtime/references)
{
using output_type = typename DpfKey::concrete_output_type<I>;
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
std::vector<output_type> tmp(n);
std::vector<std::uint8_t> ctl(n);
eval_full<I>(tmp, ctl, dpf, leaf_later{});
for (std::size_t i = 0; i < n; ++i)
{
buf[i] = buf[i] + tmp[i];
control[i] = static_cast<bool>(ctl[i]);
}
}
/// @brief Uncorrected expansion written at `(i + rot.shift) mod 2^n`, with
/// control bits rotated the same way.
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
&& !is_multilevel_key_v<DpfKey>, bool> = true>
void eval_full_add_into(Buffer & buf, Control & control, const DpfKey & dpf,
leaf_later, rotate rot) // NOLINT(runtime/references)
{
using output_type = typename DpfKey::concrete_output_type<I>;
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
const std::size_t s = rot.shift % n;
std::vector<output_type> tmp(n);
std::vector<std::uint8_t> ctl(n);
eval_full<I>(tmp, ctl, dpf, leaf_later{});
for (std::size_t i = 0; i < n; ++i)
{
const std::size_t j = (i + s) % n;
buf[j] = buf[j] + tmp[i];
control[j] = static_cast<bool>(ctl[i]);
}
}
/// @brief Rotate value and control by the same offset (`new[i] = old[(i-s) mod n]`).
/// @details Domain point `k` lands at `(k + s) mod n`, matching `dpf::rotate{s}`.
template <typename Buffer, typename Control>
void cyclic_shift_pair(Buffer & buf, Control & control, std::size_t s) // NOLINT(runtime/references)
{
const std::size_t n = buf.size();
if (n == 0 || control.size() != n)
throw std::invalid_argument("cyclic_shift_pair: size mismatch");
const std::size_t sh = s % n;
if (sh == 0)
return;
Buffer out_buf(buf);
Control out_ctl(control);
for (std::size_t i = 0; i < n; ++i)
{
out_buf[i] = buf[(i + n - sh) % n];
out_ctl[i] = control[(i + n - sh) % n];
}
buf = std::move(out_buf);
control = std::move(out_ctl);
}
/// @brief `buf[i] += F * control[i]` in the leaf's group.
/// @details XOR leaves treat `+` as XOR. `F` is a public group element.
template <typename Buffer, typename Control, typename F>
void apply_leaf_correction(Buffer & buf, const Control & control, const F & f) // NOLINT(runtime/references)
{
const std::size_t n = buf.size();
if (control.size() < n)
throw std::invalid_argument("apply_leaf_correction: control too small");
for (std::size_t i = 0; i < n; ++i)
{
if (static_cast<bool>(control[i]))
buf[i] = buf[i] + f;
}
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_LEAF_LATER_HPP__