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

911 lines
34 KiB
C++

/// @file dpf/cohort.hpp
/// @brief Many classic DPFs that share one public schedule.
/// @details Generation is the ordinary special-path walk, run level by level
/// across keys so the path bit is read once and `expand_x4` sees
/// contiguous seeds. Evaluation uses that same idea on a point, a
/// closed interval, or one `sequence_recipe`: correction words are
/// hoisted once per level, and interior nodes stay key-interleaved
/// (`position * stride + key`) down to the leaves.
///
/// Leaf layout, with `n` keys:
/// - point: `out[key]`
/// - interval: leaf node `j`, lane `p` of key `k` is
/// `out[(j * n + k) * outputs_per_leaf + p]`
/// (same lane order as one-key `eval_interval`)
/// - sequence: listed point `q` of key `k` is `out[q * n + k]`
///
/// `cohort_index(j, k, n)` is `j * n + k`.
/// Inner products use the same order and the same weights for every
/// key. Generation is the classic (not incremental) dealer key: one
/// plaintext point, one payload or `std::tuple` of payloads per key.
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details.
#ifndef LIBDPF_INCLUDE_DPF_COHORT_HPP__
#define LIBDPF_INCLUDE_DPF_COHORT_HPP__
#include "hedley/hedley.h"
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <limits>
#include <stdexcept>
#include <tuple>
#include <type_traits>
#include <utility>
#include <vector>
#include "dpf/aligned_allocator.hpp"
#include "dpf/dpf_key.hpp"
#include "dpf/eval_common.hpp"
#include "dpf/leaf_node.hpp"
#include "dpf/random.hpp"
#include "dpf/sequence_recipe.hpp"
#include "dpf/utils.hpp"
namespace dpf
{
/// @brief `j * nkeys + key`. Point, interval leaf, and sequence point all use this.
HEDLEY_CONST
HEDLEY_NO_THROW
inline constexpr std::size_t cohort_index(std::size_t j, std::size_t key,
std::size_t nkeys) noexcept
{
return j * nkeys + key;
}
/// @brief Two planes of key-interleaved interior nodes, reused across calls.
/// @tparam Key a DPF key or a `party_key` of one. Every key in a walk has this type.
template <typename Key>
class cohort_scratch
{
public:
using key_type = unwrap_party_key_t<std::decay_t<Key>>;
using node_type = typename key_type::interior_node;
using allocator = aligned_allocator<node_type>;
HEDLEY_NO_THROW
std::size_t count() const noexcept { return count_; }
HEDLEY_NO_THROW
std::size_t stride() const noexcept { return stride_; }
/// @brief Grow so each of `positions` tree slots can hold `nkeys` nodes,
/// padded to a multiple of 4 for `expand_x4`.
void fit(std::size_t nkeys, std::size_t positions)
{
count_ = nkeys;
stride_ = nkeys == 0 ? 0 : ((nkeys + 3u) & ~std::size_t{3});
const std::size_t pos = positions == 0 ? 1 : positions;
const std::size_t need = 2 * pos * stride_;
if (buf_.size() < need)
buf_.resize(need);
positions_ = pos;
if (cw0_.size() < nkeys)
{
cw0_.resize(nkeys);
cw1_.resize(nkeys);
}
}
node_type * plane(int which) noexcept
{
return buf_.data() + static_cast<std::size_t>(which) * positions_ * stride_;
}
node_type * cw0() noexcept { return cw0_.data(); }
node_type * cw1() noexcept { return cw1_.data(); }
private:
std::size_t count_ = 0;
std::size_t stride_ = 0;
std::size_t positions_ = 0;
std::vector<node_type, allocator> buf_;
std::vector<node_type, allocator> cw0_;
std::vector<node_type, allocator> cw1_;
};
namespace cohort_detail
{
template <typename T>
struct is_std_tuple : std::false_type {};
template <typename... Ts>
struct is_std_tuple<std::tuple<Ts...>> : std::true_type {};
template <typename Tree, typename Node>
HEDLEY_ALWAYS_INLINE
void expand_keys(const Node * HEDLEY_RESTRICT parents, std::size_t nkeys,
bool is_last, const Node * HEDLEY_RESTRICT cw_left,
const Node * HEDLEY_RESTRICT cw_right, int which,
Node * HEDLEY_RESTRICT dest_left, Node * HEDLEY_RESTRICT dest_right)
{
// which: 0 left, 1 right, 2 both. Both destinations are real pointers.
std::size_t k = 0;
for (; k + 4 <= nkeys; k += 4)
{
alignas(Node) Node left[4];
alignas(Node) Node right[4];
Tree::expand_x4(parents + k, left, right, is_last);
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
{
if (which != 1)
{
dest_left[k + t] = dpf::xor_if_lo_bit(left[t], cw_left[k + t],
parents[k + t]);
}
if (which != 0)
{
dest_right[k + t] = dpf::xor_if_lo_bit(right[t], cw_right[k + t],
parents[k + t]);
}
}
}
for (; k < nkeys; ++k)
{
const auto kids = Tree::expand(parents[k], is_last);
if (which != 1)
dest_left[k] = dpf::xor_if_lo_bit(kids[0], cw_left[k], parents[k]);
if (which != 0)
dest_right[k] = dpf::xor_if_lo_bit(kids[1], cw_right[k], parents[k]);
}
}
template <typename Tree, typename Node>
HEDLEY_ALWAYS_INLINE
void expand_one(const Node * HEDLEY_RESTRICT parents, std::size_t nkeys,
bool is_last, const Node * HEDLEY_RESTRICT cw, bool right,
Node * HEDLEY_RESTRICT dest)
{
Node discard;
if (right)
expand_keys<Tree>(parents, nkeys, is_last, cw, cw, 1, &discard, dest);
else
expand_keys<Tree>(parents, nkeys, is_last, cw, cw, 0, dest, &discard);
}
template <typename Range>
void require_keys(const Range & keys)
{
if (keys.size() == 0)
throw std::invalid_argument("cohort: no keys");
}
template <typename Key, typename Integral>
std::size_t nodes_at_level(std::size_t level, Integral from_node, Integral to_node)
{
const std::size_t offset = Key::depth - level;
return static_cast<std::size_t>(
utils::shift_right(static_cast<Integral>(to_node - Integral{1}), offset)
- utils::shift_right(from_node, offset)) + 1;
}
template <typename NodeT, typename Output, typename Leaf>
HEDLEY_ALWAYS_INLINE
Output lane_value(const Leaf & leaf, std::size_t lane)
{
if constexpr (utils::is_packed_subbyte_v<Output>)
{
return extract_leaf<NodeT, Output>(leaf, lane);
}
else
{
Output val;
std::memcpy(&val,
reinterpret_cast<const unsigned char *>(std::addressof(leaf))
+ lane * sizeof(Output),
sizeof(Output));
return val;
}
}
template <typename Output, typename W>
HEDLEY_ALWAYS_INLINE
Output mac_add(Output acc, Output val, W && w)
{
if constexpr (std::is_integral_v<Output> && !std::is_same_v<Output, bool>)
{
using U = std::make_unsigned_t<Output>;
return static_cast<Output>(
static_cast<U>(acc)
+ static_cast<U>(val) * static_cast<U>(w));
}
else
{
return static_cast<Output>(acc + val * Output(w));
}
}
template <typename Out>
void ensure_size(Out & out, std::size_t n)
{
if (out.size() < n)
out.resize(n);
}
template <std::size_t I, typename Range, typename Input, typename Scratch, typename Fn>
void walk_point(const Range & keys, Input walk_x, Input lane_x,
Scratch & scratch, Fn && fn)
{
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
using node = typename key_type::interior_node;
using tree = typename key_type::tree;
using output = typename key_type::template concrete_output_type<I>;
const std::size_t n = keys.size();
scratch.fit(n, 1);
int src = 0;
auto * cur = scratch.plane(src);
for (std::size_t k = 0; k < n; ++k)
cur[k] = keys[k].root();
auto mask = key_type::msb_mask;
for (std::size_t level = 0; level < key_type::depth; ++level, mask >>= 1)
{
const bool bit = !!(mask & walk_x);
const bool is_last = tree::is_last_level(level, key_type::depth);
node * cw = bit ? scratch.cw1() : scratch.cw0();
for (std::size_t k = 0; k < n; ++k)
cw[k] = keys[k].correction_word(level, bit);
const int dst = 1 - src;
node * next = scratch.plane(dst);
if (bit)
{
expand_one<tree>(cur, n, is_last, cw, true, next);
}
else
{
expand_one<tree>(cur, n, is_last, cw, false, next);
}
src = dst;
cur = next;
}
for (std::size_t k = 0; k < n; ++k)
{
auto leaf = keys[k].template traverse_exterior<I>(cur[k]);
auto handle = make_dpf_output<output>(leaf, lane_x);
fn(k, static_cast<output>(handle));
}
}
template <std::size_t I, typename Range, typename Integral, typename Scratch, typename Fn>
void walk_interval_segment(const Range & keys, Integral from_node, Integral to_node,
std::size_t leaf_base, Scratch & scratch, Fn && on_leaf)
{
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
using node = typename key_type::interior_node;
using tree = typename key_type::tree;
const std::size_t n = keys.size();
constexpr std::size_t depth = key_type::depth;
std::size_t widest = 1;
if constexpr (depth == 0)
{
widest = nodes_at_level<key_type>(0, from_node, to_node);
}
else
{
for (std::size_t level = 1; level <= depth; ++level)
{
widest = std::max(widest,
nodes_at_level<key_type>(level, from_node, to_node));
}
}
scratch.fit(n, widest);
const std::size_t stride = scratch.stride();
int src = 0;
for (std::size_t k = 0; k < n; ++k)
scratch.plane(0)[k] = keys[k].root();
for (std::size_t level = 1; level <= depth; ++level)
{
const std::size_t child_nodes
= nodes_at_level<key_type>(level, from_node, to_node);
const auto mask = utils::get_node_mask<key_type>(key_type::msb_mask, level);
const bool from_offset = static_cast<bool>(mask & from_node);
const bool to_offset = from_offset ^ static_cast<bool>(child_nodes & 1u);
const bool is_last = tree::is_last_level(level - 1, depth);
node * cw_l = scratch.cw0();
node * cw_r = scratch.cw1();
for (std::size_t k = 0; k < n; ++k)
{
cw_l[k] = keys[k].correction_word(level - 1, false);
cw_r[k] = keys[k].correction_word(level - 1, true);
}
const int dst = 1 - src;
node * prev = scratch.plane(src);
node * curr = scratch.plane(dst);
std::size_t i = 0;
std::size_t j = 0;
if (from_offset)
{
expand_one<tree>(prev + j * stride, n, is_last, cw_r, true,
curr + i * stride);
++i;
++j;
}
const std::size_t both_end = child_nodes - static_cast<std::size_t>(to_offset);
while (i < both_end)
{
expand_keys<tree>(prev + j * stride, n, is_last, cw_l, cw_r, 2,
curr + i * stride, curr + (i + 1) * stride);
i += 2;
++j;
}
if (to_offset)
{
expand_one<tree>(prev + j * stride, n, is_last, cw_l, false,
curr + i * stride);
}
src = dst;
}
const std::size_t leaves = (depth == 0)
? nodes_at_level<key_type>(0, from_node, to_node)
: nodes_at_level<key_type>(depth, from_node, to_node);
node * leaf_plane = scratch.plane(src);
for (std::size_t j = 0; j < leaves; ++j)
{
for (std::size_t k = 0; k < n; ++k)
{
auto leaf = keys[k].template traverse_exterior<I>(
leaf_plane[j * stride + k]);
on_leaf(leaf_base + j, k, leaf);
}
}
}
template <std::size_t I, typename Range, typename Scratch, typename Fn>
void walk_recipe(const Range & keys, const sequence_recipe & recipe,
Scratch & scratch, Fn && on_point)
{
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
using node = typename key_type::interior_node;
using tree = typename key_type::tree;
using output = typename key_type::template concrete_output_type<I>;
if (recipe.depth() != key_type::depth)
throw std::invalid_argument("cohort: recipe depth does not match the keys");
const std::size_t n = keys.size();
const auto nout = recipe.output_indices().size();
if (nout == 0 || recipe.num_leaf_nodes() == 0)
return;
scratch.fit(n, std::max(recipe.num_leaf_nodes(), std::size_t{1}));
const std::size_t stride = scratch.stride();
int src = 0;
for (std::size_t k = 0; k < n; ++k)
scratch.plane(0)[k] = keys[k].root();
std::size_t step = 0;
for (std::size_t level = 1; level <= key_type::depth; ++level)
{
const std::size_t step_end = recipe.level_endpoints()[level];
const bool is_last = tree::is_last_level(level - 1, key_type::depth);
node * cw_l = scratch.cw0();
node * cw_r = scratch.cw1();
for (std::size_t k = 0; k < n; ++k)
{
cw_l[k] = keys[k].correction_word(level - 1, false);
cw_r[k] = keys[k].correction_word(level - 1, true);
}
const int dst = 1 - src;
node * prev = scratch.plane(src);
node * curr = scratch.plane(dst);
std::size_t parent_i = 0;
std::size_t out_i = 0;
for (; step < step_end; ++step, ++parent_i)
{
const std::int8_t s = recipe.recipe_steps()[step];
const bool left = s > std::int8_t{-1};
const bool right = s < std::int8_t{1};
node * parent = prev + parent_i * stride;
if (left && right)
{
expand_keys<tree>(parent, n, is_last, cw_l, cw_r, 2,
curr + out_i * stride, curr + (out_i + 1) * stride);
out_i += 2;
}
else if (left)
{
expand_one<tree>(parent, n, is_last, cw_l, false,
curr + out_i * stride);
++out_i;
}
else
{
expand_one<tree>(parent, n, is_last, cw_r, true,
curr + out_i * stride);
++out_i;
}
}
src = dst;
}
node * leaf_plane = scratch.plane(src);
constexpr std::size_t opl = key_type::outputs_per_leaf;
const auto & idx = recipe.output_indices();
std::size_t cached = std::numeric_limits<std::size_t>::max();
std::vector<decltype(keys[0].template traverse_exterior<I>(node{}))> leaves(n);
for (std::size_t q = 0; q < nout; ++q)
{
const std::size_t slot = idx[q];
const std::size_t node_i = slot / opl;
const std::size_t lane = slot % opl;
if (node_i != cached)
{
for (std::size_t k = 0; k < n; ++k)
{
leaves[k] = keys[k].template traverse_exterior<I>(
leaf_plane[node_i * stride + k]);
}
cached = node_i;
}
for (std::size_t k = 0; k < n; ++k)
on_point(q, k, lane_value<typename key_type::exterior_node, output>(
leaves[k], lane));
}
}
template <typename InteriorPRG, typename ExteriorPRG, typename Input, typename Payload>
struct cohort_dpf_type
{
using type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, Input, Payload>;
};
template <typename InteriorPRG, typename ExteriorPRG, typename Input,
typename T0, typename... Ts>
struct cohort_key_pack
{
using type = dpf_key<InteriorPRG, ExteriorPRG, Input, T0, Ts...>;
};
template <typename InteriorPRG, typename ExteriorPRG, typename Input, typename... Ts>
struct cohort_dpf_type<InteriorPRG, ExteriorPRG, Input, std::tuple<Ts...>>
{
using type = typename cohort_key_pack<InteriorPRG, ExteriorPRG, Input, Ts...>::type;
};
} // namespace cohort_detail
/// @brief Evaluate every key at `x`. `out[key]` is that key's raw share of output `I`.
/// \complexity O(n m) PRG calls. m is the number of keys and n is `depth`. One batched expand per level. The scratch holds O(m) nodes.
template <std::size_t I = 0,
typename Range,
typename InputT,
typename Out>
void eval_point_cohort(const Range & keys, InputT x, Out & out,
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using stored = std::decay_t<decltype(keys[0])>;
using key_type = unwrap_party_key_t<stored>;
using output = typename key_type::template concrete_output_type<I>;
cohort_scratch<stored> local;
auto & ws = scratch != nullptr ? *scratch : local;
auto lane_x = keys[0].offset_x(x);
auto walk_x = lane_x;
utils::flip_msb_if_signed_integral(walk_x);
cohort_detail::ensure_size(out, keys.size());
cohort_detail::walk_point<I>(keys, walk_x, lane_x, ws,
[&](std::size_t k, output v) { out[k] = v; });
}
/// @brief Evaluate every key on the closed interval `[from, to]`.
/// @details Leaves are interleaved. See the file comment for the index.
/// \complexity Same node count as one `eval_interval`, times m keys. Each level batches the PRG across keys. Scratch holds O(m L) nodes, L the widest level of the interval.
template <std::size_t I = 0,
typename Range,
typename InputT,
typename Out>
void eval_interval_cohort(const Range & keys, InputT from, InputT to, Out & out,
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using stored = std::decay_t<decltype(keys[0])>;
using key_type = unwrap_party_key_t<stored>;
using output = typename key_type::template concrete_output_type<I>;
using integral = typename key_type::integral_type;
cohort_scratch<stored> local;
auto & ws = scratch != nullptr ? *scratch : local;
auto from_x = keys[0].offset_x(from);
auto to_x = keys[0].offset_x(to);
utils::flip_msb_if_signed_integral(from_x);
utils::flip_msb_if_signed_integral(to_x);
const integral from_node = utils::get_from_node<key_type>(from_x);
const integral to_node = utils::get_to_node<key_type>(to_x);
constexpr auto to_int = utils::to_integral_type<decltype(from_x)>{};
const bool wraps = utils::interval_wraps(
static_cast<integral>(to_int(from_x)),
static_cast<integral>(to_int(to_x)),
utils::bitlength_of_v<decltype(from_x)>);
auto segs = utils::split_leaf_nodes(from_node, to_node, key_type::depth, wraps);
constexpr std::size_t opl = key_type::outputs_per_leaf;
const std::size_t n = keys.size();
cohort_detail::ensure_size(out, segs.total * opl * n);
std::size_t leaf_base = 0;
for (std::size_t s = 0; s < segs.n; ++s)
{
const auto & seg = segs.seg[s];
cohort_detail::walk_interval_segment<I>(keys, seg.from_node,
seg.to_node, leaf_base, ws,
[&](std::size_t leaf_i, std::size_t k, const auto & leaf) {
if constexpr (utils::is_packed_subbyte_v<output>)
{
for (std::size_t p = 0; p < opl; ++p)
{
out[(leaf_i * n + k) * opl + p]
= cohort_detail::lane_value<
typename key_type::exterior_node, output>(leaf, p);
}
}
else
{
std::memcpy(&out[(leaf_i * n + k) * opl],
std::addressof(leaf), sizeof(output) * opl);
}
});
leaf_base += seg.count;
}
}
/// @brief `eval_interval_cohort` from `min` through `max`.
/// \complexity Same as `eval_interval_cohort` on the full domain.
template <std::size_t I = 0, typename Range, typename Out>
void eval_full_cohort(const Range & keys, Out & out,
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using input = typename unwrap_party_key_t<
std::decay_t<decltype(keys[0])>>::input_type;
eval_interval_cohort<I>(keys, std::numeric_limits<input>::min(),
std::numeric_limits<input>::max(), out, scratch);
}
/// @brief Dot every key's interval leaves with `weights`.
/// @details `weights[j]` matches one-key `eval_inner_product`: the j-th lane
/// of the covering leaves, not a clipped sub-lane. One sum per key.
/// \complexity Same interior walk as `eval_interval_cohort`, plus one multiply-add per lane per key. No leaf buffer.
template <std::size_t I = 0, typename Range, typename InputT, typename Weights>
HEDLEY_WARN_UNUSED_RESULT
auto eval_interval_inner_product_cohort(const Range & keys, InputT from, InputT to,
Weights && weights,
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using stored = std::decay_t<decltype(keys[0])>;
using key_type = unwrap_party_key_t<stored>;
using output = typename key_type::template concrete_output_type<I>;
using integral = typename key_type::integral_type;
cohort_scratch<stored> local;
auto & ws = scratch != nullptr ? *scratch : local;
const std::size_t n = keys.size();
std::vector<output> acc(n);
auto from_x = keys[0].offset_x(from);
auto to_x = keys[0].offset_x(to);
utils::flip_msb_if_signed_integral(from_x);
utils::flip_msb_if_signed_integral(to_x);
const integral from_node = utils::get_from_node<key_type>(from_x);
const integral to_node = utils::get_to_node<key_type>(to_x);
constexpr auto to_int = utils::to_integral_type<decltype(from_x)>{};
const bool wraps = utils::interval_wraps(
static_cast<integral>(to_int(from_x)),
static_cast<integral>(to_int(to_x)),
utils::bitlength_of_v<decltype(from_x)>);
auto segs = utils::split_leaf_nodes(from_node, to_node, key_type::depth, wraps);
constexpr std::size_t opl = key_type::outputs_per_leaf;
std::size_t leaf_base = 0;
for (std::size_t s = 0; s < segs.n; ++s)
{
const auto & seg = segs.seg[s];
cohort_detail::walk_interval_segment<I>(keys, seg.from_node,
seg.to_node, leaf_base, ws,
[&](std::size_t leaf_i, std::size_t k, const auto & leaf) {
for (std::size_t p = 0; p < opl; ++p)
{
const auto val = cohort_detail::lane_value<
typename key_type::exterior_node, output>(leaf, p);
const std::size_t w = leaf_i * opl + p;
acc[k] = cohort_detail::mac_add(acc[k], val, weights[w]);
}
});
leaf_base += seg.count;
}
return acc;
}
/// @brief Evaluate every key on one compiled recipe.
/// @details `out[q * n + k]` is key `k` at listed point `q`.
/// \complexity One recipe traversal. Each visited node is expanded for all m keys together. Scratch holds O(m L) nodes, L the recipe's leaf count.
template <std::size_t I = 0, typename Range, typename Out>
void eval_sequence_cohort(const Range & keys, const sequence_recipe & recipe,
Out & out, cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using stored = std::decay_t<decltype(keys[0])>;
using key_type = unwrap_party_key_t<stored>;
using output = typename key_type::template concrete_output_type<I>;
cohort_scratch<stored> local;
auto & ws = scratch != nullptr ? *scratch : local;
const std::size_t n = keys.size();
cohort_detail::ensure_size(out, recipe.output_indices().size() * n);
cohort_detail::walk_recipe<I>(keys, recipe, ws,
[&](std::size_t q, std::size_t k, output v) {
out[cohort_index(q, k, n)] = v;
});
}
/// @brief Compile `[begin, end)` once, then `eval_sequence_cohort` on that recipe.
/// @throws std::runtime_error if the range is not sorted nondecreasing.
/// \complexity One recipe build, O(k log n) in the point list, then the recipe walk.
template <std::size_t I = 0, typename Range, typename ForwardIterator, typename Out>
void eval_sequence_cohort(const Range & keys, ForwardIterator begin,
ForwardIterator end, Out & out,
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
auto recipe = make_sequence_recipe<key_type>(begin, end);
eval_sequence_cohort<I>(keys, recipe, out, scratch);
}
/// @brief Inner product of a recipe's listed points. `weights[q]` pairs with point `q`.
/// \complexity Same walk as `eval_sequence_cohort`, plus one multiply-add per point per key.
template <std::size_t I = 0, typename Range, typename Weights>
HEDLEY_WARN_UNUSED_RESULT
auto eval_sequence_inner_product_cohort(const Range & keys,
const sequence_recipe & recipe, Weights && weights,
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using stored = std::decay_t<decltype(keys[0])>;
using key_type = unwrap_party_key_t<stored>;
using output = typename key_type::template concrete_output_type<I>;
cohort_scratch<stored> local;
auto & ws = scratch != nullptr ? *scratch : local;
std::vector<output> acc(keys.size());
cohort_detail::walk_recipe<I>(keys, recipe, ws,
[&](std::size_t q, std::size_t k, output v) {
acc[k] = cohort_detail::mac_add(acc[k], v, weights[q]);
});
return acc;
}
/// @brief Compile `[begin, end)` once, then the recipe inner product.
/// @throws std::runtime_error if the range is not sorted nondecreasing.
/// \complexity One recipe build plus the recipe inner product.
template <std::size_t I = 0, typename Range, typename ForwardIterator, typename Weights>
HEDLEY_WARN_UNUSED_RESULT
auto eval_sequence_inner_product_cohort(const Range & keys, ForwardIterator begin,
ForwardIterator end, Weights && weights,
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
{
cohort_detail::require_keys(keys);
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
auto recipe = make_sequence_recipe<key_type>(begin, end);
return eval_sequence_inner_product_cohort<I>(keys, recipe,
std::forward<Weights>(weights), scratch);
}
/// @brief One party's keys plus the scratch those walks reuse.
template <typename Key>
class cohort
{
public:
using key_type = std::decay_t<Key>;
cohort() = default;
explicit cohort(std::vector<key_type> keys) : keys_(std::move(keys)) {}
std::size_t size() const noexcept { return keys_.size(); }
const std::vector<key_type> & keys() const noexcept { return keys_; }
std::vector<key_type> & keys() noexcept { return keys_; }
cohort_scratch<key_type> & scratch() noexcept { return scratch_; }
template <std::size_t I = 0, typename InputT, typename Out>
void eval_point(InputT x, Out & out)
{
eval_point_cohort<I>(keys_, x, out, &scratch_);
}
template <std::size_t I = 0, typename InputT, typename Out>
void eval_interval(InputT from, InputT to, Out & out)
{
eval_interval_cohort<I>(keys_, from, to, out, &scratch_);
}
template <std::size_t I = 0, typename Out>
void eval_full(Out & out)
{
eval_full_cohort<I>(keys_, out, &scratch_);
}
template <std::size_t I = 0, typename InputT, typename Weights>
auto eval_interval_inner_product(InputT from, InputT to, Weights && weights)
{
return eval_interval_inner_product_cohort<I>(keys_, from, to,
std::forward<Weights>(weights), &scratch_);
}
template <std::size_t I = 0, typename Out>
void eval_sequence(const sequence_recipe & recipe, Out & out)
{
eval_sequence_cohort<I>(keys_, recipe, out, &scratch_);
}
template <std::size_t I = 0, typename ForwardIterator, typename Out>
void eval_sequence(ForwardIterator begin, ForwardIterator end, Out & out)
{
eval_sequence_cohort<I>(keys_, begin, end, out, &scratch_);
}
template <std::size_t I = 0, typename Weights>
auto eval_sequence_inner_product(const sequence_recipe & recipe,
Weights && weights)
{
return eval_sequence_inner_product_cohort<I>(keys_, recipe,
std::forward<Weights>(weights), &scratch_);
}
template <std::size_t I = 0, typename ForwardIterator, typename Weights>
auto eval_sequence_inner_product(ForwardIterator begin, ForwardIterator end,
Weights && weights)
{
return eval_sequence_inner_product_cohort<I>(keys_, begin, end,
std::forward<Weights>(weights), &scratch_);
}
private:
std::vector<key_type> keys_;
cohort_scratch<key_type> scratch_;
};
/// @brief Classic keys for one point and many payloads.
/// @details The path bit is shared. Each level expands every key's seeds with
/// `expand_x4`, then writes that key's correction word. Payloads may
/// be a single output or a `std::tuple` of outputs. Comparison,
/// incremental, verifiable, and extractable tags are not part of this
/// walk; build those with `make_dpf` one key at a time.
/// @param x plaintext domain point
/// @param begin first payload
/// @param end past the last payload
/// @return party-0 cohort and party-1 cohort, in payload order
/// \complexity O(n m) PRG calls. m is the number of payloads and n is `depth`. Roots are sampled first; each level then expands contiguous seeds.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename Iter>
HEDLEY_WARN_UNUSED_RESULT
auto make_dpf_cohort(InputT x, Iter begin, Iter end)
{
using input_type = std::decay_t<InputT>;
static_assert(!is_secret_share_v<input_type>,
"make_dpf_cohort: domain point must be plaintext");
using payload = std::decay_t<decltype(*begin)>;
static_assert(cohort_detail::is_std_tuple<payload>::value
|| !is_secret_share_v<payload>,
"make_dpf_cohort: payloads must be plaintext");
std::vector<payload> ys(begin, end);
if (ys.empty())
throw std::invalid_argument("make_dpf_cohort: no payloads");
using dpf_type = typename cohort_detail::cohort_dpf_type<
InteriorPRG, ExteriorPRG, input_type, payload>::type;
using node = typename dpf_type::interior_node;
using tree = typename dpf_type::tree;
using words = typename dpf_type::correction_words_array;
using advice = typename dpf_type::correction_advice_array;
using alloc = aligned_allocator<node>;
utils::flip_msb_if_signed_integral(x);
const std::size_t m = ys.size();
constexpr std::size_t depth = dpf_type::depth;
std::vector<node, alloc> s0(m), s1(m), root0(m), root1(m);
std::vector<words> cws(m);
std::vector<advice> adv(m);
for (std::size_t k = 0; k < m; ++k)
{
node tmp[2];
tree::root_init(tmp, []() -> node {
return static_cast<node>(dpf::uniform_sample<node>());
});
root0[k] = s0[k] = tmp[0];
root1[k] = s1[k] = tmp[1];
}
auto mask = dpf_type::msb_mask;
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{
const bool bit = !!(mask & x);
const bool is_last = tree::is_last_level(level, depth);
std::size_t k = 0;
auto step = [&](std::size_t i) {
const auto kids0 = tree::expand(s0[i], is_last);
const auto kids1 = tree::expand(s1[i], is_last);
const bool c0 = static_cast<bool>(dpf::get_lo_bit(s0[i]));
const bool c1 = static_cast<bool>(dpf::get_lo_bit(s1[i]));
tree::make_cw(cws[i][level], adv[i][level], kids0, kids1,
s0[i], s1[i], bit, is_last);
const node n0 = tree::advance(s0[i], kids0, cws[i][level],
adv[i][level], bit, c0, is_last);
const node n1 = tree::advance(s1[i], kids1, cws[i][level],
adv[i][level], bit, c1, is_last);
s0[i] = n0;
s1[i] = n1;
};
for (; k + 4 <= m; k += 4)
{
alignas(node) node l0[4], r0[4], l1[4], r1[4];
tree::expand_x4(s0.data() + k, l0, r0, is_last);
tree::expand_x4(s1.data() + k, l1, r1, is_last);
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
{
const std::size_t i = k + t;
const std::array<node, 2> kids0{l0[t], r0[t]};
const std::array<node, 2> kids1{l1[t], r1[t]};
const bool c0 = static_cast<bool>(dpf::get_lo_bit(s0[i]));
const bool c1 = static_cast<bool>(dpf::get_lo_bit(s1[i]));
tree::make_cw(cws[i][level], adv[i][level], kids0, kids1,
s0[i], s1[i], bit, is_last);
const node n0 = tree::advance(s0[i], kids0, cws[i][level],
adv[i][level], bit, c0, is_last);
const node n1 = tree::advance(s1[i], kids1, cws[i][level],
adv[i][level], bit, c1, is_last);
s0[i] = n0;
s1[i] = n1;
}
}
for (; k < m; ++k)
step(k);
}
using party0 = party_key<0, dpf_type>;
using party1 = party_key<1, dpf_type>;
std::vector<party0> k0;
std::vector<party1> k1;
k0.reserve(m);
k1.reserve(m);
for (std::size_t i = 0; i < m; ++i)
{
const bool sign0 = static_cast<bool>(dpf::get_lo_bit(s0[i]));
const node seed0 = dpf::unset_lo_2bits(s0[i]);
const node seed1 = dpf::unset_lo_2bits(s1[i]);
auto built = [&]() {
if constexpr (cohort_detail::is_std_tuple<payload>::value)
{
return std::apply([&](const auto & ...p) {
return dpf::make_leaves<ExteriorPRG>(x, seed0, seed1, sign0,
std::size_t{0}, p...);
}, ys[i]);
}
else
{
return dpf::make_leaves<ExteriorPRG>(x, seed0, seed1, sign0,
std::size_t{0}, ys[i]);
}
}();
auto paired = dpf::make_party_key_pair(
dpf_type{root0[i], cws[i], adv[i], built.first.first,
built.first.second, input_type{}},
dpf_type{root1[i], cws[i], adv[i], built.second.first,
built.second.second, input_type{}});
k0.push_back(std::move(paired.first));
k1.push_back(std::move(paired.second));
}
return std::make_pair(cohort<party0>(std::move(k0)),
cohort<party1>(std::move(k1)));
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_COHORT_HPP__