libdpf/include/dpf/multipoint.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

797 lines
26 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/multipoint.hpp
/// @brief Cuckoo-packed multi-point DPF and verifiable multi-point DPF.
/// @details Packs t distinct points into m ≈ O(t) buckets (de Castro–
/// Polychroniadou, EUROCRYPT 2022, §4, ePrint 2021/580). Each bucket is an ordinary
/// point key on a smaller domain — `dpf::verifiable` selects VDPF
/// buckets. Evaluation probes κ = 3 buckets and sums the shares.
/// A batched proof is one 2λ token.
/// @note Following that section: κ = 3 cuckoo hashes, one point key per bucket.
/// @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_MULTIPOINT_HPP__
#define LIBDPF_INCLUDE_DPF_MULTIPOINT_HPP__
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <iterator>
#include <limits>
#include <random>
#include <stdexcept>
#include <type_traits>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "dpf/eval_point.hpp"
#include "dpf/incremental.hpp"
#include "dpf/prg_aes.hpp"
#include "dpf/random.hpp"
#include "dpf/secret_share.hpp"
#include "dpf/uint256_t.hpp"
#include "dpf/verifiable.hpp"
namespace dpf
{
/// @brief Knobs for cuckoo packing. `lambda` is the Remark 1 failure target.
struct multipoint_params
{
std::uint32_t lambda = 40;
std::uint32_t max_evictions = 4096;
int retries = 8;
};
/// @brief 512-bit word for the cuckoo PRP. Holds `3·2^b` for every input
/// width this library can form a point key on (up to 256 bits).
struct mpf_word
{
uint256_t lo{};
uint256_t hi{};
friend bool operator==(mpf_word a, mpf_word b) noexcept
{
return a.lo == b.lo && a.hi == b.hi;
}
friend bool operator<(mpf_word a, mpf_word b) noexcept
{
if (a.hi != b.hi)
return a.hi < b.hi;
return a.lo < b.lo;
}
};
template <typename T>
struct is_multipoint_key : std::false_type
{
};
template <std::size_t Party,
typename InputT,
typename OutputT,
typename BucketKey>
struct multipoint_key
{
static constexpr std::size_t party = Party;
static constexpr bool is_multipoint = true;
static constexpr bool is_verifiable = BucketKey::is_verifiable;
static constexpr std::size_t kappa = 3;
using input_type = InputT;
using output_type = OutputT;
using bucket_key = BucketKey;
using bucket_input = typename BucketKey::input_type;
using share_type = subtractive_share<OutputT, Party>;
simde__m128i sigma{};
std::uint64_t bucket_count = 0;
mpf_word bucket_domain{};
std::vector<party_key<Party, BucketKey>> buckets{};
};
template <std::size_t Party, typename InputT, typename OutputT, typename BucketKey>
struct is_multipoint_key<multipoint_key<Party, InputT, OutputT, BucketKey>>
: std::true_type
{
};
template <typename T>
inline constexpr bool is_multipoint_key_v =
is_multipoint_key<std::decay_t<T>>::value;
namespace detail
{
namespace mpf
{
struct prp_walk_error : std::runtime_error
{
prp_walk_error()
: std::runtime_error("multipoint PRP cycle walk exceeded its bound")
{
}
};
struct located
{
std::uint64_t bucket = 0;
mpf_word index{};
};
inline mpf_word word_add(mpf_word a, mpf_word b)
{
mpf_word r;
r.lo = a.lo + b.lo;
r.hi = a.hi + b.hi;
if (r.lo < a.lo)
r.hi = r.hi + uint256_t{1};
return r;
}
inline mpf_word word_sub(mpf_word a, mpf_word b)
{
mpf_word r;
r.lo = a.lo - b.lo;
r.hi = a.hi - b.hi;
if (a.lo < b.lo)
r.hi = r.hi - uint256_t{1};
return r;
}
inline mpf_word word_shl(mpf_word a, unsigned shift)
{
if (shift == 0)
return a;
if (shift >= 512)
return {};
if (shift >= 256)
{
mpf_word r;
r.hi = a.lo << (shift - 256);
return r;
}
mpf_word r;
r.lo = a.lo << shift;
r.hi = (a.hi << shift) | (a.lo >> (256 - shift));
return r;
}
inline mpf_word word_shr(mpf_word a, unsigned shift)
{
if (shift == 0)
return a;
if (shift >= 512)
return {};
if (shift >= 256)
{
mpf_word r;
r.lo = a.hi >> (shift - 256);
return r;
}
mpf_word r;
r.hi = a.hi >> shift;
r.lo = (a.lo >> shift) | (a.hi << (256 - shift));
return r;
}
inline mpf_word word_or(mpf_word a, mpf_word b)
{
a.lo = a.lo | b.lo;
a.hi = a.hi | b.hi;
return a;
}
inline mpf_word word_and(mpf_word a, mpf_word b)
{
a.lo = a.lo & b.lo;
a.hi = a.hi & b.hi;
return a;
}
inline bool word_bit(mpf_word a, unsigned bit)
{
if (bit >= 512)
return false;
if (bit >= 256)
return static_cast<bool>((a.hi >> (bit - 256)) & uint256_t{1});
return static_cast<bool>((a.lo >> bit) & uint256_t{1});
}
inline int word_bit_length(mpf_word a)
{
for (int i = 255; i >= 0; --i)
{
if (static_cast<bool>((a.hi >> i) & uint256_t{1}))
return i + 1 + 256;
}
for (int i = 255; i >= 0; --i)
{
if (static_cast<bool>((a.lo >> i) & uint256_t{1}))
return i + 1;
}
return 0;
}
inline mpf_word word_mul_small(mpf_word a, std::uint64_t k)
{
mpf_word r{};
while (k != 0)
{
if ((k & 1u) != 0)
r = word_add(r, a);
a = word_shl(a, 1);
k >>= 1;
}
return r;
}
inline std::pair<mpf_word, mpf_word> word_divmod(mpf_word num, mpf_word den)
{
if (den == mpf_word{})
throw std::invalid_argument("multipoint division by zero");
mpf_word q{};
mpf_word r{};
const int top = word_bit_length(num);
for (int i = top - 1; i >= 0; --i)
{
r = word_shl(r, 1);
if (word_bit(num, static_cast<unsigned>(i)))
r = word_add(r, mpf_word{uint256_t{1}, uint256_t{0}});
if (!(r < den))
{
r = word_sub(r, den);
mpf_word bit{};
if (i >= 256)
bit.hi = uint256_t{1} << static_cast<unsigned>(i - 256);
else
bit.lo = uint256_t{1} << static_cast<unsigned>(i);
q = word_or(q, bit);
}
}
return {q, r};
}
inline mpf_word domain_size(std::size_t bits)
{
mpf_word r{};
if (bits >= 512)
throw std::invalid_argument("multipoint domain shift is out of range");
if (bits >= 256)
r.hi = uint256_t{1} << (bits - 256);
else if (bits > 0)
r.lo = uint256_t{1} << bits;
return r;
}
template <typename T>
mpf_word to_word(T x)
{
constexpr std::size_t bits = utils::bitlength_of_v<T>;
mpf_word w{};
if constexpr (bits > 128)
{
w.lo = static_cast<uint256_t>(x);
}
else if constexpr (bits > 64)
{
uint128_t low{};
std::memcpy(&low, &x, sizeof(T));
w.lo = uint256_t{low};
}
else
{
w.lo = uint256_t{static_cast<std::uint64_t>(x)};
}
return w;
}
template <typename T>
T from_word(mpf_word w)
{
constexpr std::size_t bits = utils::bitlength_of_v<T>;
if constexpr (bits > 128)
{
return static_cast<T>(w.lo);
}
else if constexpr (bits > 64)
{
const uint128_t low = static_cast<uint128_t>(w.lo);
T out{};
std::memcpy(&out, &low, sizeof(T));
return out;
}
else
{
return static_cast<T>(static_cast<std::uint64_t>(w.lo));
}
}
/// @brief Low `half` bits set, as a 512-bit mask. `half <= 0` is zero.
inline mpf_word low_mask(int half)
{
if (half <= 0)
return {};
if (half >= 512)
{
mpf_word all;
all.lo = ~uint256_t{0};
all.hi = ~uint256_t{0};
return all;
}
return word_sub(domain_size(static_cast<std::size_t>(half)),
mpf_word{uint256_t{1}, uint256_t{0}});
}
inline mpf_word aes_prf(simde__m128i seed, mpf_word right, int round)
{
alignas(16) unsigned char raw[32]{};
std::memcpy(raw, &right.lo, sizeof(right.lo));
alignas(16) simde__m128i block0;
alignas(16) simde__m128i block1;
std::memcpy(&block0, raw, 16);
std::memcpy(&block1, raw + 16, 16);
block0 = simde_mm_xor_si128(block0, seed);
block0 = simde_mm_xor_si128(block0, simde_mm_set_epi32(0, 0, 0, round + 1));
const auto out0 = prg::aes128::eval(block0,
static_cast<psnip_uint32_t>(round + 1));
block1 = simde_mm_xor_si128(block1, seed);
block1 = simde_mm_xor_si128(block1,
simde_mm_set_epi32(0, 0, 0, round + 0x11));
const auto out1 = prg::aes128::eval(block1,
static_cast<psnip_uint32_t>(round + 0x21));
alignas(16) unsigned char packed[32];
std::memcpy(packed, &out0, 16);
std::memcpy(packed + 16, &out1, 16);
mpf_word f{};
std::memcpy(&f.lo, packed, sizeof(f.lo));
return f;
}
/// @brief 4-round Feistel on the next power-of-two square, then cycle-walk
/// into `[0, domain)`. AES-MMO is the round function.
/// @param seed the PRP seed
/// @param x the input, in `[0, domain)`
/// @param domain the domain size
/// @return the permuted value in `[0, domain)`
/// @throws std::invalid_argument if `x` is outside the domain
/// @throws prp_walk_error if the cycle walk exceeds its bound
inline mpf_word permute(simde__m128i seed, mpf_word x, mpf_word domain)
{
const mpf_word one{uint256_t{1}, uint256_t{0}};
if (!(one < domain))
return {};
if (!(x < domain))
throw std::invalid_argument("multipoint PRP input is outside the domain");
const int bits = word_bit_length(word_sub(domain, one));
const int half = (bits + 1) / 2;
const mpf_word mask = low_mask(half);
mpf_word val = x;
for (int guard = 0; guard < 128; ++guard)
{
mpf_word left = word_and(word_shr(val, static_cast<unsigned>(half)), mask);
mpf_word right = word_and(val, mask);
for (int round = 0; round < 4; ++round)
{
const mpf_word f = word_and(aes_prf(seed, right, round), mask);
left.lo = left.lo ^ f.lo;
left.hi = left.hi ^ f.hi;
const mpf_word tmp = left;
left = right;
right = tmp;
}
val = word_or(word_shl(left, static_cast<unsigned>(half)), right);
if (val < domain)
return val;
}
throw prp_walk_error{};
}
inline located locate(simde__m128i sigma, mpf_word x, int hash,
mpf_word n, mpf_word bucket_domain)
{
constexpr int kappa = 3;
const mpf_word y = permute(sigma,
word_add(x, word_mul_small(n, static_cast<std::uint64_t>(hash))),
word_mul_small(n, kappa));
const auto [quot, rem] = word_divmod(y, bucket_domain);
if (quot.hi != uint256_t{0})
throw std::runtime_error("multipoint bucket index does not fit");
located out;
out.bucket = static_cast<std::uint64_t>(quot.lo);
out.index = rem;
return out;
}
inline std::uint64_t bucket_count_for(std::uint64_t t, std::uint32_t lambda)
{
const double log2t = (t <= 1) ? 0.0 : std::log2(static_cast<double>(t));
const double e = (static_cast<double>(lambda) + 130.0 + log2t) / 123.5;
auto m = static_cast<std::uint64_t>(std::ceil(e * static_cast<double>(t)));
if (m < t + 1)
m = t + 1;
// Remark 1's simplification wants t ≥ 30. Below that, keep a 2t table.
if (t < 30 && m < t * 2)
m = t * 2;
// At least κ buckets so each within-bucket index fits in the input type.
if (m < 3)
m = 3;
return m;
}
inline std::uint32_t rng_seed(simde__m128i sigma)
{
const auto block = prg::aes128::eval(sigma, 0xC000u);
alignas(16) std::uint32_t words[4];
simde_mm_store_si128(reinterpret_cast<simde__m128i *>(words), block);
return words[0] ^ (words[1] * 0x9E3779B9u) ^ words[2] ^ words[3];
}
struct slot
{
std::int64_t item = -1;
int hash = -1;
};
template <typename InputT>
bool insert_cuckoo(simde__m128i sigma, const std::vector<InputT> & alphas,
std::uint64_t m, mpf_word n, mpf_word bucket_domain,
std::uint32_t max_evictions, std::vector<slot> & table)
{
table.assign(m, slot{});
std::mt19937 rng(rng_seed(sigma));
std::uniform_int_distribution<int> pick(0, 2);
const auto t = static_cast<std::int64_t>(alphas.size());
for (std::int64_t omega = 0; omega < t; ++omega)
{
std::int64_t cur = omega;
int hash = pick(rng);
std::uint32_t evictions = 0;
for (;;)
{
const auto loc = locate(sigma,
to_word(alphas[static_cast<std::size_t>(cur)]),
hash, n, bucket_domain);
if (loc.bucket >= m)
return false;
if (table[loc.bucket].item < 0)
{
table[loc.bucket] = slot{cur, hash};
break;
}
const std::int64_t evicted = table[loc.bucket].item;
table[loc.bucket] = slot{cur, hash};
cur = evicted;
hash = pick(rng);
if (++evictions > max_evictions)
return false;
}
}
return true;
}
template <bool Verifiable,
typename InteriorPRG,
typename ExteriorPRG,
typename BucketInput,
typename OutputT>
auto make_bucket(BucketInput index, const OutputT & beta)
{
if constexpr (Verifiable)
{
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(index, beta,
dpf::verifiable{});
}
else
{
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(index, beta);
}
}
template <bool Verifiable, typename InteriorPRG, typename ExteriorPRG,
typename BucketInput, typename OutputT>
struct bucket_bare
{
using type = typename decltype(make_bucket<Verifiable, InteriorPRG,
ExteriorPRG>(std::declval<BucketInput>(),
std::declval<const OutputT &>()).first)::key_type;
};
template <bool Verifiable,
typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT>
auto make_impl(std::vector<InputT> alphas, std::vector<OutputT> betas,
multipoint_params params)
{
using bare = typename bucket_bare<Verifiable, InteriorPRG, ExteriorPRG,
InputT, OutputT>::type;
using key0 = multipoint_key<0, InputT, OutputT, bare>;
using key1 = multipoint_key<1, InputT, OutputT, bare>;
static_assert(!std::is_same_v<InputT, bool>
&& (std::is_unsigned_v<InputT> || std::is_same_v<InputT, uint256_t>),
"make_multipoint: input domain must be an unsigned integer");
static_assert(utils::bitlength_of_v<InputT> <= 256,
"make_multipoint: input type is wider than a point key in this library");
if (alphas.size() != betas.size())
throw std::invalid_argument("make_multipoint: point and payload counts differ");
if (alphas.empty())
throw std::invalid_argument("make_multipoint: no points");
{
auto sorted = alphas;
std::sort(sorted.begin(), sorted.end());
if (std::adjacent_find(sorted.begin(), sorted.end()) != sorted.end())
throw std::invalid_argument("make_multipoint: duplicate points");
}
const auto t = static_cast<std::uint64_t>(alphas.size());
const auto m = bucket_count_for(t, params.lambda);
constexpr std::size_t input_bits = utils::bitlength_of_v<InputT>;
const mpf_word n = domain_size(input_bits);
constexpr int kappa = 3;
const mpf_word span = word_mul_small(n, kappa);
const mpf_word den{uint256_t{m}, uint256_t{0}};
const mpf_word numer = word_add(span,
word_sub(den, mpf_word{uint256_t{1}, uint256_t{0}}));
const mpf_word b = word_divmod(numer, den).first;
const int attempts = params.retries < 1 ? 1 : params.retries;
for (int attempt = 0; attempt < attempts; ++attempt)
{
try
{
const simde__m128i sigma = dpf::uniform_sample<simde__m128i>();
std::vector<slot> table;
if (!insert_cuckoo(sigma, alphas, m, n, b, params.max_evictions, table))
continue;
key0 left;
key1 right;
left.sigma = sigma;
right.sigma = sigma;
left.bucket_count = m;
right.bucket_count = m;
left.bucket_domain = b;
right.bucket_domain = b;
left.buckets.reserve(static_cast<std::size_t>(m));
right.buckets.reserve(static_cast<std::size_t>(m));
for (std::uint64_t i = 0; i < m; ++i)
{
InputT gamma{};
OutputT beta{};
if (table[static_cast<std::size_t>(i)].item >= 0)
{
const auto & alpha = alphas[static_cast<std::size_t>(
table[static_cast<std::size_t>(i)].item)];
const auto loc = locate(sigma, to_word(alpha),
table[static_cast<std::size_t>(i)].hash, n, b);
if (loc.bucket != i)
throw prp_walk_error{};
gamma = from_word<InputT>(loc.index);
beta = betas[static_cast<std::size_t>(
table[static_cast<std::size_t>(i)].item)];
}
auto made = make_bucket<Verifiable, InteriorPRG, ExteriorPRG>(
gamma, beta);
left.buckets.push_back(std::move(made.first));
right.buckets.push_back(std::move(made.second));
}
return std::make_pair(std::move(left), std::move(right));
}
catch (const prp_walk_error &)
{
continue;
}
}
throw std::runtime_error("make_multipoint: cuckoo hashing failed");
}
inline void absorb_proof(proof_token & acc, const proof_token & inner)
{
acc = detail::vdpf::xor_proof(acc, inner);
acc[0] = detail::vdpf::mmo(acc[0], 1);
acc[1] = detail::vdpf::mmo(acc[1], 2);
}
template <typename Key>
typename Key::share_type eval_at(const Key & key, typename Key::input_type x,
proof_token * acc)
{
using input_type = typename Key::input_type;
using bucket_input = typename Key::bucket_input;
constexpr std::size_t input_bits = utils::bitlength_of_v<input_type>;
const mpf_word n = domain_size(input_bits);
const mpf_word b = key.bucket_domain;
typename Key::share_type sum =
Key::share_type::from_raw(typename Key::output_type{});
for (int hash = 0; hash < static_cast<int>(Key::kappa); ++hash)
{
const auto loc = locate(key.sigma, to_word(x), hash, n, b);
if (loc.bucket >= key.bucket_count)
throw std::runtime_error("multipoint eval: bucket out of range");
const auto gamma = from_word<bucket_input>(loc.index);
const auto & bucket = key.buckets[loc.bucket];
if constexpr (Key::is_verifiable)
{
if (acc != nullptr)
{
proof_token inner{};
sum += *dpf::eval_point(bucket, gamma, dpf::prove(inner));
absorb_proof(*acc, inner);
continue;
}
}
sum += *dpf::eval_point(bucket, gamma);
}
return sum;
}
} // namespace mpf
} // namespace detail
/// @brief Cuckoo-pack distinct points into ordinary point-key buckets.
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam AlphaRange range of distinct domain points
/// @tparam BetaRange range of payloads, one per point
/// @param alphas the secret points
/// @param betas the payloads
/// @param params packing knobs. `lambda` is the Remark 1 failure target
/// @return the two party keys
/// @throws std::invalid_argument if the lists differ in length, are empty,
/// or contain a duplicate
/// @throws std::runtime_error if cuckoo hashing does not succeed
/// @note Following de Castro and Polychroniadou, EUROCRYPT 2022, §4 (ePrint 2021/580): κ = 3 cuckoo buckets, one point key each.
/// \complexity For t points, `bucket_count_for` sets m = ceil((lambda + 130 + log2(t)) / 123.5 * t) (at least 2t when t < 30, and at least 3). Each attempt inserts t cuckoo items (up to `max_evictions` swaps each) and then one `make_dpf` per bucket. The within-bucket domain `b` is ceil(3 * 2^{input bits} / m). Counted `insert_cuckoo` and the bucket loop. Retries are `params.retries`.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename AlphaRange,
typename BetaRange>
HEDLEY_WARN_UNUSED_RESULT
auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
multipoint_params params = {})
{
using input_type = std::decay_t<decltype(*std::begin(alphas))>;
using output_type = std::decay_t<decltype(*std::begin(betas))>;
return detail::mpf::make_impl<false, InteriorPRG, ExteriorPRG>(
std::vector<input_type>(std::begin(alphas), std::end(alphas)),
std::vector<output_type>(std::begin(betas), std::end(betas)),
params);
}
/// @brief Same packing as `make_multipoint`, with a verifiable bucket key.
/// @see `make_multipoint`
/// @param alphas the secret points
/// @param betas the payloads
/// @param params packing knobs
/// @return the two verifiable party keys
/// @throws std::invalid_argument if the lists differ in length, are empty,
/// or contain a duplicate
/// @throws std::runtime_error if cuckoo hashing does not succeed
/// @note Following de Castro and Polychroniadou, EUROCRYPT 2022, §4 (ePrint 2021/580): κ = 3 cuckoo buckets, one point key each.
/// \complexity For t points, `bucket_count_for` sets m = ceil((lambda + 130 + log2(t)) / 123.5 * t) (at least 2t when t < 30, and at least 3). Each attempt inserts t cuckoo items (up to `max_evictions` swaps each) and then one `make_dpf` per bucket. The within-bucket domain `b` is ceil(3 * 2^{input bits} / m). Counted `insert_cuckoo` and the bucket loop. Retries are `params.retries`.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename AlphaRange,
typename BetaRange>
HEDLEY_WARN_UNUSED_RESULT
auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
verifiable, multipoint_params params = {})
{
using input_type = std::decay_t<decltype(*std::begin(alphas))>;
using output_type = std::decay_t<decltype(*std::begin(betas))>;
return detail::mpf::make_impl<true, InteriorPRG, ExteriorPRG>(
std::vector<input_type>(std::begin(alphas), std::end(alphas)),
std::vector<output_type>(std::begin(betas), std::end(betas)),
params);
}
/// @brief Sum the three bucket shares at `x`.
/// @tparam Key a `multipoint_key`
/// @param key the party key
/// @param x the query point
/// @return the party's share of the payload, or of zero off the packed points
/// @throws std::runtime_error if a located bucket is outside the key
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
template <typename Key,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
auto eval_multipoint(const Key & key, typename Key::input_type x)
{
return detail::mpf::eval_at(key, x, nullptr);
}
/// @brief Evaluate `x` and fold that query into `pr`.
/// @tparam Key a verifiable `multipoint_key`
/// @param key the party key
/// @param x the query point
/// @param pr proof token replaced with this query's folded proof
/// @return the party's share of the payload
/// @throws std::runtime_error if a located bucket is outside the key
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
template <typename Key,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
auto eval_multipoint(const Key & key, typename Key::input_type x, prove_ref pr)
{
static_assert(Key::is_verifiable,
"eval_multipoint(..., prove(π)): key must be a verifiable multipoint key");
pr.token = detail::vdpf::zero_proof();
return detail::mpf::eval_at(key, x, &pr.token);
}
/// @brief Evaluate each point of `xs`, writing one share per point.
/// @tparam Key a `multipoint_key`
/// @tparam Range range of query points
/// @tparam OutIt output iterator of shares
/// @param key the party key
/// @param xs the query points
/// @param out where each share is written
/// @throws std::runtime_error if a located bucket is outside the key
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
template <typename Key, typename Range, typename OutIt,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
void eval_multipoint(const Key & key, const Range & xs, OutIt out)
{
for (const auto & x : xs)
*out++ = eval_multipoint(key, static_cast<typename Key::input_type>(x));
}
/// @brief Evaluate `xs` and fold every query into one proof.
/// @tparam Key a verifiable `multipoint_key`
/// @tparam Range range of query points
/// @tparam OutIt output iterator of shares
/// @param key the party key
/// @param xs the query points
/// @param out where each share is written
/// @param pr proof token replaced with the folded proof of `xs`
/// @throws std::runtime_error if a located bucket is outside the key
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
template <typename Key, typename Range, typename OutIt,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
void eval_multipoint(const Key & key, const Range & xs, OutIt out, prove_ref pr)
{
static_assert(Key::is_verifiable,
"eval_multipoint(..., prove(π)): key must be a verifiable multipoint key");
pr.token = detail::vdpf::zero_proof();
for (const auto & x : xs)
{
*out++ = detail::mpf::eval_at(key,
static_cast<typename Key::input_type>(x), &pr.token);
}
}
/// @brief Fold a canonical evaluation of every bucket into one proof.
/// @tparam Key a verifiable `multipoint_key`
/// @param key the party key
/// @param pr proof token replaced with the audit proof
template <typename Key,
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
void audit_multipoint(const Key & key, prove_ref pr)
{
static_assert(Key::is_verifiable,
"audit_multipoint: key must be a verifiable multipoint key");
pr.token = detail::vdpf::zero_proof();
for (const auto & bucket : key.buckets)
{
proof_token inner{};
(void)*dpf::eval_point(bucket, typename Key::bucket_input{},
dpf::prove(inner));
detail::mpf::absorb_proof(pr.token, inner);
}
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_MULTIPOINT_HPP__