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

638 lines
22 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 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`.
/// @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 <algorithm>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <iterator>
#include <stdexcept>
#include <tuple>
#include <type_traits>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "dpf/aligned_allocator.hpp"
#include "dpf/doerner_shelat.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 <typename Output, typename Leaf>
struct geneval_result
{
std::vector<Output> party0;
std::vector<Output> party1;
std::vector<simde__m128i, aligned_allocator<simde__m128i>> correction_words;
std::vector<uint8_t> correction_advice;
std::size_t live_levels = 0;
bool leaf_live = false;
Leaf leaf{};
};
namespace detail
{
template <typename T>
HEDLEY_ALWAYS_INLINE
T geneval_mod_add(T a, T b) noexcept
{
using U = std::make_unsigned_t<T>;
U sum = static_cast<U>(static_cast<U>(a) + static_cast<U>(b));
T out;
std::memcpy(&out, &sum, sizeof(out));
return out;
}
template <typename T>
HEDLEY_ALWAYS_INLINE
T geneval_mod_sub(T a, T b) noexcept
{
using U = std::make_unsigned_t<T>;
U diff = static_cast<U>(static_cast<U>(a) - static_cast<U>(b));
T out;
std::memcpy(&out, &diff, sizeof(out));
return out;
}
template <typename T>
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 <typename Dpf>
uint64_t geneval_leaf_id(typename Dpf::input_type x)
{
return static_cast<uint64_t>(utils::get_from_node<Dpf>(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<uint64_t> & 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 <typename Output, typename Leaf>
geneval_result<Output, Leaf> geneval_empty_result()
{
geneval_result<Output, Leaf> out;
std::memset(&out.leaf, 0, sizeof(out.leaf));
return out;
}
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
auto geneval_run(InputT x0, InputT x1, const std::vector<InputT> & queries,
RootSampler & root_sampler, PadRng & pads, OutputT y)
{
static_assert(std::is_integral_v<InputT>,
"geneval input shares are an integral domain");
static_assert(!dpf::is_wildcard_v<OutputT>,
"geneval output is concrete; assign a wildcard leaf on a key");
static_assert(utils::bitlength_of_v<InputT> <= 64,
"geneval leaf ids are 64-bit");
using dpf_type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, InputT, OutputT>;
using node = typename dpf_type::interior_node;
using leaf_node = leaf_node_t<node, OutputT>;
constexpr std::size_t depth = dpf_type::depth;
if (queries.empty())
return geneval_empty_result<OutputT, leaf_node>();
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<InputT> flipped;
flipped.reserve(queries.size());
std::vector<uint64_t> 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<dpf_type>(fq));
}
std::vector<uint64_t> 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<dpf_type>(alpha);
local_cw_protocol<PadRng> proto{pads};
constexpr auto to_int = utils::to_integral_type<InputT>{};
const node root0 = dpf::unset_lo_bit(static_cast<node>(root_sampler()));
const node root1 = dpf::set_lo_bit(static_cast<node>(root_sampler()));
struct slot
{
uint64_t id;
node s0;
node s1;
};
std::vector<slot> frontier;
frontier.push_back(slot{0, root0, root1});
geneval_result<OutputT, leaf_node> 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<uint8_t>(!!(to_int(mask) & to_int(x0c)));
const uint8_t bit1 = static_cast<uint8_t>(!!(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<exp> 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<uint8_t>(pads.bit() & 1u);
const uint8_t t1 = static_cast<uint8_t>(pads.bit() & 1u);
advice = static_cast<uint8_t>((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<slot> 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<ExteriorPRG>(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<std::size_t>(to_int(flipped[i]));
result.party0.push_back(extract_leaf<node, OutputT>(share0, lane));
result.party1.push_back(extract_leaf<node, OutputT>(share1, lane));
}
return result;
}
template <typename InputT>
InputT geneval_from_bits(uint64_t bits)
{
using U = std::make_unsigned_t<InputT>;
U u = static_cast<U>(bits);
InputT out;
std::memcpy(&out, &u, sizeof(out));
return out;
}
template <typename InputT>
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<InputT>)
return from > to;
else
return utils::to_integral_type<InputT>{}(from)
> utils::to_integral_type<InputT>{}(to);
}
template <typename InputT>
std::vector<InputT> geneval_full_domain()
{
constexpr std::size_t bitlen = utils::bitlength_of_v<InputT>;
if (bitlen > 20)
throw std::length_error("geneval_full domain is too large");
const uint64_t n = uint64_t{1} << bitlen;
std::vector<InputT> qs(static_cast<std::size_t>(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<std::size_t>(i)] = geneval_from_bits<InputT>(i);
return qs;
}
template <typename InputT>
std::vector<InputT> geneval_inclusive(InputT from, InputT to)
{
if (geneval_out_of_order(from, to))
{
throw std::invalid_argument("geneval_interval: from > to");
}
std::vector<InputT> qs;
InputT q = from;
const InputT one = utils::make_from_integral_value<InputT>{}(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 <typename InputT, typename TargetSampler>
InputT geneval_sample_target(TargetSampler & sample)
{
return static_cast<InputT>(sample());
}
template <typename InputT>
std::vector<InputT> geneval_shift_all(const std::vector<InputT> & qs, InputT delta)
{
std::vector<InputT> 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 <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_point(InputT x0, InputT x1, InputT query,
ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
return detail::geneval_run<InteriorPRG, ExteriorPRG>(x0, x1,
std::vector<InputT>{query}, rng.root, rng.pad, y);
}
/// Geneval on the inclusive interval `[from, to]`.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_interval(InputT x0, InputT x1, InputT from, InputT to,
ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
return detail::geneval_run<InteriorPRG, ExteriorPRG>(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 <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_full(InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
return detail::geneval_run<InteriorPRG, ExteriorPRG>(x0, x1,
detail::geneval_full_domain<InputT>(), rng.root, rng.pad, y);
}
/// Geneval on a public sequence, in the order given.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng,
typename ForwardIterator>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_sequence(InputT x0, InputT x1, ForwardIterator begin,
ForwardIterator end, ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
std::vector<InputT> qs(begin, end);
return detail::geneval_run<InteriorPRG, ExteriorPRG>(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 <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng,
typename TargetSampler>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_point(wildcard_input_t, InputT x0, InputT x1, InputT query,
ds_randomness<RootSampler, PadRng> rng, TargetSampler sample_target,
OutputT y)
{
const InputT alpha = detail::geneval_sample_target<InputT>(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<InteriorPRG, ExteriorPRG>(zero, alpha,
std::vector<InputT>{shifted}, rng.root, rng.pad, y);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_point(wildcard_input_t, InputT x0, InputT x1, InputT query,
ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
return geneval_point<InteriorPRG, ExteriorPRG>(wildcard_input, x0, x1, query,
std::move(rng), [] { return dpf::uniform_sample<InputT>(); }, y);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng,
typename TargetSampler>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_interval(wildcard_input_t, InputT x0, InputT x1, InputT from,
InputT to, ds_randomness<RootSampler, PadRng> rng, TargetSampler sample_target,
OutputT y)
{
const InputT alpha = detail::geneval_sample_target<InputT>(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<InteriorPRG, ExteriorPRG>(zero, alpha,
std::move(shifted), rng.root, rng.pad, y);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_interval(wildcard_input_t, InputT x0, InputT x1, InputT from,
InputT to, ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
return geneval_interval<InteriorPRG, ExteriorPRG>(wildcard_input, x0, x1,
from, to, std::move(rng), [] { return dpf::uniform_sample<InputT>(); }, y);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng,
typename TargetSampler>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_full(wildcard_input_t, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, TargetSampler sample_target, OutputT y)
{
const InputT alpha = detail::geneval_sample_target<InputT>(sample_target);
const InputT delta = detail::geneval_mod_sub(alpha,
detail::geneval_mod_add(x0, x1));
InputT zero{};
auto full = detail::geneval_run<InteriorPRG, ExteriorPRG>(zero, alpha,
detail::geneval_full_domain<InputT>(), rng.root, rng.pad, y);
constexpr auto to_int = utils::to_integral_type<InputT>{};
const std::size_t n = full.party0.size();
std::vector<OutputT> p0(n), p1(n);
for (std::size_t i = 0; i < n; ++i)
{
InputT q = detail::geneval_from_bits<InputT>(i);
InputT s = detail::geneval_mod_add(q, delta);
const std::size_t si = static_cast<std::size_t>(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 <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_full(wildcard_input_t, InputT x0, InputT x1,
ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
return geneval_full<InteriorPRG, ExteriorPRG>(wildcard_input, x0, x1,
std::move(rng), [] { return dpf::uniform_sample<InputT>(); }, y);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng,
typename ForwardIterator,
typename TargetSampler>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_sequence(wildcard_input_t, InputT x0, InputT x1,
ForwardIterator begin, ForwardIterator end,
ds_randomness<RootSampler, PadRng> rng, TargetSampler sample_target, OutputT y)
{
const InputT alpha = detail::geneval_sample_target<InputT>(sample_target);
const InputT delta = detail::geneval_mod_sub(alpha,
detail::geneval_mod_add(x0, x1));
std::vector<InputT> qs(begin, end);
auto shifted = detail::geneval_shift_all(qs, delta);
InputT zero{};
return detail::geneval_run<InteriorPRG, ExteriorPRG>(zero, alpha,
std::move(shifted), rng.root, rng.pad, y);
}
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename PadRng,
typename ForwardIterator>
HEDLEY_WARN_UNUSED_RESULT
auto geneval_sequence(wildcard_input_t, InputT x0, InputT x1,
ForwardIterator begin, ForwardIterator end,
ds_randomness<RootSampler, PadRng> rng, OutputT y)
{
return geneval_sequence<InteriorPRG, ExteriorPRG>(wildcard_input, x0, x1,
begin, end, std::move(rng), [] { return dpf::uniform_sample<InputT>(); }, y);
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_GENEVAL_HPP__