libdpf/include/dpf/multipoint.hpp

560 lines
20 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). 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__