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>
This commit is contained in:
parent
695f8e84f7
commit
0d22946a0e
1835 changed files with 170291 additions and 2849 deletions
911
include/dpf/cohort.hpp
Normal file
911
include/dpf/cohort.hpp
Normal file
|
|
@ -0,0 +1,911 @@
|
|||
/// @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__
|
||||
Loading…
Add table
Add a link
Reference in a new issue