/// @file dpf/geneval.hpp /// @brief Fused generation and evaluation (Doerner–Shelat on the eval trie). /// @details `make_dpf` / `make_dpf_doerner_shelat` build a reusable key, then /// `eval_*` walks it. `geneval_*` does both at once: one correction /// word per level, opened from the XOR-reduction of the nodes the /// public query actually expands. While the secret path's parent is /// still in that trie the word matches the reusable key byte for /// byte (same roots, same Beaver tape). After the path leaves, the /// word is uniform and later outputs still reconstruct — off-path /// nodes are identical across the two parties, so a dummy word /// cancels. /// /// A wildcard-input call takes additive shares of the real point and /// a public query. It samples a random target, runs geneval there, /// and shifts the query by `target - x`, which is what /// `offset_x` does after a wildcard key is bound to `x`. /// /// `geneval_cmp` is the comparison-channel form. The value-correction /// word is a function of the secret path at every level, so the walk /// stays live for the whole depth and the opened words match a /// Doerner–Shelat comparison key. Prefix shares are /// `eval_point(cmp, ...)` at each endpoint. Piecewise-cubic evaluation /// on top of that is `grotto::geneval_offset_horner`. /// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors) /// @license Released under a GNU General Public v2.0 (GPLv2) license; /// see [LICENSE.md](@ref license) for details. #ifndef LIBDPF_INCLUDE_DPF_GENEVAL_HPP__ #define LIBDPF_INCLUDE_DPF_GENEVAL_HPP__ #include #include #include #include #include #include #include #include #include #include #include "hedley/hedley.h" #include "simde/simde/x86/avx2.h" #include "dpf/aligned_allocator.hpp" #include "dpf/doerner_shelat.hpp" #include "dpf/eval_target.hpp" #include "dpf/leaf_node.hpp" namespace dpf { /// Tag for a geneval whose point is known only as additive shares. struct wildcard_input_t { }; inline constexpr wildcard_input_t wildcard_input{}; /// Shares and the correction words opened along the query trie. /// `correction_words[i]` / `correction_advice[i]` match a reusable key at /// the same target for every `i < live_levels`. `leaf_live` means the /// target's leaf was in the trie, so `leaf` is that key's leaf word. template struct geneval_result { std::vector party0; std::vector party1; std::vector> correction_words; std::vector correction_advice; std::size_t live_levels = 0; bool leaf_live = false; Leaf leaf{}; }; namespace detail { template HEDLEY_ALWAYS_INLINE T geneval_mod_add(T a, T b) noexcept { using U = std::make_unsigned_t; U sum = static_cast(static_cast(a) + static_cast(b)); T out; std::memcpy(&out, &sum, sizeof(out)); return out; } template HEDLEY_ALWAYS_INLINE T geneval_mod_sub(T a, T b) noexcept { using U = std::make_unsigned_t; U diff = static_cast(static_cast(a) - static_cast(b)); T out; std::memcpy(&out, &diff, sizeof(out)); return out; } template T geneval_flipped(T x) { utils::flip_msb_if_signed_integral(x); return x; } /// Leaf-node id of an already MSB-flipped input. The id is the high /// `depth` bits; the low `lg(outputs_per_leaf)` bits select the lane. template uint64_t geneval_leaf_id(typename Dpf::input_type x) { return static_cast(utils::get_from_node(x)); } inline uint64_t geneval_prefix(uint64_t leaf, std::size_t depth, std::size_t bits) { if (bits == 0) return 0; if (bits >= depth) return leaf; return leaf >> (depth - bits); } inline bool geneval_any_prefix(const std::vector & leaves, std::size_t depth, uint64_t id, std::size_t bits) { if (leaves.empty()) return false; if (bits == 0) return true; const std::size_t sh = depth - bits; const uint64_t lo = (sh >= 64) ? 0 : (id << sh); auto it = std::lower_bound(leaves.begin(), leaves.end(), lo); if (it == leaves.end()) return false; return geneval_prefix(*it, depth, bits) == id; } template geneval_result geneval_empty_result() { geneval_result out; std::memset(&out.leaf, 0, sizeof(out.leaf)); return out; } template auto geneval_run(InputT x0, InputT x1, const std::vector & queries, RootSampler & root_sampler, PadRng & pads, OutputT y) { static_assert(std::is_integral_v, "geneval input shares are an integral domain"); static_assert(!dpf::is_wildcard_v, "geneval output is concrete; assign a wildcard leaf on a key"); static_assert(utils::bitlength_of_v <= 64, "geneval leaf ids are 64-bit"); using dpf_type = utils::dpf_type_t; using node = typename dpf_type::interior_node; using leaf_node = leaf_node_t; constexpr std::size_t depth = dpf_type::depth; if (queries.empty()) return geneval_empty_result(); if (queries.size() > (std::size_t{1} << 22)) throw std::length_error("geneval query is too large"); InputT x0c = x0; InputT x1c = x1; utils::flip_msb_if_signed_integral(x0c); const InputT alpha = utils::xor_input_shares(x0c, x1c); std::vector flipped; flipped.reserve(queries.size()); std::vector leaves; leaves.reserve(queries.size()); for (const InputT & q : queries) { InputT fq = geneval_flipped(q); flipped.push_back(fq); leaves.push_back(geneval_leaf_id(fq)); } std::vector unique_leaves = leaves; std::sort(unique_leaves.begin(), unique_leaves.end()); unique_leaves.erase(std::unique(unique_leaves.begin(), unique_leaves.end()), unique_leaves.end()); if (unique_leaves.size() > (std::size_t{1} << 20)) throw std::length_error("geneval trie is too large"); const uint64_t secret_leaf = geneval_leaf_id(alpha); local_cw_protocol proto{pads}; constexpr auto to_int = utils::to_integral_type{}; const node root0 = dpf::unset_lo_bit(static_cast(root_sampler())); const node root1 = dpf::set_lo_bit(static_cast(root_sampler())); struct slot { uint64_t id; node s0; node s1; }; std::vector frontier; frontier.push_back(slot{0, root0, root1}); geneval_result result; std::memset(&result.leaf, 0, sizeof(result.leaf)); result.correction_words.reserve(depth); result.correction_advice.reserve(depth); auto mask = dpf_type::msb_mask; bool still_live = true; for (std::size_t level = 0; level < depth; ++level, mask >>= 1) { const uint8_t bit0 = static_cast(!!(to_int(mask) & to_int(x0c))); const uint8_t bit1 = static_cast(!!(to_int(mask) & to_int(x1c))); const uint64_t parent_id = geneval_prefix(secret_leaf, depth, level); node L0 = simde_mm_setzero_si128(); node R0 = simde_mm_setzero_si128(); node L1 = simde_mm_setzero_si128(); node R1 = simde_mm_setzero_si128(); bool level_live = false; struct exp { uint64_t id; node s0, s1, L0, R0, L1, R1; }; std::vector exps; exps.reserve(frontier.size()); for (const slot & n : frontier) { if (n.id == parent_id) level_live = true; const auto c0 = InteriorPRG::eval01(dpf::unset_lo_2bits(n.s0)); const auto c1 = InteriorPRG::eval01(dpf::unset_lo_2bits(n.s1)); L0 = ds_xor(L0, c0[0]); R0 = ds_xor(R0, c0[1]); L1 = ds_xor(L1, c1[0]); R1 = ds_xor(R1, c1[1]); exps.push_back(exp{n.id, n.s0, n.s1, c0[0], c0[1], c1[0], c1[1]}); } node cw; uint8_t advice; if (still_live && level_live) { auto blinds = proto.prepare_level(L0, R0, bit0, L1, R1, bit1); auto opened = proto.open_cw(blinds); cw = opened.first; advice = opened.second; ++result.live_levels; } else { still_live = false; cw = pads.block(); const uint8_t t0 = static_cast(pads.bit() & 1u); const uint8_t t1 = static_cast(pads.bit() & 1u); advice = static_cast((t1 << 1) | t0); } result.correction_words.push_back(cw); result.correction_advice.push_back(advice); const node cw0 = dpf::set_lo_bit(cw, advice & 1u); const node cw1 = dpf::set_lo_bit(cw, (advice >> 1) & 1u); const std::size_t child_bits = level + 1; std::vector next; next.reserve(exps.size() * 2); for (const exp & e : exps) { const uint64_t left = e.id << 1; const uint64_t right = left | 1ull; if (geneval_any_prefix(unique_leaves, depth, left, child_bits)) { next.push_back(slot{left, dpf::xor_if_lo_bit(e.L0, cw0, e.s0), dpf::xor_if_lo_bit(e.L1, cw0, e.s1)}); } if (geneval_any_prefix(unique_leaves, depth, right, child_bits)) { next.push_back(slot{right, dpf::xor_if_lo_bit(e.R0, cw1, e.s0), dpf::xor_if_lo_bit(e.R1, cw1, e.s1)}); } } frontier = std::move(next); } result.leaf_live = geneval_any_prefix(unique_leaves, depth, secret_leaf, depth); if (result.leaf_live) { const slot * on = nullptr; for (const slot & n : frontier) { if (n.id == secret_leaf) { on = &n; break; } } if (on == nullptr) throw std::logic_error("geneval: secret leaf missing from trie"); const bool sign0 = dpf::get_lo_bit(on->s0); auto built = dpf::make_leaves(alpha, dpf::unset_lo_2bits(on->s0), dpf::unset_lo_2bits(on->s1), sign0, std::size_t{0}, y); result.leaf = std::get<0>(built.first.first); } result.party0.reserve(flipped.size()); result.party1.reserve(flipped.size()); for (std::size_t i = 0; i < flipped.size(); ++i) { const uint64_t id = leaves[i]; const slot * n = nullptr; for (const slot & s : frontier) { if (s.id == id) { n = &s; break; } } if (n == nullptr) throw std::logic_error("geneval: query leaf missing from trie"); auto share0 = dpf_type::template traverse_exterior<0>(n->s0, result.leaf); auto share1 = dpf_type::template traverse_exterior<0>(n->s1, result.leaf); const auto lane = static_cast(to_int(flipped[i])); result.party0.push_back(extract_leaf(share0, lane)); result.party1.push_back(extract_leaf(share1, lane)); } return result; } template InputT geneval_from_bits(uint64_t bits) { using U = std::make_unsigned_t; U u = static_cast(bits); InputT out; std::memcpy(&out, &u, sizeof(out)); return out; } template bool geneval_out_of_order(InputT from, InputT to) { // Numeric order. An unsigned compare of a signed value treats a negative // `from` as larger than a positive `to`, and would reject `[-1, 1]`. if constexpr (std::is_signed_v) return from > to; else return utils::to_integral_type{}(from) > utils::to_integral_type{}(to); } template std::vector geneval_full_domain() { constexpr std::size_t bitlen = utils::bitlength_of_v; if (bitlen > 20) throw std::length_error("geneval_full domain is too large"); const uint64_t n = uint64_t{1} << bitlen; std::vector qs(static_cast(n)); // Index `i` is the input's bit pattern, including the sign bit. A // narrowing cast of `i` to a signed type is implementation-defined. for (uint64_t i = 0; i < n; ++i) qs[static_cast(i)] = geneval_from_bits(i); return qs; } template std::vector geneval_inclusive(InputT from, InputT to) { if (geneval_out_of_order(from, to)) { throw std::invalid_argument("geneval_interval: from > to"); } std::vector qs; InputT q = from; const InputT one = utils::make_from_integral_value{}(1); for (;;) { qs.push_back(q); if (q == to) break; q = geneval_mod_add(q, one); if (qs.size() > (std::size_t{1} << 22)) throw std::length_error("geneval_interval is too large"); } return qs; } template InputT geneval_sample_target(TargetSampler & sample) { return static_cast(sample()); } template std::vector geneval_shift_all(const std::vector & qs, InputT delta) { std::vector out; out.reserve(qs.size()); for (const InputT & q : qs) out.push_back(geneval_mod_add(q, delta)); return out; } } // namespace detail /// Geneval at one public point. The secret point is `x0 XOR x1`. template HEDLEY_WARN_UNUSED_RESULT auto geneval_point(InputT x0, InputT x1, InputT query, ds_randomness rng, OutputT y) { return detail::geneval_run(x0, x1, std::vector{query}, rng.root, rng.pad, y); } /// Geneval on the inclusive interval `[from, to]`. template HEDLEY_WARN_UNUSED_RESULT auto geneval_interval(InputT x0, InputT x1, InputT from, InputT to, ds_randomness rng, OutputT y) { return detail::geneval_run(x0, x1, detail::geneval_inclusive(from, to), rng.root, rng.pad, y); } /// Geneval on the whole domain. Refuses a domain above 2^20 inputs. template HEDLEY_WARN_UNUSED_RESULT auto geneval_full(InputT x0, InputT x1, ds_randomness rng, OutputT y) { return detail::geneval_run(x0, x1, detail::geneval_full_domain(), rng.root, rng.pad, y); } /// Geneval on a public sequence, in the order given. template HEDLEY_WARN_UNUSED_RESULT auto geneval_sequence(InputT x0, InputT x1, ForwardIterator begin, ForwardIterator end, ds_randomness rng, OutputT y) { std::vector qs(begin, end); return detail::geneval_run(x0, x1, std::move(qs), rng.root, rng.pad, y); } /// Wildcard-input geneval. `x0 + x1` is the real point (additive shares). /// `sample_target()` is the random DPF target; the public query is shifted /// by `target - (x0 + x1)` before the walk. template HEDLEY_WARN_UNUSED_RESULT auto geneval_point(wildcard_input_t, InputT x0, InputT x1, InputT query, ds_randomness rng, TargetSampler sample_target, OutputT y) { const InputT alpha = detail::geneval_sample_target(sample_target); const InputT delta = detail::geneval_mod_sub(alpha, detail::geneval_mod_add(x0, x1)); const InputT shifted = detail::geneval_mod_add(query, delta); InputT zero{}; return detail::geneval_run(zero, alpha, std::vector{shifted}, rng.root, rng.pad, y); } template HEDLEY_WARN_UNUSED_RESULT auto geneval_point(wildcard_input_t, InputT x0, InputT x1, InputT query, ds_randomness rng, OutputT y) { return geneval_point(wildcard_input, x0, x1, query, std::move(rng), [] { return dpf::uniform_sample(); }, y); } template HEDLEY_WARN_UNUSED_RESULT auto geneval_interval(wildcard_input_t, InputT x0, InputT x1, InputT from, InputT to, ds_randomness rng, TargetSampler sample_target, OutputT y) { const InputT alpha = detail::geneval_sample_target(sample_target); const InputT delta = detail::geneval_mod_sub(alpha, detail::geneval_mod_add(x0, x1)); auto shifted = detail::geneval_shift_all( detail::geneval_inclusive(from, to), delta); InputT zero{}; return detail::geneval_run(zero, alpha, std::move(shifted), rng.root, rng.pad, y); } template HEDLEY_WARN_UNUSED_RESULT auto geneval_interval(wildcard_input_t, InputT x0, InputT x1, InputT from, InputT to, ds_randomness rng, OutputT y) { return geneval_interval(wildcard_input, x0, x1, from, to, std::move(rng), [] { return dpf::uniform_sample(); }, y); } template HEDLEY_WARN_UNUSED_RESULT auto geneval_full(wildcard_input_t, InputT x0, InputT x1, ds_randomness rng, TargetSampler sample_target, OutputT y) { const InputT alpha = detail::geneval_sample_target(sample_target); const InputT delta = detail::geneval_mod_sub(alpha, detail::geneval_mod_add(x0, x1)); InputT zero{}; auto full = detail::geneval_run(zero, alpha, detail::geneval_full_domain(), rng.root, rng.pad, y); constexpr auto to_int = utils::to_integral_type{}; const std::size_t n = full.party0.size(); std::vector p0(n), p1(n); for (std::size_t i = 0; i < n; ++i) { InputT q = detail::geneval_from_bits(i); InputT s = detail::geneval_mod_add(q, delta); const std::size_t si = static_cast(to_int(s)); p0[i] = full.party0[si]; p1[i] = full.party1[si]; } full.party0 = std::move(p0); full.party1 = std::move(p1); return full; } template HEDLEY_WARN_UNUSED_RESULT auto geneval_full(wildcard_input_t, InputT x0, InputT x1, ds_randomness rng, OutputT y) { return geneval_full(wildcard_input, x0, x1, std::move(rng), [] { return dpf::uniform_sample(); }, y); } template HEDLEY_WARN_UNUSED_RESULT auto geneval_sequence(wildcard_input_t, InputT x0, InputT x1, ForwardIterator begin, ForwardIterator end, ds_randomness rng, TargetSampler sample_target, OutputT y) { const InputT alpha = detail::geneval_sample_target(sample_target); const InputT delta = detail::geneval_mod_sub(alpha, detail::geneval_mod_add(x0, x1)); std::vector qs(begin, end); auto shifted = detail::geneval_shift_all(qs, delta); InputT zero{}; return detail::geneval_run(zero, alpha, std::move(shifted), rng.root, rng.pad, y); } template HEDLEY_WARN_UNUSED_RESULT auto geneval_sequence(wildcard_input_t, InputT x0, InputT x1, ForwardIterator begin, ForwardIterator end, ds_randomness rng, OutputT y) { return geneval_sequence(wildcard_input, x0, x1, begin, end, std::move(rng), [] { return dpf::uniform_sample(); }, y); } /// Opened comparison key material and one prefix share per endpoint. /// `live_levels` is the full depth: a comparison value word depends on the /// secret path at every level, so there is no early dummy-word tail. struct geneval_cmp_result { std::vector party0; std::vector party1; std::vector> correction_words; std::vector correction_advice; std::vector value_cw; uint64_t cw_last = 0; uint64_t addend0 = 0; uint64_t addend1 = 0; uint64_t mask = 0; std::size_t live_levels = 0; }; /// Doerner–Shelat comparison geneval. `x0 XOR x1` is the secret point, in the /// same share convention as `geneval_point`. `spec` is an `lt` / `leq` / `gt` /// / `geq` pack. Each endpoint is returned in order as the two parties' /// `eval_point(cmp, ...)` shares. An empty range opens nothing. template HEDLEY_WARN_UNUSED_RESULT geneval_cmp_result geneval_cmp(InputT x0, InputT x1, ForwardIterator begin, ForwardIterator end, ds_randomness rng, Spec spec) { geneval_cmp_result out; if (begin == end) return out; auto keys = make_dpf_doerner_shelat(std::move(x0), std::move(x1), std::move(rng), std::move(spec)); const auto & k0 = keys.first; const auto & k1 = keys.second; using key_type = std::decay_t; constexpr std::size_t depth = key_type::depth; out.live_levels = depth; out.mask = k0.cmp().mask; out.cw_last = k0.cw_last(); out.addend0 = k0.cmp_addend().raw(); out.addend1 = k1.cmp_addend().raw(); out.correction_words.resize(depth); out.correction_advice.resize(depth); out.value_cw.resize(depth); for (std::size_t level = 0; level < depth; ++level) { out.correction_words[level] = k0.correction_word(level); out.correction_advice[level] = static_cast(k0.correction_advice(level)); out.value_cw[level] = k0.value_cw(level); } for (auto it = begin; it != end; ++it) { out.party0.push_back(eval_point(dpf::cmp, k0, *it).raw()); out.party1.push_back(eval_point(dpf::cmp, k1, *it).raw()); } return out; } /// `gt(beta)` comparison geneval. `if_false` is 0. template HEDLEY_WARN_UNUSED_RESULT geneval_cmp_result geneval_cmp(InputT x0, InputT x1, ForwardIterator begin, ForwardIterator end, ds_randomness rng, uint64_t beta) { return geneval_cmp(std::move(x0), std::move(x1), begin, end, std::move(rng), dpf::gt(beta)); } } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_GENEVAL_HPP__