libdpf/include/dpf/grow.hpp

1144 lines
47 KiB
C++
Raw Normal View History

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