libdpf/include/dpf/doerner_shelat.hpp
Ryan Henry e4e666f459 Initial import of libdpf.
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-24 14:08:32 -06:00

542 lines
18 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/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 <cstddef>
#include <cstdint>
#include <type_traits>
#include <utility>
#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 <typename RootSampler, typename PadRng>
struct ds_randomness
{
RootSampler root;
PadRng pad;
};
namespace detail
{
struct urandom_pad_rng
{
simde__m128i block()
{
return dpf::uniform_sample<simde__m128i>();
}
uint8_t bit()
{
return static_cast<uint8_t>(dpf::uniform_sample<unsigned char>() & 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 <typename PadRng>
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<uint8_t>(pad.bit() & 1u);
p.p1.bit = static_cast<uint8_t>(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 <typename PadRng>
ds_and_pads ds_sample_and(PadRng & pad)
{
ds_and_pads p{};
const uint8_t a = static_cast<uint8_t>(pad.bit() & 1u);
const simde__m128i B = pad.block();
const simde__m128i C = ds_gate(a, B);
p.a0 = static_cast<uint8_t>(pad.bit() & 1u);
p.a1 = static_cast<uint8_t>(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<uint8_t>(bit0 ^ p.p0.bit);
b1.bit = static_cast<uint8_t>(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<uint8_t>(dpf::get_lo_bit(L0) ^ bit0);
const uint8_t a01 = static_cast<uint8_t>(dpf::get_lo_bit(R0) ^ bit0);
const uint8_t a10 = static_cast<uint8_t>(dpf::get_lo_bit(L1) ^ bit1);
const uint8_t a11 = static_cast<uint8_t>(dpf::get_lo_bit(R1) ^ bit1);
const uint8_t t0 = static_cast<uint8_t>(a00 ^ a10 ^ 1u);
const uint8_t t1 = static_cast<uint8_t>(a01 ^ a11);
return static_cast<uint8_t>((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<uint8_t>(tpack & 1u);
const uint8_t t1 = static_cast<uint8_t>((tpack >> 1) & 1u);
const simde__m128i lo = dpf::set_lo_bit(simde_mm_setzero_si128(), 1);
const simde__m128i DT = ds_gate(static_cast<uint8_t>(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<uint8_t>((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 <typename PadRng>
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<simde__m128i, uint8_t> 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 <typename BlockSampler>
uint64_t sample_addend_blind(uint64_t mask, BlockSampler && sample) noexcept
{
return dcf_impl::sample_addend_blind(mask,
std::forward<BlockSampler>(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 <typename InputT, typename LeafFn>
void open_leaf_group(InputT x0, InputT x1, LeafFn && leaf_fn)
{
std::forward<LeafFn>(leaf_fn)(utils::xor_input_shares(x0, x1));
}
};
/// Generation-side level state (seeds / home bits). Not an eval path memoizer.
template <typename NodeT>
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 <typename InteriorPRG, typename CwProtocol, typename NodeT,
typename InputT, typename AdviceT>
void ds_advance_level(ds_gen_state<NodeT> & 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<InputT>{};
const auto mi = to_int(mask);
const uint8_t bit0 = static_cast<uint8_t>(!!(mi & to_int(x0)));
const uint8_t bit1 = static_cast<uint8_t>(!!(mi & to_int(x1)));
NodeT s0 = st.seed0();
NodeT s1 = st.seed1();
const uint8_t adv0 = static_cast<uint8_t>(
dpf::get_lo_bit_and_clear_lo_2bits(s0));
const uint8_t adv1 = static_cast<uint8_t>(
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<int>(
(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 <typename T>
struct is_ds_randomness : std::false_type {};
template <typename RootSampler, typename PadRng>
struct is_ds_randomness<ds_randomness<RootSampler, PadRng>> : std::true_type {};
template <typename T, typename = void>
struct is_cw_protocol : std::false_type {};
template <typename PadRng>
struct is_cw_protocol<local_cw_protocol<PadRng>, void> : std::true_type {};
template <typename T>
struct is_cw_protocol<T,
std::void_t<decltype(std::declval<T &>().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<T &>().open_cw(
std::declval<const ds_level_blinds &>()))>>
: std::true_type {};
template <typename ...Ts>
struct first_is_cw_protocol : std::false_type {};
template <typename T, typename ...Rest>
struct first_is_cw_protocol<T, Rest...>
: is_cw_protocol<std::decay_t<T>> {};
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename RootSampler,
typename CwProtocol>
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<InputT>,
"Doerner–Shelat gen takes XOR shares of a concrete point");
static_assert(!dpf::is_secret_share_v<InputT>,
"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<InteriorPRG, ExteriorPRG, InputT,
OutputT, OutputTs...>;
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<node>(root_sampler()));
const node root1 = dpf::set_lo_bit(static_cast<node>(root_sampler()));
ds_gen_state<node> 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<InteriorPRG>(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<ExteriorPRG>(x,
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), sign0,
std::size_t{0}, std::forward<OutputT>(y), std::forward<OutputTs>(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 <typename PadRng>
using local_cw_protocol = detail::local_cw_protocol<PadRng>;
template <typename NodeT>
using ds_gen_state = detail::ds_gen_state<NodeT>;
// 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__