libdpf/include/dpf/eval_walk.hpp

450 lines
19 KiB
C++
Raw Permalink Normal View History

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