/// @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 #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/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 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::uint64_t bucket_count = 0; mpf_word bucket_domain{}; 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::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((a.hi >> (bit - 256)) & uint256_t{1}); return static_cast((a.lo >> bit) & uint256_t{1}); } inline int word_bit_length(mpf_word a) { for (int i = 255; i >= 0; --i) { if (static_cast((a.hi >> i) & uint256_t{1})) return i + 1 + 256; } for (int i = 255; i >= 0; --i) { if (static_cast((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 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(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(i - 256); else bit.lo = uint256_t{1} << static_cast(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 mpf_word to_word(T x) { constexpr std::size_t bits = utils::bitlength_of_v; mpf_word w{}; if constexpr (bits > 128) { w.lo = static_cast(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(x)}; } return w; } template T from_word(mpf_word w) { constexpr std::size_t bits = utils::bitlength_of_v; if constexpr (bits > 128) { return static_cast(w.lo); } else if constexpr (bits > 64) { const uint128_t low = static_cast(w.lo); T out{}; std::memcpy(&out, &low, sizeof(T)); return out; } else { return static_cast(static_cast(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(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(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(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(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(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(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(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(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; // 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(words), block); return words[0] ^ (words[1] * 0x9E3779B9u) ^ words[2] ^ words[3]; } struct slot { std::int64_t item = -1; int hash = -1; }; template bool insert_cuckoo(simde__m128i sigma, const std::vector & alphas, std::uint64_t m, mpf_word n, mpf_word 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 auto t = static_cast(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(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 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_same_v && (std::is_unsigned_v || std::is_same_v), "make_multipoint: input domain must be an unsigned integer"); static_assert(utils::bitlength_of_v <= 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(alphas.size()); const auto m = bucket_count_for(t, params.lambda); constexpr std::size_t input_bits = utils::bitlength_of_v; 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(); 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 = b; right.bucket_domain = b; left.buckets.reserve(static_cast(m)); right.buckets.reserve(static_cast(m)); for (std::uint64_t i = 0; i < m; ++i) { InputT gamma{}; OutputT beta{}; if (table[static_cast(i)].item >= 0) { const auto & alpha = alphas[static_cast( table[static_cast(i)].item)]; const auto loc = locate(sigma, to_word(alpha), table[static_cast(i)].hash, n, b); if (loc.bucket != i) throw prp_walk_error{}; gamma = from_word(loc.index); beta = betas[static_cast( table[static_cast(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); acc[1] = detail::vdpf::mmo(acc[1], 2); } 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 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(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(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 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, /// 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 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 /// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth. 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 /// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth. 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 /// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth. 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 /// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth. 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__