libdpf/include/grotto/prefix_parity.hpp
Ryan Henry e4e666f459 Initial import of libdpf.
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-24 14:08:32 -06:00

494 lines
20 KiB
C++
Raw 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 grotto/prefix_parity.hpp
/// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @brief
/// @details
/// @copyright Copyright (c) 2019-2023 Ryan Henry and others
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref GPLv2) for details.
#ifndef LIBDPF_INCLUDE_GROTTO_PREFIX_PARITY_HPP__
#define LIBDPF_INCLUDE_GROTTO_PREFIX_PARITY_HPP__
#include "hedley/hedley.h"
#include "dpf/dpf_key.hpp"
#include "dpf/bit.hpp"
#include "dpf/bitstring.hpp"
#include "dpf/twiddle.hpp"
#include "dpf/leaf_node.hpp"
#include "dpf/dcf.hpp"
#include "grotto/offset_iterable.hpp"
#include "dpf/path_memoizer.hpp"
#include "dpf/utils.hpp"
#include <stdexcept>
#include <type_traits>
namespace grotto
{
template <typename NodeT,
typename InputT>
auto parity_of_substring_prefix(const NodeT & node, InputT x) noexcept
{
static constexpr auto bits_per_limb = dpf::utils::bitlength_of_v<decltype(node[0])>;
std::size_t off = dpf::offset_within_block<dpf::bit, NodeT>(x);
auto div = std::lldiv(off, bits_per_limb);
auto parity = node[div.quot] & ((1ul << div.rem) - 1ul);
for (std::size_t i = 0; i < div.quot; ++i) parity ^= node[i];
return psnip_builtin_parity64(parity);
}
template <bool use_early_terminate_optimization = true,
typename InputT,
typename DpfKey,
std::size_t NumParts,
std::enable_if_t<std::is_same_v<InputT, typename DpfKey::input_type>, bool> = false>
static auto prefix_parities(const DpfKey & dpf, const std::array<InputT, NumParts> & endpoints)
{
static constexpr std::size_t num_parts = NumParts;
static constexpr std::size_t depth = DpfKey::depth;
using input_type = typename DpfKey::input_type;
static constexpr std::size_t input_bits = dpf::utils::bitlength_of_v<input_type>;
using interior_node = typename DpfKey::interior_node;
using exterior_node = typename DpfKey::exterior_node;
exterior_node leaf;
auto path = make_basic_path_memoizer(dpf);
std::array<uint_fast8_t, depth+1> direction = { 0 }; // always "traverse left" to get to the root
std::array<uint_fast8_t, depth+1> parity = { 0 };
std::array<bool, num_parts> prefix_parities;
// A prefix that ends on a node boundary (trailing zeros past the leaf
// width) is the parity accumulated at that node. Stopping any earlier
// would cut through a leaf, so those endpoints still open the exterior
// node. The path memoizer's high-water mark keeps a later endpoint from
// resuming past nodes this walk never filled.
std::size_t new_first = for_each_offset(std::begin(endpoints), std::end(endpoints), dpf.offset_x(0),
[&](std::size_t which_part, input_type current_endpoint)
{
std::size_t to_level = depth;
if constexpr (use_early_terminate_optimization)
{
// Stop where the remaining suffix is all zeros, but never
// above a partial leaf. `parity_of_substring_prefix` indexes
// bits inside the exterior node, which is wider than
// `lg_outputs_per_leaf` when several outputs share that node.
constexpr std::size_t leaf_bit_lg
= dpf::lg_outputs_per_leaf_v<dpf::bit, exterior_node>;
constexpr std::size_t align = std::max(
DpfKey::lg_outputs_per_leaf, leaf_bit_lg);
const std::size_t tz = dpf::utils::countr_zero<input_type>{}(current_endpoint);
if (tz >= input_bits)
to_level = 0;
else if (tz > align)
to_level = input_bits - tz;
}
// assign_x sets path[0] = dpf.root on the first point and returns
// the first level that differs from the previous point. direction[0]
// and parity[0] stay 0, which is the empty-prefix parity.
const std::size_t next_level = dpf::detail::path_resume_for_level(
path, dpf, current_endpoint, to_level);
std::size_t level_index = next_level - 1;
if (level_index < to_level)
{
DPF_UNROLL_LOOP
for (auto mask = dpf.msb_mask >> level_index;
level_index < to_level; ++level_index, mask>>=1)
{
bool bit = !!(mask & current_endpoint);
direction[level_index+1] = bit;
path[level_index+1] = DpfKey::traverse_interior(path[level_index], dpf.correction_word(level_index, direction[level_index+1]), direction[level_index+1]);
parity[level_index+1] = parity[level_index] ^ ((direction[level_index] ^ direction[level_index+1]) & dpf::get_lo_bit(path[level_index]));
}
}
dpf::detail::path_note_filled_to(path, to_level);
if (HEDLEY_UNLIKELY(to_level != depth))
{
// Same quantity as a full walk that then steps only left into
// an empty leaf substring: parity[L] ^ (direction[L] & t[L]).
prefix_parities[which_part] = parity[to_level]
^ (direction[to_level] & dpf::get_lo_bit(path[to_level]));
}
else
{
if (next_level <= depth)
{
leaf = dpf.template traverse_exterior<0>(path[depth]);
}
prefix_parities[which_part] = parity[depth]
^ (direction[depth] & dpf::get_lo_bit(path[depth])
^ parity_of_substring_prefix(leaf, current_endpoint));
}
});
return std::make_tuple(prefix_parities, new_first);
}
template <typename DpfKey,
std::size_t NumParts>
static auto all_segment_parities_from_prefix_parities(const DpfKey & dpf,
const std::array<bool, NumParts> & prefix_parities, std::size_t new_first)
{
static constexpr std::size_t num_parts = NumParts;
std::array<bool, num_parts> segment_parities;
// currently setup so that for endpoints = {A, B, C}, segment_parities[0] = [A, B) while segment_parities[2] = [C, A)
// to change to having the wrapping segment first, xor with previous prefix_parity and not subsequent one
// also need to change final xor to simply using new_first
// segment_parities[0] = prefix_parities[0] ^ prefix_parities[num_parts-1];
for (std::size_t i = 0; i < num_parts; ++i)
{
segment_parities[i] = prefix_parities[i] ^ prefix_parities[(i+1)%num_parts];
}
segment_parities[(new_first+num_parts-1)%num_parts] ^= dpf::get_lo_bit(dpf.root());
return segment_parities;
}
template <typename DpfKey,
std::size_t NumSegments,
std::size_t NumParts>
static auto specific_segment_parities_from_prefix_parities(const DpfKey & dpf,
const std::array<std::size_t, NumSegments> & segment_indices,
const std::array<bool, NumParts> & prefix_parities, std::size_t new_first)
{
assert(std::is_sorted(std::begin(segment_indices), std::end(segment_indices)));
static constexpr std::size_t num_segments = NumSegments;
static constexpr std::size_t num_parts = NumParts;
std::array<bool, num_segments> segment_parities;
for (std::size_t i = 0; i < num_segments; ++i)
{
segment_parities[i] = prefix_parities[segment_indices[i]] ^ prefix_parities[segment_indices[(i+1)%num_segments]];
}
auto begin = std::begin(segment_indices), end = std::end(segment_indices);
segment_parities[(std::distance(begin, std::lower_bound(begin, end, new_first))+num_segments-1)%num_segments] ^= dpf::get_lo_bit(dpf.root());
return segment_parities;
}
template <bool use_early_terminate_optimization = true,
typename InputT,
typename DpfKey,
std::size_t NumParts,
std::enable_if_t<std::is_same_v<InputT, typename DpfKey::input_type>, bool> = false>
static auto segment_parities(const DpfKey & dpf, const std::array<InputT, NumParts> & endpoints)
{
if constexpr (NumParts == 0)
{
(void)dpf;
return std::array<bool, 0>{};
}
else if constexpr (NumParts == 1)
{
(void)endpoints;
// One segment wraps the whole domain. Its parity share is the root
// control bit (the general n-way formula reduces to the same value).
return std::array<bool, 1>{
static_cast<bool>(dpf::get_lo_bit(dpf.root()))};
}
else
{
auto [prefix_parities, new_first] = grotto::prefix_parities<use_early_terminate_optimization>(dpf, endpoints);
return all_segment_parities_from_prefix_parities(dpf, prefix_parities, new_first);
}
}
namespace detail
{
template <typename T, typename = void>
struct key_has_cmp : std::false_type {};
template <typename T>
struct key_has_cmp<T, std::void_t<decltype(std::declval<const T &>().has_cmp())>>
: std::true_type {};
template <typename KeyT>
uint64_t cmp_addend_raw(const KeyT & key) noexcept
{
if constexpr (dpf::is_party_key_v<KeyT>)
return key.cmp_addend().raw();
else
return key.cmp_addend();
}
} // namespace detail
/// Additive prefix indicators from the key's DCF value sums.
///
/// Same multi-knot walk as `prefix_parities` (shared path, MSB-first), but each
/// level contributes the Boyle/Guo value correction instead of an advice-bit
/// XOR. Party 0 + party 1 equals `eval` of the comparison at that endpoint, in
/// the comparison group: for `dpf::gt(1)` that is `1` when `α < endpoint` and
/// `0` otherwise. A published per-level sign is not used; that bit is the keep
/// child's control and would leak the path (elementary2
/// `notes/signed-prefix-parity.md` §5).
///
/// Endpoints are visited in the order given. The running sum is reused from the
/// longest common prefix. The zero-suffix early stop of the XOR walk is not
/// applied: later levels still add a Convert word.
template <typename InputT,
typename DpfKey,
std::size_t NumParts,
std::enable_if_t<std::is_same_v<InputT, typename DpfKey::input_type>, bool> = false>
static auto signed_prefix_parities(const DpfKey & dpf,
const std::array<InputT, NumParts> & endpoints)
{
if constexpr (!detail::key_has_cmp<DpfKey>::value)
{
(void)dpf;
(void)endpoints;
throw std::invalid_argument("signed_prefix_parities: key has no comparison channel");
}
else if (!dpf.has_cmp())
{
throw std::invalid_argument("signed_prefix_parities: key has no comparison channel");
}
else if (!dpf.cmp_assigned())
{
throw std::invalid_argument(
"signed_prefix_parities: wildcard comparison payload not assigned");
}
else
{
using namespace dpf::detail::dcf_impl;
using key_type = dpf::unwrap_party_key_t<DpfKey>;
constexpr std::size_t depth = key_type::depth;
const auto & ch = dpf.cmp();
const std::size_t nbits = static_cast<std::size_t>(ch.nbits);
if (nbits > depth)
throw std::invalid_argument("signed_prefix_parities: comparison is deeper than the key");
const uint64_t mask = ch.mask;
const int party = dpf::get_lo_bit(dpf.root()) ? 1 : 0;
const uint64_t addend = detail::cmp_addend_raw(dpf);
auto path = dpf::make_basic_path_memoizer(dpf);
// `sum_at[i]` is the path sum after the first `i` levels, before `cw_last`.
std::array<uint64_t, depth + 1> sum_at{};
std::array<uint64_t, NumParts> prefixes{};
for (std::size_t which = 0; which < NumParts; ++which)
{
auto tx = dpf.offset_x(endpoints[which]);
dpf::utils::flip_msb_if_signed_integral(tx);
if (ch.trivial == dpf::cmp_trivial::always_true
|| ch.trivial == dpf::cmp_trivial::always_false)
{
prefixes[which] = addend & mask;
continue;
}
const std::size_t resume = dpf::detail::path_resume_for_level(
path, dpf, tx, nbits);
for (std::size_t level = resume, shift = level == 0 ? 0 : level - 1;
level <= nbits; ++level, ++shift)
{
const auto bit_mask = key_type::msb_mask >> shift;
const bool bit = !!(bit_mask & tx);
path[level] = DpfKey::traverse_interior(path[level - 1],
dpf.correction_word(level - 1, bit), bit);
}
dpf::detail::path_note_filled_to(path, nbits);
const std::size_t start = resume > nbits ? nbits : resume - 1;
uint64_t V = sum_at[start];
if (resume <= nbits)
{
auto bit_mask = key_type::msb_mask >> start;
for (std::size_t i = start; i < nbits; ++i, bit_mask >>= 1)
{
const bool xi = !!(bit_mask & tx);
const auto & parent = path[i];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(parent));
auto kids = key_type::interior_prg::eval01(dpf::unset_lo_2bits(parent));
const uint64_t v = convert_node(kids[xi ? 1 : 0], mask);
const uint64_t contrib = (v + (t ? dpf.value_cw(i) : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
sum_at[i + 1] = V;
}
}
const auto & leaf = path[nbits];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(leaf));
const uint64_t c = convert_node(leaf, mask);
const uint64_t contrib = (c + (t ? dpf.cw_last() : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
if (ch.eval_as_ge)
V = neg_m(V, mask);
prefixes[which] = (V + addend) & mask;
}
return prefixes;
}
}
/// Runtime-length form of `signed_prefix_parities`. `out[i]` receives the same
/// share a one-element call would return for `endpoints[i]`.
template <typename InputT,
typename DpfKey>
static void signed_prefix_parities_into(const DpfKey & dpf,
const InputT * endpoints, std::size_t n, uint64_t * out)
{
if constexpr (!detail::key_has_cmp<DpfKey>::value)
{
(void)dpf;
(void)endpoints;
(void)n;
(void)out;
throw std::invalid_argument("signed_prefix_parities: key has no comparison channel");
}
else if (!dpf.has_cmp())
{
throw std::invalid_argument("signed_prefix_parities: key has no comparison channel");
}
else if (!dpf.cmp_assigned())
{
throw std::invalid_argument(
"signed_prefix_parities: wildcard comparison payload not assigned");
}
else if (n != 0 && (endpoints == nullptr || out == nullptr))
{
throw std::invalid_argument("signed_prefix_parities: null endpoint buffer");
}
else
{
using namespace dpf::detail::dcf_impl;
using key_type = dpf::unwrap_party_key_t<DpfKey>;
constexpr std::size_t depth = key_type::depth;
const auto & ch = dpf.cmp();
const std::size_t nbits = static_cast<std::size_t>(ch.nbits);
if (nbits > depth)
throw std::invalid_argument("signed_prefix_parities: comparison is deeper than the key");
const uint64_t mask = ch.mask;
const int party = dpf::get_lo_bit(dpf.root()) ? 1 : 0;
const uint64_t addend = detail::cmp_addend_raw(dpf);
auto path = dpf::make_basic_path_memoizer(dpf);
std::array<uint64_t, depth + 1> sum_at{};
for (std::size_t which = 0; which < n; ++which)
{
auto tx = dpf.offset_x(endpoints[which]);
dpf::utils::flip_msb_if_signed_integral(tx);
if (ch.trivial == dpf::cmp_trivial::always_true
|| ch.trivial == dpf::cmp_trivial::always_false)
{
out[which] = addend & mask;
continue;
}
const std::size_t resume = dpf::detail::path_resume_for_level(
path, dpf, tx, nbits);
for (std::size_t level = resume, shift = level == 0 ? 0 : level - 1;
level <= nbits; ++level, ++shift)
{
const auto bit_mask = key_type::msb_mask >> shift;
const bool bit = !!(bit_mask & tx);
path[level] = DpfKey::traverse_interior(path[level - 1],
dpf.correction_word(level - 1, bit), bit);
}
dpf::detail::path_note_filled_to(path, nbits);
const std::size_t start = resume > nbits ? nbits : resume - 1;
uint64_t V = sum_at[start];
if (resume <= nbits)
{
auto bit_mask = key_type::msb_mask >> start;
for (std::size_t i = start; i < nbits; ++i, bit_mask >>= 1)
{
const bool xi = !!(bit_mask & tx);
const auto & parent = path[i];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(parent));
auto kids = key_type::interior_prg::eval01(dpf::unset_lo_2bits(parent));
const uint64_t v = convert_node(kids[xi ? 1 : 0], mask);
const uint64_t contrib = (v + (t ? dpf.value_cw(i) : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
sum_at[i + 1] = V;
}
}
const auto & leaf = path[nbits];
const uint8_t t = static_cast<uint8_t>(dpf::get_lo_bit(leaf));
const uint64_t c = convert_node(leaf, mask);
const uint64_t contrib = (c + (t ? dpf.cw_last() : 0ULL)) & mask;
V = (V + (party ? neg_m(contrib, mask) : contrib)) & mask;
if (ch.eval_as_ge)
V = neg_m(V, mask);
out[which] = (V + addend) & mask;
}
}
}
/// One-hot segment shares for a unit `gt` comparison (`if_true = 1`, `if_false = 0`).
///
/// `endpoints` is sorted ascending. Piece `i < n-1` is `[endpoints[i], endpoints[i+1])`
/// and the last piece wraps. Party 0 + party 1 is `1` on the piece that contains
/// the key's target and `0` on the others, so a public dot product with LUT
/// constants is an additive share of the function value (not of its negation).
template <typename InputT,
typename DpfKey,
std::size_t NumParts,
std::enable_if_t<std::is_same_v<InputT, typename DpfKey::input_type>, bool> = false>
static auto signed_segment_parities(const DpfKey & dpf,
const std::array<InputT, NumParts> & endpoints)
{
using namespace dpf::detail::dcf_impl;
if constexpr (!detail::key_has_cmp<DpfKey>::value)
{
(void)dpf;
(void)endpoints;
throw std::invalid_argument("signed_segment_parities: key has no comparison channel");
}
else if (!dpf.has_cmp())
{
throw std::invalid_argument("signed_segment_parities: key has no comparison channel");
}
else
{
const auto & ch = dpf.cmp();
if (ch.kind != dpf::cmp_kind::gt)
throw std::invalid_argument(
"signed_segment_parities: comparison must be gt");
const uint64_t mask = ch.mask;
// Public +1 sits with party 0 (root control bit clear), matching a dealer
// split of the constant 1 with a zero blind on party 1.
const uint64_t one = dpf::get_lo_bit(dpf.root()) ? 0ULL : (1ULL & mask);
if constexpr (NumParts == 0)
{
(void)dpf;
(void)endpoints;
return std::array<uint64_t, 0>{};
}
else if constexpr (NumParts == 1)
{
(void)endpoints;
return std::array<uint64_t, 1>{one};
}
else
{
const auto prefix = signed_prefix_parities(dpf, endpoints);
std::array<uint64_t, NumParts> segments{};
for (std::size_t i = 0; i < NumParts; ++i)
{
const uint64_t nxt = prefix[(i + 1) % NumParts];
segments[i] = (nxt + neg_m(prefix[i], mask)) & mask;
}
// `segments[n-1]` is `p[0] - p[n-1]`. The wrap piece is that difference
// plus the public 1, which makes the pieces sum to 1.
segments[NumParts - 1] = (segments[NumParts - 1] + one) & mask;
return segments;
}
}
}
} // namespace grotto
#endif // LIBDPF_INCLUDE_GROTTO_PREFIX_PARITY_HPP__