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

1143 lines
47 KiB
C++
Raw Permalink 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/grow.hpp
/// @brief Dealer-side growth of an existing `dpf_key`: one more level, or one
/// more output on a level the tree already has.
/// @details `extend` and `add_output` take both parties' keys (same dealer view
/// as `make_dpf`), return a new `party_key` pair whose type is the old
/// key plus the new material, and keep earlier correction words /
/// leaves / comparison words. Specs are the same objects `make_dpf`
/// accepts. Memoizer overloads skip the rewalk when both path
/// memoizers are already filled through the frontier. Interactive
/// Doerner–Shelat growth is `extend_ds` / `add_output_ds` in
/// grow_ds.hpp.
/// @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_GROW_HPP__
#define LIBDPF_INCLUDE_DPF_GROW_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <tuple>
#include <type_traits>
#include <utility>
#include "hedley/hedley.h"
#include "dpf/dpf_key.hpp"
#include "dpf/incremental.hpp"
#include "dpf/path_memoizer.hpp"
#include "dpf/placement.hpp"
#include "dpf/secret_share.hpp"
#include "dpf/tree_traits.hpp"
#include "dpf/twiddle.hpp"
#include "dpf/utils.hpp"
#include "dpf/verifiable.hpp"
namespace dpf
{
namespace detail
{
namespace grow_impl
{
template <typename T, typename = void>
struct has_placed_tuple : std::false_type
{
};
template <typename T>
struct has_placed_tuple<T, std::void_t<typename T::placed_tuple>>
: std::true_type
{
};
template <typename T>
inline constexpr bool has_placed_tuple_v =
has_placed_tuple<std::decay_t<T>>::value;
template <typename Key>
using bare_key_t = unwrap_party_key_t<Key>;
template <typename Node, typename Tree>
HEDLEY_ALWAYS_INLINE
bool try_advance(Node & s0, Node & s1, Node cw, psnip_uint8_t advice, bool bit)
{
const bool t0 = static_cast<bool>(get_lo_bit(s0));
const bool t1 = static_cast<bool>(get_lo_bit(s1));
const bool is_last = false;
const auto kids0 = Tree::expand(s0, is_last);
const auto kids1 = Tree::expand(s1, is_last);
const Node n0 = Tree::advance(s0, kids0, cw, advice, bit, t0, is_last);
const Node n1 = Tree::advance(s1, kids1, cw, advice, bit, t1, is_last);
if (get_lo_bit(n0) == get_lo_bit(n1))
return false;
s0 = n0;
s1 = n1;
return true;
}
/// @brief Rewalk Gen to `stop` using stored CWs. Path bits are recovered: only
/// the programmed direction keeps the parties' control bits distinct.
template <typename Key>
void rewalk_to(const Key & k0, const Key & k1,
typename Key::interior_node & s0, typename Key::interior_node & s1,
std::size_t stop, bool * path_out = nullptr)
{
using tree = typename Key::tree;
using node = typename Key::interior_node;
s0 = k0.root();
s1 = k1.root();
const auto & cws = k0.correction_words();
const auto & adv = k0.correction_advice();
for (std::size_t level = 0; level < stop; ++level)
{
bool found = false;
node a0 = s0;
node a1 = s1;
if (try_advance<node, tree>(a0, a1, cws[level], adv[level], false))
{
s0 = a0;
s1 = a1;
if (path_out)
path_out[level] = false;
found = true;
}
else
{
a0 = s0;
a1 = s1;
if (try_advance<node, tree>(a0, a1, cws[level], adv[level], true))
{
s0 = a0;
s1 = a1;
if (path_out)
path_out[level] = true;
found = true;
}
}
if (!found)
throw std::logic_error("grow: path bit recovery failed");
}
}
template <typename InputT>
HEDLEY_NO_THROW
constexpr bool bit_at(InputT x, std::size_t level) noexcept
{
constexpr auto bitlen = utils::bitlength_of_v<InputT>;
const auto to_int = utils::to_integral_type<InputT>{};
using I = typename utils::to_integral_type<InputT>::integral_type;
const I xi = static_cast<I>(to_int(x));
return static_cast<bool>(
utils::shift_right(xi, bitlen - 1 - level) & I{1});
}
/// @brief Read both parties' seeds at `need` from path memoizers. The caller
/// must already have filled each memo through `need` (e.g. via
/// `eval_point` / `ensure_level`); this does not walk the tree.
template <typename Key, typename Memo0, typename Memo1>
void frontier_from_memos(const Key & k0, const Key & k1, Memo0 & m0, Memo1 & m1,
typename Key::input_type xx, std::size_t need,
typename Key::interior_node & s0, typename Key::interior_node & s1,
bool * path_out)
{
if (need > Key::depth)
throw std::invalid_argument("grow: memoizer level past key depth");
static_assert(detail::has_path_high_water<Memo0>::value
&& detail::has_path_high_water<Memo1>::value,
"grow: memoizer overloads require a path memoizer with filled_to");
if (m0.filled_to() < need)
throw std::invalid_argument(
"grow: party-0 memoizer is not filled to the required level");
if (m1.filled_to() < need)
throw std::invalid_argument(
"grow: party-1 memoizer is not filled to the required level");
(void)k0;
(void)k1;
for (std::size_t level = 0; level < Key::depth && path_out != nullptr;
++level)
path_out[level] = bit_at(xx, level);
s0 = m0[need];
s1 = m1[need];
}
template <typename Memo0, typename Memo1, typename Node>
void read_snap_from_memos(Memo0 & m0, Memo1 & m1, std::size_t depth, Node * snap0,
Node * snap1, bool * snap_sign)
{
for (std::size_t lvl = 0; lvl <= depth; ++lvl)
{
snap0[lvl] = m0[lvl];
snap1[lvl] = m1[lvl];
snap_sign[lvl] = static_cast<bool>(get_lo_bit(snap0[lvl]));
}
}
template <typename OldPlaced, typename NewPlaced, std::size_t... OldIs,
std::size_t... NewIs>
auto cat_placed(const OldPlaced &, NewPlaced && neu,
std::index_sequence<OldIs...>, std::index_sequence<NewIs...>)
{
// Old payload values are not stored on the key; only types matter for the
// concatenated placed tuple used to name the new key. Dummy defaults.
return std::tuple_cat(
std::make_tuple(std::tuple_element_t<OldIs, OldPlaced>{}...),
std::make_tuple(std::get<NewIs>(std::forward<NewPlaced>(neu))...));
}
template <typename OldWraps, typename NewWraps, std::size_t... OldIs,
std::size_t... NewIs>
auto cat_wrappers(OldWraps && old_w, NewWraps && new_w,
std::index_sequence<OldIs...>, std::index_sequence<NewIs...>)
{
return std::make_tuple(
std::get<OldIs>(std::forward<OldWraps>(old_w))...,
std::get<NewIs>(std::forward<NewWraps>(new_w))...);
}
template <typename OldAdds, typename NewAdds, std::size_t... OldIs,
std::size_t... NewIs>
auto cat_addends(const OldAdds & old_a, NewAdds && new_a,
std::index_sequence<OldIs...>, std::index_sequence<NewIs...>)
{
return std::make_tuple(
std::get<OldIs>(old_a)...,
std::get<NewIs>(std::forward<NewAdds>(new_a))...);
}
template <typename OldKey, typename NewKey>
void copy_cmp_common(const OldKey & src, typename NewKey::value_cw_array & vcw,
typename NewKey::value_cw_array & vcoeff,
typename NewKey::tail_array & tail,
typename NewKey::tail_array & tcoeff,
typename NewKey::prefix_cw_array & prefix,
typename NewKey::prefix_cw_array & pcoeff,
typename NewKey::value_cw_word & cw_last,
typename NewKey::value_cw_word & cw_last_coeff)
{
if constexpr (OldKey::cmp_depth > 0 && NewKey::cmp_depth > 0)
{
const auto & old_vcw = src.value_cw();
const auto & old_coeff = src.value_cw_coeff();
constexpr std::size_t ncopy =
OldKey::value_cw_len < NewKey::value_cw_len
? OldKey::value_cw_len
: NewKey::value_cw_len;
for (std::size_t i = 0; i < ncopy; ++i)
{
vcw[i] = old_vcw[i];
if constexpr (OldKey::cmp_is_wildcard && NewKey::cmp_is_wildcard)
vcoeff[i] = old_coeff[i];
}
const auto & ot = src.tail_cw();
const auto & otc = src.tail_coeff();
constexpr std::size_t nt =
OldKey::cmp_tail < NewKey::cmp_tail ? OldKey::cmp_tail
: NewKey::cmp_tail;
for (std::size_t i = 0; i < nt; ++i)
{
tail[i] = ot[i];
if constexpr (OldKey::cmp_is_wildcard && NewKey::cmp_is_wildcard)
tcoeff[i] = otc[i];
}
if constexpr (OldKey::cmp_idcf && NewKey::cmp_idcf)
{
const auto & op = src.prefix_cws();
const auto & opc = src.prefix_cw_coeff();
constexpr std::size_t np =
OldKey::prefix_cw_len < NewKey::prefix_cw_len
? OldKey::prefix_cw_len
: NewKey::prefix_cw_len;
for (std::size_t i = 0; i < np; ++i)
{
prefix[i] = op[i];
if constexpr (OldKey::cmp_is_wildcard && NewKey::cmp_is_wildcard)
pcoeff[i] = opc[i];
}
}
cw_last = src.cw_last_word();
cw_last_coeff = src.cw_last_coeff_word();
}
}
template <typename OldKey, typename NewPlacedPart, std::size_t NewCmpDepth,
std::size_t NewCmpOutBits, bool NewCmpWild, std::size_t NewCmpBlock,
bool NewCmpIdcf, bool NewVerifiable, bool NewExtractable>
struct grown_key_type
{
using old_placed = typename OldKey::placed_tuple;
using new_placed = decltype(std::tuple_cat(
std::declval<old_placed>(), std::declval<NewPlacedPart>()));
static constexpr std::size_t cmp_depth =
NewCmpDepth > 0 ? NewCmpDepth : OldKey::cmp_depth;
static constexpr std::size_t cmp_out_bits =
NewCmpDepth > 0 ? NewCmpOutBits : OldKey::cmp_out_bits;
static constexpr bool cmp_wild =
NewCmpDepth > 0 ? NewCmpWild : OldKey::cmp_is_wildcard;
static constexpr std::size_t cmp_block =
NewCmpDepth > 0 ? NewCmpBlock : OldKey::cmp_block;
static constexpr bool cmp_idcf =
NewCmpDepth > 0 ? NewCmpIdcf : OldKey::cmp_idcf;
static constexpr bool is_verifiable =
NewVerifiable || OldKey::is_verifiable;
static constexpr bool is_extractable =
NewExtractable || OldKey::is_extractable;
using type = incr::incr_dpf_key_of_t<typename OldKey::interior_prg,
typename OldKey::exterior_prg, typename OldKey::raw_input_type,
new_placed, cmp_depth, cmp_out_bits, cmp_wild, cmp_block, cmp_idcf,
is_verifiable, is_extractable>;
};
template <bool DoExtend, bool UseMemo, typename K0, typename K1,
typename InputT, typename Memo0, typename Memo1, typename... Specs>
auto grow_impl(const K0 & pk0, const K1 & pk1, Memo0 * m0, Memo1 * m1, bool bit,
InputT x, bool have_pre,
typename bare_key_t<K0>::interior_node * pre_cw, psnip_uint8_t * pre_advice,
typename bare_key_t<K0>::interior_node * pre_s0,
typename bare_key_t<K0>::interior_node * pre_s1, Specs &&... specs)
{
using old_key = bare_key_t<K0>;
static_assert(std::is_same_v<old_key, bare_key_t<K1>>,
"grow: both keys must have the same type");
static_assert(has_placed_tuple_v<old_key>,
"grow: key must be a multi-level / incremental dpf_key");
static_assert(old_key::num_outputs > 0 || old_key::cmp_depth > 0,
"grow: empty key");
using input_type = typename old_key::input_type;
static_assert(std::is_same_v<std::decay_t<InputT>, input_type>
|| std::is_convertible_v<InputT, input_type>,
"grow: programmed point type mismatch");
constexpr auto bitlen = utils::bitlength_of_v<input_type>;
const old_key & k0 = static_cast<const old_key &>(pk0);
const old_key & k1 = static_cast<const old_key &>(pk1);
if constexpr (old_key::depth > 0)
{
if (std::memcmp(k0.correction_words().data(),
k1.correction_words().data(),
sizeof(typename old_key::correction_words_array)) != 0
|| std::memcmp(k0.correction_advice().data(),
k1.correction_advice().data(),
sizeof(typename old_key::correction_advice_array)) != 0)
throw std::invalid_argument("grow: keys are not a matching pair");
}
dcf_runtime_spec dcf_spec{};
bool has_cmp = false;
auto new_placed_part = detail::incr::flatten_args_and_cmp<bitlen>(dcf_spec,
has_cmp, std::forward<Specs>(specs)...);
using new_placed_part_t = decltype(new_placed_part);
constexpr std::size_t arg_cmp_depth =
forced_cmp_depth_v<bitlen, Specs...>;
constexpr std::size_t arg_cmp_bits =
forced_cmp_out_bits_v<bitlen, Specs...>;
constexpr bool arg_cmp_wild = forced_cmp_wild_v<Specs...>;
constexpr std::size_t arg_cmp_block = forced_cmp_block_v<Specs...>;
constexpr bool arg_cmp_idcf = forced_cmp_idcf_v<Specs...>;
constexpr bool arg_verifiable = args_have_verifiable_v<Specs...>;
constexpr bool arg_extractable = args_have_extractable_v<Specs...>;
if constexpr (old_key::cmp_depth > 0 && arg_cmp_depth > 0)
{
static_assert(old_key::cmp_depth == arg_cmp_depth
&& old_key::cmp_block == arg_cmp_block
&& old_key::cmp_idcf == arg_cmp_idcf,
"grow: comparison channel already present with a different shape");
}
using grown = grown_key_type<old_key, new_placed_part_t, arg_cmp_depth,
arg_cmp_bits, arg_cmp_wild, arg_cmp_block, arg_cmp_idcf, arg_verifiable,
arg_extractable>;
using new_key = typename grown::type;
using new_placed = typename grown::new_placed;
using node = typename new_key::interior_node;
using tree = typename new_key::tree;
constexpr std::size_t old_n = old_key::num_outputs;
constexpr std::size_t new_n = new_key::num_outputs;
constexpr std::size_t added = new_n - old_n;
if constexpr (DoExtend)
{
static_assert(new_key::depth == old_key::depth + 1,
"extend: expected depth to grow by exactly one");
}
else
{
static_assert(new_key::depth == old_key::depth,
"add_output: comparison or output would deepen the tree; "
"use extend or a shorter lt_at/gt_at/eq_at prefix");
}
if constexpr (!DoExtend)
{
if constexpr (added == 0 && arg_cmp_depth == 0 && !arg_verifiable
&& !arg_extractable)
throw std::invalid_argument("add_output: no new material");
}
else
{
if (old_key::depth + 1 >= bitlen)
throw std::invalid_argument(
"extend: key is already as deep as the input type");
if (bit != bit_at(static_cast<input_type>(x), old_key::depth))
throw std::invalid_argument(
"extend: bit does not match the programmed point");
}
// Reject specs that belong on an older level (extend) or on a missing
// level (add_output).
if constexpr (added > 0)
{
constexpr auto new_meta = new_key::meta;
for (std::size_t i = old_n; i < new_n; ++i)
{
const auto lvl = new_meta[i].tree_level;
if constexpr (DoExtend)
{
if (lvl != new_key::depth)
throw std::invalid_argument(
"extend: output prefix is not the new depth");
}
else
{
if (lvl > old_key::depth)
throw std::invalid_argument(
"add_output: prefix needs a deeper tree");
}
}
}
input_type xx = static_cast<input_type>(x);
utils::flip_msb_if_signed_integral(xx);
node s0{};
node s1{};
std::array<bool, old_key::depth == 0 ? 1 : old_key::depth> path{};
if constexpr (UseMemo)
{
static_assert(!std::is_void_v<Memo0> && !std::is_void_v<Memo1>,
"grow: memoizer overload requires two memoizers");
frontier_from_memos(k0, k1, *m0, *m1, xx, old_key::depth, s0, s1,
old_key::depth == 0 ? nullptr : path.data());
}
else
{
rewalk_to(k0, k1, s0, s1, old_key::depth,
old_key::depth == 0 ? nullptr : path.data());
(void)m0;
(void)m1;
}
typename new_key::correction_words_array correction_words{};
typename new_key::correction_advice_array correction_advice{};
typename new_key::correction_seeds_array correction_seeds{};
for (std::size_t i = 0; i < old_key::depth; ++i)
{
correction_words[i] = k0.correction_words()[i];
correction_advice[i] = k0.correction_advice()[i];
}
if constexpr (old_key::is_verifiable && new_key::is_verifiable)
{
for (std::size_t i = 0; i < old_key::depth; ++i)
correction_seeds[i] = k0.correction_seeds()[i];
}
else if constexpr (!old_key::is_verifiable && new_key::is_verifiable)
{
// Opting in: sample seeds for levels that already exist.
node p0 = k0.root();
node p1 = k1.root();
for (std::size_t level = 0; level < old_key::depth; ++level)
{
const bool pb = path[level];
const bool is_last = tree::is_last_level(level, new_key::depth);
const bool a0 = static_cast<bool>(get_lo_bit(p0));
const bool a1 = static_cast<bool>(get_lo_bit(p1));
auto c0 = tree::expand(p0, is_last);
auto c1 = tree::expand(p1, is_last);
p0 = tree::advance(p0, c0, correction_words[level],
correction_advice[level], pb, a0, is_last);
p1 = tree::advance(p1, c1, correction_words[level],
correction_advice[level], pb, a1, is_last);
const auto prefix = static_cast<psnip_uint64_t>(
utils::to_integral_type<input_type>{}(xx)
>> (bitlen - (level + 1)));
if constexpr (new_key::cmp_block > 0)
{
correction_seeds[level] = detail::vdpf::make_cs(
detail::blocked::fold_spine_tag | level, prefix, p0, p1);
}
else
{
correction_seeds[level] =
detail::vdpf::make_cs(level, prefix, p0, p1);
}
}
s0 = p0;
s1 = p1;
}
if constexpr (DoExtend)
{
const std::size_t level = old_key::depth;
if (have_pre)
{
correction_words[level] = *pre_cw;
correction_advice[level] = *pre_advice;
s0 = *pre_s0;
s1 = *pre_s1;
}
else
{
const bool is_last = tree::is_last_level(level, new_key::depth);
const bool a0 = static_cast<bool>(get_lo_bit(s0));
const bool a1 = static_cast<bool>(get_lo_bit(s1));
auto c0 = tree::expand(s0, is_last);
auto c1 = tree::expand(s1, is_last);
node cw{};
psnip_uint8_t advice = 0;
tree::make_cw(cw, advice, c0, c1, s0, s1, bit, is_last);
s0 = tree::advance(s0, c0, cw, advice, bit, a0, is_last);
s1 = tree::advance(s1, c1, cw, advice, bit, a1, is_last);
correction_words[level] = cw;
correction_advice[level] = advice;
}
if constexpr (new_key::is_verifiable)
{
const auto prefix = static_cast<psnip_uint64_t>(
utils::to_integral_type<input_type>{}(xx)
>> (bitlen - (level + 1)));
if constexpr (new_key::cmp_block > 0)
{
correction_seeds[level] = detail::vdpf::make_cs(
detail::blocked::fold_spine_tag | level, prefix, s0, s1);
}
else
{
correction_seeds[level] =
detail::vdpf::make_cs(level, prefix, s0, s1);
}
}
}
else
{
(void)have_pre;
(void)pre_cw;
(void)pre_advice;
(void)pre_s0;
(void)pre_s1;
}
// Full placed tuple for naming / meta; only the new slots' values are used
// when planting leaves.
new_placed placed = cat_placed(typename old_key::placed_tuple{},
std::move(new_placed_part), std::make_index_sequence<old_n>{},
std::make_index_sequence<added>{});
auto leaves0 = detail::incr::empty_leaves<typename new_key::exterior_prg,
new_placed>(std::make_index_sequence<new_n>{});
auto beavers0 = detail::incr::empty_beavers<typename new_key::exterior_prg,
new_placed>(std::make_index_sequence<new_n>{});
auto leaves1 = detail::incr::empty_leaves<typename new_key::exterior_prg,
new_placed>(std::make_index_sequence<new_n>{});
auto beavers1 = detail::incr::empty_beavers<typename new_key::exterior_prg,
new_placed>(std::make_index_sequence<new_n>{});
// Copy old raw leaves / beavers into the front slots.
if constexpr (old_n > 0)
{
[&]<std::size_t... Is>(std::index_sequence<Is...>) {
((std::get<Is>(leaves0) = k0.template leaf<Is>(),
std::get<Is>(beavers0) = k0.template beaver<Is>(),
std::get<Is>(leaves1) = k1.template leaf<Is>(),
std::get<Is>(beavers1) = k1.template beaver<Is>()),
...);
}(std::make_index_sequence<old_n>{});
}
// When depth grows, groups that sat at the old deepest level move from
// pos_base-start-0 to pos_base-start-2. Retarget their shared correction
// words so eval of those slots stays share-identical.
if constexpr (DoExtend && old_n > 0)
{
using exterior = typename new_key::exterior_prg;
using node_ex = typename exterior::block_type;
node p0 = k0.root();
node p1 = k1.root();
// Walk to each old-deepest group's level once.
for (std::size_t level = 0; level < old_key::depth; ++level)
{
const bool pb = path[level];
const bool is_last = tree::is_last_level(level, new_key::depth);
const bool a0 = static_cast<bool>(get_lo_bit(p0));
const bool a1 = static_cast<bool>(get_lo_bit(p1));
auto c0 = tree::expand(p0, is_last);
auto c1 = tree::expand(p1, is_last);
p0 = tree::advance(p0, c0, correction_words[level],
correction_advice[level], pb, a0, is_last);
p1 = tree::advance(p1, c1, correction_words[level],
correction_advice[level], pb, a1, is_last);
}
const auto seed0 = dpf::unset_lo_2bits(p0);
const auto seed1 = dpf::unset_lo_2bits(p1);
const bool sign0 = static_cast<bool>(get_lo_bit(p0));
[&]<std::size_t... Is>(std::index_sequence<Is...>) {
(([&] {
constexpr std::size_t old_pos = old_key::meta[Is].pos_base;
constexpr std::size_t new_pos = new_key::meta[Is].pos_base;
if constexpr (old_pos == new_pos)
return;
using Out = typename new_key::template output_type_t<Is>;
using Concrete = concrete_type_t<Out>;
// Build a one-slot outputs tuple so make_leaf_mask indices match.
using outs = std::tuple<Out>;
const auto mask_old =
dpf::make_leaf_mask<exterior, 0, outs, node_ex>(
seed0, seed1, old_pos);
const auto mask_new =
dpf::make_leaf_mask<exterior, 0, outs, node_ex>(
seed0, seed1, new_pos);
auto & L0 = std::get<Is>(leaves0);
auto & L1 = std::get<Is>(leaves1);
if (sign0)
{
// L = naked - mask => L' = L + mask_old - mask_new
L0 = dpf::subtract_leaf<Concrete>(
dpf::add_leaf<Concrete>(L0, mask_old), mask_new);
L1 = dpf::subtract_leaf<Concrete>(
dpf::add_leaf<Concrete>(L1, mask_old), mask_new);
}
else
{
// L = mask - naked => L' = L - mask_old + mask_new
L0 = dpf::add_leaf<Concrete>(
dpf::subtract_leaf<Concrete>(L0, mask_old), mask_new);
L1 = dpf::add_leaf<Concrete>(
dpf::subtract_leaf<Concrete>(L1, mask_old), mask_new);
}
}()), ...);
}(std::make_index_sequence<old_n>{});
(void)sign0;
}
// Plant only groups that contain a newly added slot. A group that also
// holds an older slot would need that slot's payload to rebuild the
// packed leaf; reject that case.
if constexpr (added > 0)
{
using MetaHolder = detail::incr::meta_holder<new_key>;
constexpr auto meta = new_key::meta;
constexpr std::size_t ngroups = [] {
std::size_t m = 0;
for (std::size_t i = 0; i < new_n; ++i)
m = std::max(m, new_key::meta[i].group_id + 1);
return m;
}();
constexpr auto order =
detail::incr::build_group_order<decltype(meta), ngroups>(meta,
new_n);
std::array<node, new_key::depth + 1> snap0{};
std::array<node, new_key::depth + 1> snap1{};
std::array<bool, new_key::depth + 1> snap_sign{};
if constexpr (UseMemo)
{
// Seeds through the old frontier come from the memoizers. After an
// extend the new depth's node is already in s0/s1.
read_snap_from_memos(*m0, *m1, old_key::depth, snap0.data(),
snap1.data(), snap_sign.data());
if constexpr (DoExtend)
{
snap0[new_key::depth] = s0;
snap1[new_key::depth] = s1;
snap_sign[new_key::depth] =
static_cast<bool>(get_lo_bit(s0));
}
}
else
{
node p0 = k0.root();
node p1 = k1.root();
snap0[0] = p0;
snap1[0] = p1;
snap_sign[0] = static_cast<bool>(get_lo_bit(p0));
for (std::size_t level = 0; level < new_key::depth; ++level)
{
const bool pb = (level < old_key::depth) ? path[level] : bit;
const bool is_last =
tree::is_last_level(level, new_key::depth);
const bool a0 = static_cast<bool>(get_lo_bit(p0));
const bool a1 = static_cast<bool>(get_lo_bit(p1));
auto c0 = tree::expand(p0, is_last);
auto c1 = tree::expand(p1, is_last);
p0 = tree::advance(p0, c0, correction_words[level],
correction_advice[level], pb, a0, is_last);
p1 = tree::advance(p1, c1, correction_words[level],
correction_advice[level], pb, a1, is_last);
snap0[level + 1] = p0;
snap1[level + 1] = p1;
snap_sign[level + 1] = static_cast<bool>(get_lo_bit(p0));
}
}
detail::incr::for_each_index(std::make_index_sequence<ngroups>{},
[&](auto oi) {
constexpr std::size_t G = order[decltype(oi)::value];
constexpr bool touches_new = [] {
for (std::size_t i = old_n; i < new_n; ++i)
if (new_key::meta[i].group_id == G)
return true;
return false;
}();
constexpr bool touches_old = [] {
for (std::size_t i = 0; i < old_n; ++i)
if (new_key::meta[i].group_id == G)
return true;
return false;
}();
if constexpr (!touches_new)
return;
if constexpr (touches_old)
{
throw std::invalid_argument(
"grow: new output shares a packing group with an "
"existing slot; use make_dpf for that pack");
}
constexpr std::size_t lvl = [] {
for (std::size_t i = 0; i < new_n; ++i)
if (new_key::meta[i].group_id == G)
return new_key::meta[i].tree_level;
return std::size_t{0};
}();
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using leaf_prg = std::conditional_t<new_key::is_extractable,
detail::vdpf::extractable_leaf_prg<
typename new_key::exterior_prg>,
typename new_key::exterior_prg>;
detail::incr::gen_group<typename new_key::interior_prg,
leaf_prg, input_type, new_placed, MetaHolder, G>(xx,
dpf::unset_lo_2bits(snap0[lvl]),
dpf::unset_lo_2bits(snap1[lvl]),
static_cast<uint8_t>(snap_sign[lvl]),
static_cast<uint8_t>(dpf::get_lo_bit(snap1[lvl])), placed,
leaves0, beavers0, leaves1, beavers1);
HEDLEY_PRAGMA(GCC diagnostic pop)
});
}
if constexpr (!old_key::is_extractable && new_key::is_extractable
&& old_n > 0 && added == 0)
{
throw std::invalid_argument(
"grow: opting into extractable requires new outputs or a full "
"rebuild via make_dpf");
}
auto wrap0 = detail::incr::wrap_leaves<typename new_key::interior_prg,
typename new_key::exterior_prg, input_type, new_placed>(leaves0,
beavers0, std::make_index_sequence<new_n>{});
auto wrap1 = detail::incr::wrap_leaves<typename new_key::interior_prg,
typename new_key::exterior_prg, input_type, new_placed>(leaves1,
beavers1, std::make_index_sequence<new_n>{});
typename new_key::value_cw_array value_cws{};
typename new_key::value_cw_array value_cw_coeff{};
typename new_key::tail_array tail{};
typename new_key::tail_array tail_coeff{};
typename new_key::prefix_cw_array prefix_cw{};
typename new_key::prefix_cw_array prefix_coeff{};
typename new_key::value_cw_word cw_last{};
typename new_key::value_cw_word cw_last_coeff{};
typename new_key::value_cw_word cmp_add0{};
typename new_key::value_cw_word cmp_add1{};
copy_cmp_common<old_key, new_key>(k0, value_cws, value_cw_coeff, tail,
tail_coeff, prefix_cw, prefix_coeff, cw_last, cw_last_coeff);
if constexpr (old_key::cmp_depth > 0)
{
cmp_add0 = k0.cmp_addend_word();
cmp_add1 = k1.cmp_addend_word();
}
detail::cmp_meta cmp = k0.cmp();
// Add a comparison onto a key that had none (integral per-level path).
if (has_cmp && old_key::cmp_depth == 0)
{
using namespace detail::dcf_impl;
cmp.nbits = static_cast<int>(dcf_spec.prefix ? dcf_spec.prefix
: bitlen);
cmp.mask = dcf_spec.mask;
cmp.kind = dcf_spec.kind;
cmp.active = true;
cmp.incremental = dcf_spec.incremental;
cmp.eval_as_ge = false;
cmp.trivial = cmp_trivial::none;
cmp.block_width = static_cast<int>(new_key::cmp_block);
cmp.tail_bits = static_cast<int>(new_key::cmp_q);
cmp.include_eq = false;
const uint64_t delta = dcf_spec.beta & dcf_spec.mask;
const uint64_t false_value = dcf_spec.false_value & dcf_spec.mask;
unsigned __int128 thresh =
static_cast<unsigned __int128>(
utils::to_integral_type<input_type>{}(xx));
if (dcf_spec.prefix)
thresh >>= (bitlen - dcf_spec.prefix);
detail::incr::adjust_cmp_threshold(cmp, thresh,
static_cast<std::size_t>(cmp.nbits));
const std::size_t cmp_nbits = static_cast<std::size_t>(cmp.nbits);
if (cmp.trivial != cmp_trivial::none || new_key::cmp_block > 0
|| dcf_spec.custom || dcf_spec.use_payload_group
|| is_paint_kind(dcf_spec.kind))
{
throw std::invalid_argument(
"add_output: comparison shape not supported on grow yet");
}
uint64_t Va = 0;
uint64_t Va1 = 0;
node p0 = k0.root();
node p1 = k1.root();
for (std::size_t level = 0; level < new_key::depth; ++level)
{
const bool pb = (level < old_key::depth) ? path[level]
: (DoExtend ? bit : bit_at(xx, level));
const bool is_last = tree::is_last_level(level, new_key::depth);
const bool advice0 = static_cast<bool>(get_lo_bit(p0));
const bool advice1 = static_cast<bool>(get_lo_bit(p1));
const auto val0 = tree::expand_value(p0);
const auto val1 = tree::expand_value(p1);
auto c0 = tree::expand(p0, is_last);
auto c1 = tree::expand(p1, is_last);
if (level < cmp_nbits)
{
const int ai = static_cast<int>(
(thresh >> (cmp_nbits - 1 - level)) & 1);
const uint64_t base = make_value_cw(val0[0], val0[1],
val1[0], val1[1], static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va, delta, cmp.mask);
value_cws[level] =
static_cast<typename new_key::value_cw_word>(base);
if constexpr (new_key::cmp_is_wildcard)
{
const uint64_t v1 = make_value_cw(val0[0], val0[1],
val1[0], val1[1], static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va1, 1ULL,
cmp.mask);
value_cw_coeff[level] =
static_cast<typename new_key::value_cw_word>(
(v1 + neg_m(base, cmp.mask)) & cmp.mask);
}
}
p0 = tree::advance(p0, c0, correction_words[level],
correction_advice[level], pb, advice0, is_last);
p1 = tree::advance(p1, c1, correction_words[level],
correction_advice[level], pb, advice1, is_last);
if (level + 1 == cmp_nbits)
{
cw_last = static_cast<typename new_key::value_cw_word>(
make_final_cw(p0, p1,
static_cast<uint8_t>(get_lo_bit(p1)), Va, cmp.mask,
true));
if constexpr (new_key::cmp_is_wildcard)
{
const uint64_t l1 = make_final_cw(p0, p1,
static_cast<uint8_t>(get_lo_bit(p1)), Va1, cmp.mask,
true);
cw_last_coeff =
static_cast<typename new_key::value_cw_word>(
(l1 + neg_m(static_cast<uint64_t>(cw_last),
cmp.mask))
& cmp.mask);
}
}
}
uint64_t target = false_value;
if (cmp.eval_as_ge)
target = (delta + false_value) & cmp.mask;
const uint64_t rblind = sample_addend_blind(cmp.mask,
[] { return dpf::uniform_sample<node>(); });
uint64_t a0 = 0, a1 = 0;
detail::incr::split_cmp_addend(target, cmp.mask, rblind, a0, a1);
cmp_add0 = static_cast<typename new_key::value_cw_word>(a0);
cmp_add1 = static_cast<typename new_key::value_cw_word>(a1);
}
else if constexpr (DoExtend && old_key::cmp_depth > 0
&& new_key::cmp_block == 0 && new_key::depth > old_key::depth)
{
// Extending a key that already has a comparison: append a value word
// when the new level is still inside the comparison. Requires the
// caller to re-pass the comparison spec so δ is known.
if (has_cmp && cmp.active && cmp.trivial == cmp_trivial::none)
{
using namespace detail::dcf_impl;
const std::size_t level = old_key::depth;
const std::size_t cmp_nbits = static_cast<std::size_t>(cmp.nbits);
if (level < cmp_nbits)
{
const uint64_t delta = dcf_spec.beta & dcf_spec.mask;
unsigned __int128 thresh =
static_cast<unsigned __int128>(
utils::to_integral_type<input_type>{}(xx));
if (cmp_nbits < bitlen)
thresh >>= (bitlen - cmp_nbits);
uint64_t Va = 0;
uint64_t Va1 = 0;
node p0 = k0.root();
node p1 = k1.root();
for (std::size_t L = 0; L < level; ++L)
{
const bool advice0 =
static_cast<bool>(get_lo_bit(p0));
const bool advice1 =
static_cast<bool>(get_lo_bit(p1));
const auto val0 = tree::expand_value(p0);
const auto val1 = tree::expand_value(p1);
const int ai = static_cast<int>(
(thresh >> (cmp_nbits - 1 - L)) & 1);
(void)make_value_cw(val0[0], val0[1], val1[0], val1[1],
static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va, delta,
cmp.mask);
if constexpr (new_key::cmp_is_wildcard)
{
(void)make_value_cw(val0[0], val0[1], val1[0],
val1[1], static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va1, 1ULL,
cmp.mask);
}
const bool pb = path[L];
const bool is_last =
tree::is_last_level(L, new_key::depth);
auto c0 = tree::expand(p0, is_last);
auto c1 = tree::expand(p1, is_last);
p0 = tree::advance(p0, c0, correction_words[L],
correction_advice[L], pb, advice0, is_last);
p1 = tree::advance(p1, c1, correction_words[L],
correction_advice[L], pb, advice1, is_last);
}
const bool advice0 = static_cast<bool>(get_lo_bit(p0));
const bool advice1 = static_cast<bool>(get_lo_bit(p1));
const auto val0 = tree::expand_value(p0);
const auto val1 = tree::expand_value(p1);
const int ai = static_cast<int>(
(thresh >> (cmp_nbits - 1 - level)) & 1);
const uint64_t base = make_value_cw(val0[0], val0[1],
val1[0], val1[1], static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va, delta, cmp.mask);
value_cws[level] =
static_cast<typename new_key::value_cw_word>(base);
if constexpr (new_key::cmp_is_wildcard)
{
const uint64_t v1 = make_value_cw(val0[0], val0[1],
val1[0], val1[1], static_cast<uint8_t>(advice0),
static_cast<uint8_t>(advice1), ai, Va1, 1ULL,
cmp.mask);
value_cw_coeff[level] =
static_cast<typename new_key::value_cw_word>(
(v1 + neg_m(base, cmp.mask)) & cmp.mask);
}
if (level + 1 == cmp_nbits)
{
const bool is_last =
tree::is_last_level(level, new_key::depth);
auto c0 = tree::expand(p0, is_last);
auto c1 = tree::expand(p1, is_last);
p0 = tree::advance(p0, c0, correction_words[level],
correction_advice[level], bit, advice0, is_last);
p1 = tree::advance(p1, c1, correction_words[level],
correction_advice[level], bit, advice1, is_last);
cw_last = static_cast<typename new_key::value_cw_word>(
make_final_cw(p0, p1,
static_cast<uint8_t>(get_lo_bit(p1)), Va,
cmp.mask, true));
}
}
}
}
auto adds_new = detail::incr::extract_addends(placed,
std::make_index_sequence<new_n>{});
if constexpr (old_n > 0)
{
[&]<std::size_t... Is>(std::index_sequence<Is...>) {
((std::get<Is>(adds_new) = std::get<Is>(k0.public_addends)), ...);
}(std::make_index_sequence<old_n>{});
}
input_type off0 = k0.offset_x.raw();
input_type off1 = k1.offset_x.raw();
new_key key0{k0.root(), correction_words, correction_advice,
std::move(wrap0), off0, cmp, value_cws,
static_cast<uint64_t>(cw_last), static_cast<uint64_t>(cmp_add0),
adds_new, value_cw_coeff, static_cast<uint64_t>(cw_last_coeff), tail,
tail_coeff, prefix_cw, prefix_coeff, correction_seeds};
new_key key1{k1.root(), correction_words, correction_advice,
std::move(wrap1), off1, cmp, value_cws,
static_cast<uint64_t>(cw_last), static_cast<uint64_t>(cmp_add1),
adds_new, value_cw_coeff, static_cast<uint64_t>(cw_last_coeff), tail,
tail_coeff, prefix_cw, prefix_coeff, correction_seeds};
if constexpr (old_key::cmp_depth > 0 || (arg_cmp_depth > 0))
{
key0.set_cmp_scalars(cw_last, cmp_add0, cw_last_coeff);
key1.set_cmp_scalars(cw_last, cmp_add1, cw_last_coeff);
if constexpr (old_key::cmp_depth > 0)
{
key0.set_cmp_assigned(k0.cmp_assigned());
key1.set_cmp_assigned(k1.cmp_assigned());
}
}
return dpf::make_party_key_pair(std::move(key0), std::move(key1));
}
} // namespace grow_impl
} // namespace detail
/// @brief Append one interior level and plant specs whose prefix is the new
/// depth. `bit` is the next MSB-side path bit; `x` is the programmed
/// point (must agree with `bit` at the new level) used for leaf lanes.
/// \complexity O(d) rewalk of the existing spine plus one `make_cw` / `advance`
/// and one exterior leaf plant per new packing group. d is the old
/// depth. Copying prior correction words and leaves is O(d + m) for
/// m existing outputs.
/// \rounds none (dealer / joint view)
/// \communication none
template <typename K0, typename K1, typename InputT, typename... Specs>
HEDLEY_WARN_UNUSED_RESULT
auto extend(const K0 & k0, const K1 & k1, bool bit, InputT x,
Specs &&... specs)
{
return detail::grow_impl::grow_impl<true, false>(k0, k1,
static_cast<void *>(nullptr), static_cast<void *>(nullptr), bit, x,
false, static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<psnip_uint8_t *>(nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
std::forward<Specs>(specs)...);
}
/// @brief Append one interior level. Path bit taken from `x` at the old depth.
template <typename K0, typename K1, typename InputT, typename... Specs,
typename = std::enable_if_t<
!std::is_same_v<std::decay_t<InputT>, bool>
&& !is_at_v<std::decay_t<InputT>>
&& !is_cmp_spec_v<std::decay_t<InputT>>
&& !is_verifiable_tag_v<std::decay_t<InputT>>
&& !is_extractable_tag_v<std::decay_t<InputT>>>>
HEDLEY_WARN_UNUSED_RESULT
auto extend(const K0 & k0, const K1 & k1, InputT x, Specs &&... specs)
{
using old_key = detail::grow_impl::bare_key_t<K0>;
const bool bit =
detail::grow_impl::bit_at(static_cast<typename old_key::input_type>(x),
old_key::depth);
return extend(k0, k1, bit, x, std::forward<Specs>(specs)...);
}
/// @brief Plant outputs on levels the tree already has. No new correction word.
/// \complexity O(d) rewalk (or O(1) seed reads with filled memoizers) plus one
/// exterior leaf plant per new packing group. No new interior CW.
/// \rounds none
/// \communication none
template <typename K0, typename K1, typename InputT, typename... Specs>
HEDLEY_WARN_UNUSED_RESULT
auto add_output(const K0 & k0, const K1 & k1, InputT x, Specs &&... specs)
{
return detail::grow_impl::grow_impl<false, false>(k0, k1,
static_cast<void *>(nullptr), static_cast<void *>(nullptr),
/*bit=*/false, x, false,
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<psnip_uint8_t *>(nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
std::forward<Specs>(specs)...);
}
/// @brief Like `extend`, but on-path seeds come from path memoizers (joint /
/// 2+1 view). Memoizers must already be filled through the old depth
/// for `x` (e.g. after `eval_point`); they are not walked here.
/// \complexity O(1) seed reads at the frontier, then the same `make_cw` and
/// leaf work as dealer `extend`. No O(d) rewalk.
/// \rounds none
/// \communication none
template <typename K0, typename K1, typename Memo0, typename Memo1,
typename InputT, typename... Specs>
HEDLEY_WARN_UNUSED_RESULT
auto extend(const K0 & k0, const K1 & k1, Memo0 & m0, Memo1 & m1, bool bit,
InputT x, Specs &&... specs)
{
return detail::grow_impl::grow_impl<true, true>(k0, k1, &m0, &m1, bit, x,
false,
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<psnip_uint8_t *>(nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
std::forward<Specs>(specs)...);
}
template <typename K0, typename K1, typename Memo0, typename Memo1,
typename InputT, typename... Specs,
typename = std::enable_if_t<
!std::is_same_v<std::decay_t<InputT>, bool>
&& !is_at_v<std::decay_t<InputT>>
&& !is_cmp_spec_v<std::decay_t<InputT>>
&& !is_verifiable_tag_v<std::decay_t<InputT>>
&& !is_extractable_tag_v<std::decay_t<InputT>>>>
HEDLEY_WARN_UNUSED_RESULT
auto extend(const K0 & k0, const K1 & k1, Memo0 & m0, Memo1 & m1, InputT x,
Specs &&... specs)
{
using old_key = detail::grow_impl::bare_key_t<K0>;
const bool bit =
detail::grow_impl::bit_at(static_cast<typename old_key::input_type>(x),
old_key::depth);
return extend(k0, k1, m0, m1, bit, x, std::forward<Specs>(specs)...);
}
/// @brief Like `add_output`, reading on-path seeds from path memoizers.
/// \complexity O(1) seed reads per planted level plus leaf plants. Memoizers
/// must already be filled through each new slot's tree level.
/// \rounds none
/// \communication none
template <typename K0, typename K1, typename Memo0, typename Memo1,
typename InputT, typename... Specs>
HEDLEY_WARN_UNUSED_RESULT
auto add_output(const K0 & k0, const K1 & k1, Memo0 & m0, Memo1 & m1, InputT x,
Specs &&... specs)
{
return detail::grow_impl::grow_impl<false, true>(k0, k1, &m0, &m1,
/*bit=*/false, x, false,
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<psnip_uint8_t *>(nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
static_cast<typename detail::grow_impl::bare_key_t<K0>::interior_node *>(
nullptr),
std::forward<Specs>(specs)...);
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_GROW_HPP__