libdpf/include/dpf/leaf_later.hpp

233 lines
8.4 KiB
C++
Raw Permalink Normal View History

/// @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__