/// @file dpf/doerner_shelat.hpp /// @brief Doerner–Shelat generation of a dealer DPF key. /// @details Two XOR shares of the point are walked level by level. Correction /// words, advice bits, seeds, and leaves are the ones `make_dpf` /// would emit for the XOR of those shares, the same roots, and the /// same beaver coins. Beaver pads used to hide the path bit cancel /// and are not part of the key. Pad randomness must not come from /// `uniform_fill` if the beaver tape is being matched. /// @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_DOERNER_SHELAT_HPP__ #define LIBDPF_INCLUDE_DPF_DOERNER_SHELAT_HPP__ #include #include #include #include #include "hedley/hedley.h" #include "simde/simde/x86/avx2.h" #include "dpf/dpf_key.hpp" #include "dpf/random.hpp" #include "dpf/dcf.hpp" namespace dpf { /// Roots and the Beaver-pad stream for one Doerner–Shelat generation. /// `root` is called twice, same as `make_dpf`: party 0 clears the low bit of /// the first sample, party 1 sets the low bit of the second. template struct ds_randomness { RootSampler root; PadRng pad; }; namespace detail { struct urandom_pad_rng { simde__m128i block() { return dpf::uniform_sample(); } uint8_t bit() { return static_cast(dpf::uniform_sample() & 1u); } }; struct ds_cw_party { simde__m128i rand; simde__m128i gamma; uint8_t bit; }; struct ds_cw_pads { ds_cw_party p0; ds_cw_party p1; }; struct ds_blind { simde__m128i msg; uint8_t bit; }; struct ds_and_pads { uint8_t a0; uint8_t a1; simde__m128i b0_share, b1_share, c0_share, c1_share; }; struct ds_and_shares { simde__m128i z0; simde__m128i z1; }; HEDLEY_ALWAYS_INLINE simde__m128i ds_xor(simde__m128i a, simde__m128i b) noexcept { return simde_mm_xor_si128(a, b); } HEDLEY_ALWAYS_INLINE simde__m128i ds_gate(uint8_t bit, simde__m128i block) noexcept { return dpf::get_if(block, bit & 1u); } template ds_cw_pads ds_sample_cw(PadRng & pad) { ds_cw_pads p{}; const simde__m128i zero = simde_mm_setzero_si128(); p.p0.rand = pad.block(); p.p1.rand = pad.block(); p.p0.bit = static_cast(pad.bit() & 1u); p.p1.bit = static_cast(pad.bit() & 1u); p.p0.gamma = p.p1.bit ? p.p0.rand : zero; p.p1.gamma = p.p0.bit ? p.p1.rand : zero; return p; } template ds_and_pads ds_sample_and(PadRng & pad) { ds_and_pads p{}; const uint8_t a = static_cast(pad.bit() & 1u); const simde__m128i B = pad.block(); const simde__m128i C = ds_gate(a, B); p.a0 = static_cast(pad.bit() & 1u); p.a1 = static_cast(a ^ p.a0); p.b0_share = pad.block(); p.b1_share = ds_xor(B, p.b0_share); p.c0_share = pad.block(); p.c1_share = ds_xor(C, p.c0_share); return p; } HEDLEY_ALWAYS_INLINE simde__m128i ds_cw_share(simde__m128i L, simde__m128i R, uint8_t my_bit, const ds_cw_party & mine, const ds_blind & their) noexcept { simde__m128i out = ds_xor(R, mine.gamma); if (my_bit & 1u) { out = ds_xor(out, ds_xor(ds_xor(L, R), their.msg)); } if (their.bit & 1u) { out = ds_xor(out, mine.rand); } return out; } inline void ds_cw_blinds(const ds_cw_pads & p, simde__m128i L0, simde__m128i R0, uint8_t bit0, simde__m128i L1, simde__m128i R1, uint8_t bit1, ds_blind & b0, ds_blind & b1) noexcept { b0.bit = static_cast(bit0 ^ p.p0.bit); b1.bit = static_cast(bit1 ^ p.p1.bit); b0.msg = ds_xor(ds_xor(L0, R0), p.p0.rand); b1.msg = ds_xor(ds_xor(L1, R1), p.p1.rand); } inline simde__m128i ds_cw_outs(const ds_cw_pads & p, simde__m128i L0, simde__m128i R0, uint8_t bit0, simde__m128i L1, simde__m128i R1, uint8_t bit1, const ds_blind & b0, const ds_blind & b1) noexcept { return ds_xor( ds_cw_share(L0, R0, bit0, p.p0, b1), ds_cw_share(L1, R1, bit1, p.p1, b0)); } inline uint8_t ds_open_advice(simde__m128i L0, simde__m128i R0, uint8_t bit0, simde__m128i L1, simde__m128i R1, uint8_t bit1) noexcept { const uint8_t a00 = static_cast(dpf::get_lo_bit(L0) ^ bit0); const uint8_t a01 = static_cast(dpf::get_lo_bit(R0) ^ bit0); const uint8_t a10 = static_cast(dpf::get_lo_bit(L1) ^ bit1); const uint8_t a11 = static_cast(dpf::get_lo_bit(R1) ^ bit1); const uint8_t t0 = static_cast(a00 ^ a10 ^ 1u); const uint8_t t1 = static_cast(a01 ^ a11); return static_cast((t1 << 1) | (t0 & 1u)); } inline void ds_next_terms(simde__m128i L, simde__m128i R, uint8_t advice, simde__m128i cw, uint8_t tpack, simde__m128i & M, simde__m128i & base) noexcept { const simde__m128i D = ds_xor(L, R); const uint8_t t0 = static_cast(tpack & 1u); const uint8_t t1 = static_cast((tpack >> 1) & 1u); const simde__m128i lo = dpf::set_lo_bit(simde_mm_setzero_si128(), 1); const simde__m128i DT = ds_gate(static_cast(t0 ^ t1), lo); const simde__m128i cw_base = ds_xor(dpf::unset_lo_bit(cw), ds_gate(t0, lo)); M = (advice & 1u) ? ds_xor(D, DT) : D; base = (advice & 1u) ? ds_xor(L, cw_base) : L; } inline ds_and_shares ds_and_open(const ds_and_pads & p, simde__m128i M, uint8_t b_recv) noexcept { const simde__m128i e = ds_xor(ds_xor(M, p.b0_share), p.b1_share); const uint8_t d = static_cast((b_recv ^ p.a1) ^ p.a0); ds_and_shares z; z.z0 = ds_xor(ds_xor(ds_xor(ds_gate(d, e), ds_gate(d, p.b0_share)), ds_gate(p.a0, e)), p.c0_share); z.z1 = ds_xor(ds_xor(ds_gate(d, p.b1_share), ds_gate(p.a1, e)), p.c1_share); return z; } inline simde__m128i ds_deliver(uint8_t b_exp, simde__m128i base, simde__m128i M, const ds_and_shares & z) noexcept { const simde__m128i local = ds_xor(base, ds_gate(b_exp, M)); return ds_xor(ds_xor(local, z.z0), z.z1); } /// Per-level messages prepared before the CW protocol runs (blinds + pads). struct ds_level_blinds { ds_cw_pads cwp; ds_and_pads and0; ds_and_pads and1; ds_blind b0; ds_blind b1; simde__m128i L0, R0, L1, R1; uint8_t bit0; uint8_t bit1; }; /// Opened CW, advice, and AND products delivered by a `CwProtocol`. struct ds_level_open { simde__m128i cw; uint8_t advice; ds_and_shares z0; ds_and_shares z1; uint64_t value_cw = 0; // public after open when cmp is active at this level }; /// Running comparison-gen state shared across DS levels (Va residual). /// When `track_coeff` is set (wildcard cmp payload), a parallel β = 1 /// accumulator `Va1` is advanced alongside `Va` so the gen can stash /// `value_cw(1) − value_cw(β)` coefficients for a later `assign_cmp`. struct ds_cmp_gen_state { bool active = false; std::size_t nbits = 0; uint64_t mask = 0; uint64_t beta = 0; bool include_eq = false; cmp_trivial trivial = cmp_trivial::none; uint64_t Va = 0; unsigned __int128 thresh = 0; bool track_coeff = false; uint64_t Va1 = 0; uint64_t last_vcw_coeff = 0; }; /// Local joint simulation: today's `ds_cw_outs` / `ds_open_advice` / `ds_and_open`. /// An MPC backend would send `blinds` and return the same `ds_level_open` shape. template struct local_cw_protocol { PadRng & pads; ds_level_blinds prepare_level(simde__m128i L0, simde__m128i R0, uint8_t bit0, simde__m128i L1, simde__m128i R1, uint8_t bit1) { ds_level_blinds b; b.cwp = ds_sample_cw(pads); b.and0 = ds_sample_and(pads); b.and1 = ds_sample_and(pads); b.L0 = L0; b.R0 = R0; b.L1 = L1; b.R1 = R1; b.bit0 = bit0; b.bit1 = bit1; ds_cw_blinds(b.cwp, L0, R0, bit0, L1, R1, bit1, b.b0, b.b1); return b; } ds_level_open complete_level(std::size_t /*level*/, const ds_level_blinds & b, simde__m128i M0, simde__m128i M1, uint8_t rec0, uint8_t rec1) { ds_level_open out; out.advice = ds_open_advice(b.L0, b.R0, b.bit0, b.L1, b.R1, b.bit1); out.cw = ds_cw_outs(b.cwp, b.L0, b.R0, b.bit0, b.L1, b.R1, b.bit1, b.b0, b.b1); out.z0 = ds_and_open(b.and0, M0, rec0); out.z1 = ds_and_open(b.and1, M1, rec1); return out; } /// Open CW + advice only (AND pads stay in `blinds` for a later open). std::pair open_cw(const ds_level_blinds & b) noexcept { return {ds_cw_outs(b.cwp, b.L0, b.R0, b.bit0, b.L1, b.R1, b.bit1, b.b0, b.b1), ds_open_advice(b.L0, b.R0, b.bit0, b.L1, b.R1, b.bit1)}; } /// Open the public value CW for this level (local: clear convert+make_value_cw). /// MPC backends open additive shares of the same word. uint64_t open_value_cw(const ds_level_blinds & b, uint8_t adv0, uint8_t adv1, int ai, uint64_t & Va, uint64_t beta, uint64_t mask) noexcept { return dcf_impl::make_value_cw(b.L0, b.R0, b.L1, b.R1, adv0, adv1, ai, Va, beta, mask); } ds_and_shares open_and(const ds_and_pads & p, simde__m128i M, uint8_t b_recv) noexcept { return ds_and_open(p, M, b_recv); } /// Open the final comparison leaf CW. Wraps `make_final_cw` so the /// Doerner–Shelat gen does not call it directly on reconstructed seeds; /// an MPC backend would open additive shares of the same word. uint64_t open_final_cw(simde__m128i s0, simde__m128i s1, uint8_t t1, uint64_t Va, uint64_t mask, uint64_t on_path) noexcept { return dcf_impl::make_final_cw(s0, s1, t1, Va, mask, on_path); } /// Draw the group-width `cmp_addend` blind. Local joint simulation reuses /// the shared root sampler so the blind matches the dealer's; an MPC /// backend would instead pull a group-width element from the pad stream. template uint64_t sample_addend_blind(uint64_t mask, BlockSampler && sample) noexcept { return dcf_impl::sample_addend_blind(mask, std::forward(sample)); } /// Open a group of leaf correction words for one prefix group. In this /// local joint simulation both XOR shares of the point are present, so the /// point is reconstructed *inside* the protocol and handed to `leaf_fn` /// (which runs `make_leaves` for the group). The Doerner–Shelat gen never /// forms `x = x0 ^ x1` at its own call site; an MPC backend would instead /// run a per-group leaf CW exchange that never reveals `x`. template void open_leaf_group(InputT x0, InputT x1, LeafFn && leaf_fn) { std::forward(leaf_fn)(utils::xor_input_shares(x0, x1)); } }; /// Generation-side level state (seeds / home bits). Not an eval path memoizer. template struct ds_gen_state { NodeT inbox[2]; int home[2]; NodeT root0; NodeT root1; void init(NodeT r0, NodeT r1) noexcept { root0 = r0; root1 = r1; inbox[0] = r0; inbox[1] = r1; home[0] = 0; home[1] = 1; } NodeT & seed0() noexcept { return inbox[home[0]]; } NodeT & seed1() noexcept { return inbox[home[1]]; } const NodeT & seed0() const noexcept { return inbox[home[0]]; } const NodeT & seed1() const noexcept { return inbox[home[1]]; } }; /// One interior level: expand, protocol open, advance both party seeds. /// When `cmp` is non-null and active for `level`, also opens `value_cw` via /// the protocol (no second PRG expand outside). template void ds_advance_level(ds_gen_state & st, InputT x0, InputT x1, InputT mask, std::size_t level, CwProtocol & proto, NodeT & cw_out, AdviceT & advice_out, uint64_t * value_cw_out = nullptr, ds_cmp_gen_state * cmp = nullptr) { // Integral bridge so bit extraction works for `keyword` / `modint` / // signed / bitstring the same way dealer gen does via `mask & x`. constexpr auto to_int = utils::to_integral_type{}; const auto mi = to_int(mask); const uint8_t bit0 = static_cast(!!(mi & to_int(x0))); const uint8_t bit1 = static_cast(!!(mi & to_int(x1))); NodeT s0 = st.seed0(); NodeT s1 = st.seed1(); const uint8_t adv0 = static_cast( dpf::get_lo_bit_and_clear_lo_2bits(s0)); const uint8_t adv1 = static_cast( dpf::get_lo_bit_and_clear_lo_2bits(s1)); const auto c0 = InteriorPRG::eval01(s0); const auto c1 = InteriorPRG::eval01(s1); auto blinds = proto.prepare_level(c0[0], c0[1], bit0, c1[0], c1[1], bit1); if (value_cw_out != nullptr && cmp != nullptr && cmp->active && cmp->trivial == cmp_trivial::none && level < cmp->nbits) { const int ai = static_cast( (cmp->thresh >> (cmp->nbits - 1 - level)) & 1); *value_cw_out = proto.open_value_cw(blinds, adv0, adv1, ai, cmp->Va, cmp->beta, cmp->mask); if (cmp->track_coeff) { // Affine coefficient: same level with β = 1 on a parallel Va. const uint64_t v1 = proto.open_value_cw(blinds, adv0, adv1, ai, cmp->Va1, 1ULL, cmp->mask); cmp->last_vcw_coeff = (v1 + dcf_impl::neg_m(*value_cw_out, cmp->mask)) & cmp->mask; } } auto [cw, tpack] = proto.open_cw(blinds); const uint8_t exp0 = st.home[0] == 0 ? bit0 : bit1; const uint8_t rec0 = st.home[0] == 0 ? bit1 : bit0; const uint8_t exp1 = st.home[1] == 0 ? bit0 : bit1; const uint8_t rec1 = st.home[1] == 0 ? bit1 : bit0; NodeT M0, base0, M1, base1; ds_next_terms(c0[0], c0[1], adv0, cw, tpack, M0, base0); ds_next_terms(c1[0], c1[1], adv1, cw, tpack, M1, base1); const NodeT nxt0 = ds_deliver(exp0, base0, M0, proto.open_and(blinds.and0, M0, rec0)); const NodeT nxt1 = ds_deliver(exp1, base1, M1, proto.open_and(blinds.and1, M1, rec1)); st.home[0] ^= 1; st.home[1] ^= 1; st.inbox[st.home[0]] = nxt0; st.inbox[st.home[1]] = nxt1; cw_out = cw; advice_out = tpack; } template struct is_ds_randomness : std::false_type {}; template struct is_ds_randomness> : std::true_type {}; template struct is_cw_protocol : std::false_type {}; template struct is_cw_protocol, void> : std::true_type {}; template struct is_cw_protocol().prepare_level( simde_mm_setzero_si128(), simde_mm_setzero_si128(), uint8_t{}, simde_mm_setzero_si128(), simde_mm_setzero_si128(), uint8_t{})), decltype(std::declval().open_cw( std::declval()))>> : std::true_type {}; template struct first_is_cw_protocol : std::false_type {}; template struct first_is_cw_protocol : is_cw_protocol> {}; template auto make_dpf_doerner_shelat_impl(InputT x0, InputT x1, RootSampler & root_sampler, CwProtocol & proto, OutputT && y, OutputTs && ...ys) { static_assert(!dpf::is_wildcard_v, "Doerner–Shelat gen takes XOR shares of a concrete point"); static_assert(!dpf::is_secret_share_v, "Doerner–Shelat: pass additive_share of xor_wrapper, or raw XOR shares"); static_assert(sizeof(typename InteriorPRG::block_type) == sizeof(simde__m128i), "Doerner–Shelat gen uses the AES-block interior node"); using dpf_type = utils::dpf_type_t; using node = typename dpf_type::interior_node; using input_type = typename dpf_type::input_type; constexpr auto depth = dpf_type::depth; utils::flip_msb_if_signed_integral(x0); const node root0 = dpf::unset_lo_bit(static_cast(root_sampler())); const node root1 = dpf::set_lo_bit(static_cast(root_sampler())); ds_gen_state st; st.init(root0, root1); typename dpf_type::correction_words_array correction_words{}; typename dpf_type::correction_advice_array correction_advice{}; auto mask = dpf_type::msb_mask; for (std::size_t level = 0; level < depth; ++level, mask >>= 1) { ds_advance_level(st, x0, x1, mask, level, proto, correction_words[level], correction_advice[level]); } const node parent0 = st.seed0(); const node parent1 = st.seed1(); const bool sign0 = dpf::get_lo_bit(parent0); input_type x = utils::xor_input_shares(x0, x1); auto built = dpf::make_leaves(x, dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), sign0, std::size_t{0}, std::forward(y), std::forward(ys)...); input_type off0{}; input_type off1{}; return dpf::make_party_key_pair( dpf_type{root0, correction_words, correction_advice, built.first.first, built.first.second, off0}, dpf_type{root1, correction_words, correction_advice, built.second.first, built.second.second, off1}); } } // namespace detail /// Local CW protocol (pads cancel; same keys as dealer when roots match). template using local_cw_protocol = detail::local_cw_protocol; template using ds_gen_state = detail::ds_gen_state; // Public `make_dpf_doerner_shelat(x0, x1, ...)` lives in incremental.hpp so // classic and `at<>` / mixed-width packs share one entry point. } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_DOERNER_SHELAT_HPP__