/// @file grotto/prefix_parity.hpp /// @author Ryan Henry /// @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 #include namespace grotto { template auto parity_of_substring_prefix(const NodeT & node, InputT x) noexcept { static constexpr auto bits_per_limb = dpf::utils::bitlength_of_v; std::size_t off = dpf::offset_within_block(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> = false> static auto prefix_parities(const DpfKey & dpf, const std::array & 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; 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 direction = { 0 }; // always "traverse left" to get to the root std::array parity = { 0 }; std::array 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; constexpr std::size_t align = std::max( DpfKey::lg_outputs_per_leaf, leaf_bit_lg); const std::size_t tz = dpf::utils::countr_zero{}(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 static auto all_segment_parities_from_prefix_parities(const DpfKey & dpf, const std::array & prefix_parities, std::size_t new_first) { static constexpr std::size_t num_parts = NumParts; std::array 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 static auto specific_segment_parities_from_prefix_parities(const DpfKey & dpf, const std::array & segment_indices, const std::array & 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 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> = false> static auto segment_parities(const DpfKey & dpf, const std::array & endpoints) { if constexpr (NumParts == 0) { (void)dpf; return std::array{}; } 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{ static_cast(dpf::get_lo_bit(dpf.root()))}; } else { auto [prefix_parities, new_first] = grotto::prefix_parities(dpf, endpoints); return all_segment_parities_from_prefix_parities(dpf, prefix_parities, new_first); } } namespace detail { template struct key_has_cmp : std::false_type {}; template struct key_has_cmp().has_cmp())>> : std::true_type {}; template uint64_t cmp_addend_raw(const KeyT & key) noexcept { if constexpr (dpf::is_party_key_v) 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 , bool> = false> static auto signed_prefix_parities(const DpfKey & dpf, const std::array & endpoints) { if constexpr (!detail::key_has_cmp::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; constexpr std::size_t depth = key_type::depth; const auto & ch = dpf.cmp(); const std::size_t nbits = static_cast(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 sum_at{}; std::array 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(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(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 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::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; constexpr std::size_t depth = key_type::depth; const auto & ch = dpf.cmp(); const std::size_t nbits = static_cast(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 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(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(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 , bool> = false> static auto signed_segment_parities(const DpfKey & dpf, const std::array & endpoints) { using namespace dpf::detail::dcf_impl; if constexpr (!detail::key_has_cmp::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{}; } else if constexpr (NumParts == 1) { (void)endpoints; return std::array{one}; } else { const auto prefix = signed_prefix_parities(dpf, endpoints); std::array 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__