Document the new DPF surfaces in one command set, and test the field, half-tree, and multipoint edges.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Ryan Henry 2026-09-24 23:18:10 -06:00
parent 0d8a5a8131
commit 0dff6df8ed
250 changed files with 12199 additions and 1981 deletions

560
include/dpf/multipoint.hpp Normal file
View file

@ -0,0 +1,560 @@
/// @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). 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.
/// @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/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;
};
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::uint32_t bucket_count = 0;
std::uint64_t bucket_domain = 0;
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::uint32_t bucket = 0;
std::uint64_t index = 0;
};
using wide = unsigned __int128;
inline wide domain_size(std::size_t bits)
{
return wide{1} << bits;
}
/// @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 wide permute(simde__m128i seed, wide x, wide domain)
{
if (domain <= 1)
return 0;
if (x >= domain)
throw std::invalid_argument("multipoint PRP input is outside the domain");
int bits = 0;
for (wide v = domain - 1; v > 0; v >>= 1)
++bits;
const int half = (bits + 1) / 2;
const wide mask = (half >= 128)
? ~wide{0}
: (wide{1} << half) - 1;
wide val = x;
for (int guard = 0; guard < 128; ++guard)
{
unsigned __int128 left = (val >> half) & mask;
unsigned __int128 right = val & mask;
for (int round = 0; round < 4; ++round)
{
alignas(16) std::uint64_t lanes[2] = {
static_cast<std::uint64_t>(right),
static_cast<std::uint64_t>(right >> 64)};
auto msg = simde_mm_load_si128(
reinterpret_cast<const simde__m128i *>(lanes));
msg = simde_mm_xor_si128(msg, seed);
msg = simde_mm_xor_si128(msg,
simde_mm_set_epi32(0, 0, 0, round + 1));
const auto out = prg::aes128::eval(msg,
static_cast<psnip_uint32_t>(round + 1));
simde_mm_store_si128(reinterpret_cast<simde__m128i *>(lanes), out);
wide f = lanes[0] | (wide{lanes[1]} << 64);
f &= mask;
left ^= f;
const wide tmp = left;
left = right;
right = tmp;
}
val = (left << half) | right;
if (val < domain)
return val;
}
throw prp_walk_error{};
}
inline located locate(simde__m128i sigma, wide x, int hash,
wide n, wide bucket_domain)
{
constexpr int kappa = 3;
const wide y = permute(sigma,
x + n * static_cast<unsigned>(hash), n * kappa);
located out;
out.bucket = static_cast<std::uint32_t>(y / bucket_domain);
out.index = static_cast<std::uint64_t>(y % bucket_domain);
return out;
}
inline std::uint32_t bucket_count_for(std::uint32_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::uint32_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;
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
{
int item = -1;
int hash = -1;
};
template <typename InputT>
bool insert_cuckoo(simde__m128i sigma, const std::vector<InputT> & alphas,
std::uint32_t m, wide n, wide 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 int t = static_cast<int>(alphas.size());
for (int omega = 0; omega < t; ++omega)
{
int cur = omega;
int hash = pick(rng);
std::uint32_t evictions = 0;
for (;;)
{
const auto loc = locate(sigma,
static_cast<wide>(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 int 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 BucketInput,
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,
BucketInput, OutputT>::type;
using key0 = multipoint_key<0, InputT, OutputT, bare>;
using key1 = multipoint_key<1, InputT, OutputT, bare>;
static_assert(std::is_unsigned_v<InputT> && !std::is_same_v<InputT, bool>,
"make_multipoint: input domain must be an unsigned integer");
static_assert(utils::bitlength_of_v<InputT> <= 32,
"make_multipoint: input domain wider than 32 bits is not supported");
static_assert(std::is_unsigned_v<BucketInput>
&& !std::is_same_v<BucketInput, bool>,
"make_multipoint: BucketInput must be an unsigned integer");
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");
if (alphas.size() > static_cast<std::size_t>(std::numeric_limits<std::uint32_t>::max()))
throw std::invalid_argument("make_multipoint: too many 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::uint32_t>(alphas.size());
const auto m = bucket_count_for(t, params.lambda);
constexpr std::size_t input_bits = utils::bitlength_of_v<InputT>;
const wide n = domain_size(input_bits);
constexpr int kappa = 3;
const wide b = (n * kappa + m - 1) / m;
constexpr std::size_t bucket_bits = utils::bitlength_of_v<BucketInput>;
const wide bucket_cap = domain_size(bucket_bits);
if (b > bucket_cap)
{
throw std::invalid_argument(
"make_multipoint: bucket domain does not fit in BucketInput");
}
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 = static_cast<std::uint64_t>(b);
right.bucket_domain = static_cast<std::uint64_t>(b);
left.buckets.reserve(m);
right.buckets.reserve(m);
for (std::uint32_t i = 0; i < m; ++i)
{
BucketInput gamma{};
OutputT beta{};
if (table[i].item >= 0)
{
const auto & alpha = alphas[static_cast<std::size_t>(table[i].item)];
const auto loc = locate(sigma,
static_cast<wide>(alpha), table[i].hash, n, b);
if (loc.bucket != i)
throw prp_walk_error{};
gamma = static_cast<BucketInput>(loc.index);
beta = betas[static_cast<std::size_t>(table[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);
}
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 wide n = domain_size(input_bits);
const wide 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, static_cast<wide>(x), hash, n, b);
if (loc.bucket >= key.bucket_count)
throw std::runtime_error("multipoint eval: bucket out of range");
const auto gamma = static_cast<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 BucketInput unsigned type of a bucket index. Defaults to `uint32_t`
/// @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,
/// contain a duplicate, or a bucket index does not fit `BucketInput`
/// @throws std::runtime_error if cuckoo hashing does not succeed
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename BucketInput = std::uint32_t,
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, BucketInput>(
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,
/// contain a duplicate, or a bucket index does not fit `BucketInput`
/// @throws std::runtime_error if cuckoo hashing does not succeed
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename BucketInput = std::uint32_t,
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, BucketInput>(
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
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
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
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
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__