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

449 lines
19 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_walk.hpp
/// @brief Walk helpers that fold a small operation into an existing DPF walk.
/// @details These are the calls the protocol mockups in
/// `examples/applications/` had to build by hand around a walk:
/// - `dpf::rotate{s}` weights the walk by `w[(i + s) mod 2^n]`
/// (Duoram's read, Pika's lookup) with no second vector.
/// - `eval_full_add_into(buf, key)` adds a full-domain expansion
/// into a buffer the caller already holds (Prio's histogram,
/// Express's mailbox, Duoram's update). A `rotate` overload shifts
/// the write; a `sketch` overload folds the audit in the same pass.
/// - `dpf::cyclic_shift(buf, s)` rotates a materialized share buffer,
/// `new[i] = old[(i - s) mod n]`.
/// - `dpf::pack_bit_columns(keys...)` runs the full-domain bit walk
/// once per key and packs one integer per row (lane `e` is key `e`),
/// the digit BitMore reads when the server count is a power of two.
/// - `dpf::mod_bit_columns<ℓ>(keys...)` is that digit modulo `ℓ` when
/// the server count is not a power of two. The first key is still
/// the low bit. The running residue stays in a byte through `ℓ = 128`
/// and in a 16-bit lane through `ℓ = 32768`.
/// - `eval_prefixes(out<I,N>, key)` returns the `2^N` prefix shares,
/// and `eval_prefix_inner_product(out<I,N>, key, values)` dots them
/// with `values` in one walk to depth `N` (Poplar, PRAC).
/// - `idpf_eval_ctx` / `eval_until` (see `dpf/eval_until.hpp`) resume
/// under a live prefix list instead of materializing `2^N` nodes.
/// @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_WALK_HPP__
#define LIBDPF_INCLUDE_DPF_EVAL_WALK_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <iterator>
#include <limits>
#include <tuple>
#include <type_traits>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/bitmore_mod.hpp"
#include "dpf/eval_target.hpp"
#include "dpf/eval_full.hpp"
#include "dpf/eval_inner_product.hpp"
#include "dpf/eval_unified.hpp"
#include "dpf/output_buffer.hpp"
#include "dpf/utils.hpp"
#include "dpf/verifiable.hpp"
namespace dpf
{
/// @brief Rotation offset applied inside a walk: domain point `i` uses lane
/// `(i + shift) mod 2^n`. Pass to `eval_full_inner_product` /
/// `eval_full_add_into` so the caller does not build a second vector.
struct rotate
{
std::size_t shift;
};
namespace detail_walk
{
/// @brief Number of input points `2^n` of `KeyT`, as a `std::size_t`.
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),
"walk helper: input domain does not fit a std::size_t index");
return std::size_t{1} << bits;
}
/// @brief `w[(i + shift) mod n]`, a read-only rotated view of `w`.
template <typename Weights>
struct rotated_weights
{
const Weights & w;
std::size_t shift;
std::size_t n;
HEDLEY_ALWAYS_INLINE
decltype(auto) operator[](std::size_t i) const
{
return w[(i + shift) % n];
}
};
template <typename T, typename = void>
struct has_raw : std::false_type {};
template <typename T>
struct has_raw<T, std::void_t<decltype(std::declval<const T &>().raw())>>
: std::true_type {};
/// @brief The group element carried by an eval buffer slot. A one-word share
/// (`additive`, `subtractive`, or `additive3`) exposes it through
/// `.raw()`. A replicated share has two components and is returned
/// as itself. A raw output type is returned unchanged.
template <typename T>
HEDLEY_ALWAYS_INLINE
auto group_value(const T & v)
{
if constexpr (has_raw<T>::value)
return v.raw();
else
return v;
}
/// @brief Call `fn` on `keys` from the last key down to the first.
/// @details `mod_bit_columns` inserts the high bit first. Key 0 stays the low
/// bit, matching `pack_bit_columns`.
template <typename Fn, typename Tuple, std::size_t... I>
void insert_keys_msb_first(Fn && fn, Tuple && keys, std::index_sequence<I...>)
{
constexpr std::size_t n = sizeof...(I);
(fn(std::get<n - 1 - I>(keys)), ...);
}
} // namespace detail_walk
// ---------------------------------------------------------------------------
// dpf::rotate on the paired full-domain inner product
// ---------------------------------------------------------------------------
/// @brief `sum_i DPF_I(i) * rows[(i + rot.shift) mod 2^n]` over the whole domain.
/// @details The same paired walk as `eval_full_inner_product(paired, key, rows)`,
/// but the weight at domain point `i` is read from `rows` rotated by
/// `rot.shift`. Duoram's read and Pika's lookup pass the unrotated
/// table and this offset instead of materializing a rotated copy.
template <std::size_t I = 0, std::size_t... Is, typename DpfKey, typename Rows>
HEDLEY_WARN_UNUSED_RESULT
auto eval_full_inner_product(paired_t, const DpfKey & dpf, Rows && rows,
rotate rot)
{
constexpr std::size_t n = detail_walk::domain_size<DpfKey>();
detail_walk::rotated_weights<std::remove_reference_t<Rows>> view{
rows, rot.shift % n, n};
return eval_full_inner_product<I, Is...>(paired, dpf, view);
}
// ---------------------------------------------------------------------------
// eval_full_add_into(buf, key [, rotate | sketch])
// ---------------------------------------------------------------------------
/// @brief Add a full-domain expansion of output `I` into `buf` in place.
/// @details `buf[i] += DPF_I(i)` (the leaf share's group element) for every
/// domain point `i`. `buf` already holds the caller's running shares
/// (Prio's histogram, Express's mailbox, Duoram's update); its element
/// type must support `+` with the leaf share's group element.
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one `+` per point.
template <std::size_t I = 0, typename Buffer, 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, const DpfKey & dpf) // NOLINT(runtime/references)
{
auto result = eval_full<I>(dpf);
auto & iter = result.second;
std::size_t i = 0;
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
buf[i] = buf[i] + detail_walk::group_value(*it);
}
/// @brief Add a full-domain expansion into `buf`, shifted by `rot`.
/// @details `buf[(i + rot.shift) mod 2^n] += DPF_I(i)`. Duoram's update writes
/// the payload placed at `r` into the memory slot `r + shift`.
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one `+` per point.
template <std::size_t I = 0, typename Buffer, 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, const DpfKey & dpf, rotate rot) // NOLINT(runtime/references)
{
constexpr std::size_t n = detail_walk::domain_size<DpfKey>();
const std::size_t s = rot.shift % n;
auto result = eval_full<I>(dpf);
auto & iter = result.second;
std::size_t i = 0;
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
{
const std::size_t j = (i + s) % n;
buf[j] = buf[j] + detail_walk::group_value(*it);
}
}
/// @brief Add a full-domain expansion into `buf` and fold the audit sketch.
/// @details `buf[i] += DPF_I(i)` for every point, and each written share is
/// absorbed into `sk` — the same one-hot audit as
/// `eval_full(key, sketch(σ))`, in the same pass. Express's mailbox
/// write becomes one call instead of a point loop and a second fold.
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one `+` and one
/// absorb per point.
template <std::size_t I = 0, typename Buffer, 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, const DpfKey & dpf, sketch_ref sk) // NOLINT(runtime/references)
{
static_assert(DpfKey::is_extractable,
"eval_full_add_into(..., sketch(σ)): key must carry dpf::extractable");
auto result = eval_full<I>(dpf);
auto & iter = result.second;
std::size_t i = 0;
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
{
auto g = detail_walk::group_value(*it);
buf[i] = buf[i] + g;
sk.absorb(g); // fold the raw share, as eval_point(..., sketch) does
}
}
// ---------------------------------------------------------------------------
// dpf::cyclic_shift(buffer, s)
// ---------------------------------------------------------------------------
/// @brief A copy of `buf` rotated so `result[i] == buf[(i - s) mod n]`.
/// @details The buffer analogue of `dpf::rotate{s}`: the value at index `k`
/// moves to `k + s`. Duoram's update shifts an expanded share buffer
/// this way before adding it into memory.
template <typename Container>
HEDLEY_WARN_UNUSED_RESULT
Container cyclic_shift(const Container & buf, std::size_t s)
{
Container out(buf);
const std::size_t n = buf.size();
if (n == 0)
return out;
const std::size_t sh = s % n;
for (std::size_t i = 0; i < n; ++i)
out[i] = buf[(i + n - sh) % n];
return out;
}
// ---------------------------------------------------------------------------
// dpf::pack_bit_columns(keys...)
// ---------------------------------------------------------------------------
/// @brief Pack the full-domain bit expansions of `keys` into one integer per row.
/// @details Runs `eval_full` once per key over a shared bit domain and sets bit
/// `e` of `result[row]` when key `e` opens to 1 at that row. `L` keys
/// need `L <= 8*sizeof(Int)`. BitMore reads `result[row]` as the
/// server's `L`-bit digit instead of unpacking one `int` per bit.
/// @tparam Int packed digit type (defaults to `std::uint64_t`)
/// \complexity One full-domain bit expansion per key, `Θ(L · 2^n)`.
template <typename Int = std::uint64_t, typename First, typename... Rest>
HEDLEY_WARN_UNUSED_RESULT
std::vector<Int> pack_bit_columns(const First & first, const Rest &... rest)
{
static_assert(1 + sizeof...(Rest) <= 8 * sizeof(Int),
"pack_bit_columns: more keys than bits in the packed digit type");
constexpr std::size_t n = detail_walk::domain_size<First>();
std::vector<Int> out(n, Int{0});
std::size_t e = 0;
auto do_one = [&](const auto & key)
{
auto result = eval_full(key);
auto & iter = result.second;
std::size_t row = 0;
for (auto it = std::begin(iter); it != std::end(iter) && row < n;
++it, ++row)
{
if (static_cast<bool>(*it))
out[row] |= static_cast<Int>(Int{1} << e);
}
++e;
};
do_one(first);
(do_one(rest), ...);
return out;
}
// ---------------------------------------------------------------------------
// dpf::mod_bit_columns<ℓ>(keys...)
// ---------------------------------------------------------------------------
/// @brief One residue per row: the packed bit columns, modulo `Modulus`.
/// @details The same full-domain bit walk as `pack_bit_columns`. Key 0 is the
/// low bit, so row `r` opens to `(sum_e bit_e(r) · 2^e) mod Modulus`.
/// Bits are folded high-bit first. Through modulus 128 each running
/// slot is a byte; through 32768 it is a 16-bit lane. Both use the
/// high-nibble partial reduction, and only the finished row is fully
/// reduced. Hafiz and Henry §5.3 read `result[row]` as the server's
/// digit when the server count is not a power of two.
/// @tparam Modulus server count, `2` through `32768`
/// @tparam Int residue type (defaults to `std::uint16_t`)
/// \complexity One full-domain bit expansion per key, `Θ(L · 2^n)`, and one
/// lane insertion per row per key.
template <unsigned Modulus, typename Int = std::uint16_t, typename First, typename... Rest>
HEDLEY_WARN_UNUSED_RESULT
std::vector<Int> mod_bit_columns(const First & first, const Rest &... rest)
{
static_assert(Modulus >= 2u && Modulus <= 32768u,
"mod_bit_columns: modulus must be in 2..32768");
static_assert(std::is_unsigned_v<Int>,
"mod_bit_columns: residue type must be unsigned");
static_assert(Modulus - 1u <= static_cast<std::uintmax_t>(std::numeric_limits<Int>::max()),
"mod_bit_columns: residue type cannot hold a value modulo Modulus");
constexpr unsigned lane_bits = Modulus <= 128u ? 8u : 16u;
using reg = simde__m256i;
using acc = bitmore_accumulator<Modulus, reg, lane_bits>;
constexpr std::size_t lanes = lane_bits == 8u ? 32u : 16u;
constexpr std::size_t n = detail_walk::domain_size<First>();
const std::size_t chunks = (n + lanes - 1u) / lanes;
std::vector<acc> running(chunks);
const auto keys = std::forward_as_tuple(first, rest...);
detail_walk::insert_keys_msb_first([&](const auto & key)
{
using key_t = std::decay_t<decltype(key)>;
static_assert(detail_walk::domain_size<key_t>() == n,
"mod_bit_columns: every key must share the first key's domain");
auto result = eval_full(key);
auto & iter = result.second;
auto it = std::begin(iter);
const auto end = std::end(iter);
for (std::size_t c = 0; c < chunks; ++c)
{
reg bits{};
if constexpr (lane_bits == 8u)
{
alignas(32) unsigned char raw[32]{};
for (std::size_t i = 0; i < lanes && it != end; ++i, ++it)
raw[i] = static_cast<bool>(*it) ? 1u : 0u;
std::memcpy(&bits, raw, sizeof(bits));
}
else
{
alignas(32) std::uint16_t raw[16]{};
for (std::size_t i = 0; i < lanes && it != end; ++i, ++it)
raw[i] = static_cast<bool>(*it) ? 1u : 0u;
std::memcpy(&bits, raw, sizeof(bits));
}
running[c].insert_bit(bits);
}
}, keys, std::make_index_sequence<1u + sizeof...(Rest)>{});
std::vector<Int> out(n);
std::size_t row = 0;
for (std::size_t c = 0; c < chunks; ++c)
{
const reg reduced = running[c].reduced();
if constexpr (lane_bits == 8u)
{
alignas(32) unsigned char raw[32]{};
std::memcpy(raw, &reduced, sizeof(raw));
for (std::size_t i = 0; i < lanes && row < n; ++i, ++row)
out[row] = static_cast<Int>(raw[i]);
}
else
{
alignas(32) std::uint16_t raw[16]{};
std::memcpy(raw, &reduced, sizeof(raw));
for (std::size_t i = 0; i < lanes && row < n; ++i, ++row)
out[row] = static_cast<Int>(raw[i]);
}
}
return out;
}
// ---------------------------------------------------------------------------
// eval_prefixes / eval_prefix_inner_product (idpf prefix walk)
// ---------------------------------------------------------------------------
/// @brief The `2^N` prefix shares of output `I` (prefix length `N`).
/// @details One walk to depth `N`; slot `p` is the share on prefix `p`. Poplar
/// reads these to score every node at a depth in one pass instead of
/// one `eval_point` per node. Returns the `(buffer, iterable)` pair of
/// `eval_full(out<I,N>, key)`.
/// \complexity One walk to depth `N`: `Θ(N + 2^N)` interior traversals.
template <std::size_t I, std::size_t N, typename KeyT,
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
HEDLEY_WARN_UNUSED_RESULT
auto eval_prefixes(out_t<I, N>, const KeyT & key)
{
return eval_full(out_t<I, N>{}, key);
}
/// @brief `sum_p DPF_I(p) * values[p]` over the `2^N` prefixes of output `I`.
/// @details Walks to depth `N` once and dots the prefix shares with `values`
/// (indexed by prefix `0 .. 2^N - 1`). PRAC's strides and Poplar's
/// "is this prefix heavy?" are this one call, replacing an
/// `eval_point` per node.
/// \complexity One walk to depth `N`, `Θ(N + 2^N)`, plus one multiply-add per
/// prefix into an `O(1)` accumulator.
template <std::size_t I, std::size_t N, typename KeyT, typename Values,
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
HEDLEY_WARN_UNUSED_RESULT
auto eval_prefix_inner_product(out_t<I, N>, const KeyT & key, Values && values)
{
using lane_t = typename KeyT::input_type;
constexpr lane_t lo = lane_t{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_inner_product(out_t<I, N>{}, key, lo, hi,
std::forward<Values>(values));
}
/// @brief XOR of `records[j]` for each listed point where the bit share is set.
/// @details Same walk as `eval_sequence` on a bit key, but folds the selected
/// records into one accumulator instead of materializing a bit vector.
/// Keyword PIR's server response is this one call.
template <typename DpfKey, typename ForwardIterator, typename Records,
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<DpfKey>>, int> = 0>
HEDLEY_WARN_UNUSED_RESULT
auto eval_sequence_xor(const DpfKey & key, ForwardIterator begin,
ForwardIterator end, Records && records)
{
using record_t = std::decay_t<decltype(records[0])>;
record_t acc{};
auto [buf, iter] = eval_sequence(key, begin, end);
std::size_t i = 0;
for (auto bit : iter)
{
if (static_cast<bool>(bit))
acc = static_cast<record_t>(acc ^ records[i]);
++i;
}
return acc;
}
/// @brief Evaluate a walk while invoking `fold(index, share)` once per written output.
/// @details The fold is a template — inlined the way `sketch_ref::absorb` is.
/// Express's audit, Pika's SZ check, and a SNIP input vector each
/// supply their own fold over the group they write.
template <typename Fold, typename Buffer, typename DpfKey,
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
&& !is_multilevel_key_v<DpfKey>
&& !std::is_same_v<std::decay_t<Fold>, sketch_ref>
&& !std::is_same_v<std::decay_t<Fold>, rotate>, bool> = true>
void eval_full_add_into(Buffer & buf, const DpfKey & dpf, Fold && fold) // NOLINT(runtime/references)
{
auto result = eval_full(dpf);
auto & iter = result.second;
std::size_t i = 0;
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
{
auto g = detail_walk::group_value(*it);
buf[i] = buf[i] + g;
fold(i, g);
}
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_EVAL_WALK_HPP__