/// @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 #include #include #include #include #include #include #include #include #include #include #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 struct is_multipoint_key : std::false_type { }; template 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; simde__m128i sigma{}; std::uint32_t bucket_count = 0; std::uint64_t bucket_domain = 0; std::vector> buckets{}; }; template struct is_multipoint_key> : std::true_type { }; template inline constexpr bool is_multipoint_key_v = is_multipoint_key>::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(right), static_cast(right >> 64)}; auto msg = simde_mm_load_si128( reinterpret_cast(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(round + 1)); simde_mm_store_si128(reinterpret_cast(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(hash), n * kappa); located out; out.bucket = static_cast(y / bucket_domain); out.index = static_cast(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(t)); const double e = (static_cast(lambda) + 130.0 + log2t) / 123.5; auto m = static_cast(std::ceil(e * static_cast(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(words), block); return words[0] ^ (words[1] * 0x9E3779B9u) ^ words[2] ^ words[3]; } struct slot { int item = -1; int hash = -1; }; template bool insert_cuckoo(simde__m128i sigma, const std::vector & alphas, std::uint32_t m, wide n, wide bucket_domain, std::uint32_t max_evictions, std::vector & table) { table.assign(m, slot{}); std::mt19937 rng(rng_seed(sigma)); std::uniform_int_distribution pick(0, 2); const int t = static_cast(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(alphas[static_cast(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 auto make_bucket(BucketInput index, const OutputT & beta) { if constexpr (Verifiable) { return dpf::make_dpf(index, beta, dpf::verifiable{}); } else { return dpf::make_dpf(index, beta); } } template struct bucket_bare { using type = typename decltype(make_bucket(std::declval(), std::declval()).first)::key_type; }; template auto make_impl(std::vector alphas, std::vector betas, multipoint_params params) { using bare = typename bucket_bare::type; using key0 = multipoint_key<0, InputT, OutputT, bare>; using key1 = multipoint_key<1, InputT, OutputT, bare>; static_assert(std::is_unsigned_v && !std::is_same_v, "make_multipoint: input domain must be an unsigned integer"); static_assert(utils::bitlength_of_v <= 32, "make_multipoint: input domain wider than 32 bits is not supported"); static_assert(std::is_unsigned_v && !std::is_same_v, "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::numeric_limits::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(alphas.size()); const auto m = bucket_count_for(t, params.lambda); constexpr std::size_t input_bits = utils::bitlength_of_v; 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; 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(); std::vector 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(b); right.bucket_domain = static_cast(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(table[i].item)]; const auto loc = locate(sigma, static_cast(alpha), table[i].hash, n, b); if (loc.bucket != i) throw prp_walk_error{}; gamma = static_cast(loc.index); beta = betas[static_cast(table[i].item)]; } auto made = make_bucket( 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::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; 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(Key::kappa); ++hash) { const auto loc = locate(key.sigma, static_cast(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(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 HEDLEY_WARN_UNUSED_RESULT auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas, multipoint_params params = {}) { using input_type = std::decay_t; using output_type = std::decay_t; return detail::mpf::make_impl( std::vector(std::begin(alphas), std::end(alphas)), std::vector(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 HEDLEY_WARN_UNUSED_RESULT auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas, verifiable, multipoint_params params = {}) { using input_type = std::decay_t; using output_type = std::decay_t; return detail::mpf::make_impl( std::vector(std::begin(alphas), std::end(alphas)), std::vector(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 , 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 , 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 , int> = 0> void eval_multipoint(const Key & key, const Range & xs, OutIt out) { for (const auto & x : xs) *out++ = eval_multipoint(key, static_cast(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 , 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(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 , 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__