libdpf/include/dpf/doerner_shelat.hpp

1059 lines
37 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 shares of the point are walked level by level — XOR shares by
/// default, or additive shares when tagged with `arith_input`.
/// Correction words, advice bits, seeds, and leaves are the ones
/// `make_dpf` would emit for the reconstructed point, 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"
#include "dpf/constrained_cmp.hpp"
namespace dpf
{
/// @brief Tag: Doerner–Shelat / geneval takes additive shares of the point
/// (`x0 + x1` in the input ring). Default calls take XOR shares.
struct arith_input_t
{
};
inline constexpr arith_input_t arith_input{};
/// @brief Tag: payload β is additively shared (`y0 + y1`). Leaf CW is opened via
/// Π_CCMP on the on-path control bits (see `open_arith_leaf`).
/// @see `open_arith_leaf`
struct arith_output_t
{
};
inline constexpr arith_output_t arith_output{};
/// @brief Additive (or XOR) shares of one concrete payload for dealerless leaf open.
/// @details Use as a placed value / `at<>` element when several outputs are shared.
/// @tparam T value type
template <typename T>
struct arith_beta
{
using payload_type = T;
T y0{};
T y1{};
};
template <typename T>
struct is_arith_beta : std::false_type
{
};
template <typename T>
struct is_arith_beta<arith_beta<T>> : std::true_type
{
};
template <typename T>
inline constexpr bool is_arith_beta_v = is_arith_beta<std::decay_t<T>>::value;
namespace detail
{
namespace incr
{
/// @brief `placed<N, arith_beta<T>>::output_type` is `T` (see placement.hpp).
/// @tparam T value type
template <typename T>
struct unwrap_placed_output<arith_beta<T>>
{
using type = T;
};
} // namespace incr
} // namespace detail
/// @brief Roots and the Beaver-pad stream for one Doerner–Shelat generation.
/// @details `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.
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
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_NO_THROW
HEDLEY_ALWAYS_INLINE
simde__m128i ds_xor(simde__m128i a, simde__m128i b) noexcept
{
return simde_mm_xor_si128(a, b);
}
HEDLEY_NO_THROW
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_NO_THROW
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;
}
HEDLEY_NO_THROW
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);
}
HEDLEY_NO_THROW
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));
}
HEDLEY_NO_THROW
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));
}
HEDLEY_NO_THROW
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;
}
HEDLEY_NO_THROW
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;
}
HEDLEY_NO_THROW
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);
}
/// @brief 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;
};
/// @brief 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
};
/// @brief Running comparison-gen state shared across DS levels (Va residual).
/// @details 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;
cmp_kind kind = cmp_kind::lt;
bool paint = false;
std::size_t length_bits = 0;
paint_callback paint_cb = nullptr;
const void * paint_ctx = nullptr;
};
/// @brief Local joint simulation: today's `ds_cw_outs` / `ds_open_advice` / `ds_and_open`.
/// @details An MPC backend would send `blinds` and return the same `ds_level_open` shape.
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
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;
}
/// @brief Open CW + advice only (AND pads stay in `blinds` for a later open).
/// @param b the `b`
/// @return the returned `std::pair<simde__m128i, uint8_t>`
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
HEDLEY_NO_THROW
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)};
}
HEDLEY_PRAGMA(GCC diagnostic pop)
/// @brief Open the public value CW for this level (local: clear convert+make_value_cw).
/// @details MPC backends open additive shares of the same word.
/// @param b the `b`
/// @param adv0 the `adv0`
/// @param adv1 the `adv1`
/// @param ai the `ai`
/// @param Va the `Va`
/// @param beta the payload
/// @param mask the bit mask
/// @return the returned `uint64_t`
HEDLEY_NO_THROW
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);
}
/// @brief Open a path-paint value CW. `plant` is the scaled lose-subtree constant.
/// @param b the `b`
/// @param adv0 the `adv0`
/// @param adv1 the `adv1`
/// @param ai the `ai`
/// @param Va the `Va`
/// @param plant the unit plant
/// @param mask the bit mask
/// @return the returned `uint64_t`
HEDLEY_NO_THROW
uint64_t open_planted_cw(const ds_level_blinds & b, uint8_t adv0, uint8_t adv1,
int ai, uint64_t & Va, uint64_t plant, uint64_t mask) noexcept
{
return dcf_impl::make_value_cw_planted(b.L0, b.R0, b.L1, b.R1, adv0,
adv1, ai, Va, plant, mask);
}
HEDLEY_NO_THROW
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);
}
/// @brief 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.
/// @param s0 the `s0`
/// @param s1 the `s1`
/// @param t1 the `t1`
/// @param Va the `Va`
/// @param mask the bit mask
/// @param on_path the value reconstructed on the secret path
/// @return the returned `uint64_t`
HEDLEY_NO_THROW
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);
}
/// @brief 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.
/// @tparam BlockSampler block sampler
/// @param mask the bit mask
/// @param sample the `sample`
/// @return the returned `uint64_t`
template <typename BlockSampler>
HEDLEY_NO_THROW
uint64_t sample_addend_blind(uint64_t mask, BlockSampler && sample) noexcept
{
return dcf_impl::sample_addend_blind(mask,
std::forward<BlockSampler>(sample));
}
/// @brief Majority of three bits (next carry of a full adder).
/// @param a the `a`
/// @param b the `b`
/// @param c the `c`
/// @return Majority of three bits (next carry of a full adder)
HEDLEY_NO_THROW
static constexpr uint8_t majority(uint8_t a, uint8_t b, uint8_t c) noexcept
{
return static_cast<uint8_t>((a & b) | (a & c) | (b & c));
}
/// @brief One additive digit: sum bit `a XOR b XOR cin`, carry out = majority.
/// @param a the `a`
/// @param b the `b`
/// @param cin the `cin`
/// @param cout the `cout`
/// @return One additive digit: sum bit `a XOR b XOR cin`, carry out = majority
HEDLEY_NO_THROW
static constexpr uint8_t open_sum_bit(uint8_t a, uint8_t b, uint8_t cin,
uint8_t & cout) noexcept
{
cout = majority(a, b, cin);
return static_cast<uint8_t>(a ^ b ^ cin);
}
/// @brief Reconstruct `a0 + a1` via a LSB→MSB carry chain, then flip the MSB
/// when the domain is signed — matching `make_dpf` on the sum. The call
/// site never forms the sum; an MPC backend would open the same bits.
/// @tparam InputT input domain type
/// @param a0 the `a0`
/// @param a1 the `a1`
/// @return Reconstruct `a0 + a1` via a LSB→MSB carry chain, then flip the MSB when the domain
/// is signed — matching `make_dpf` on the sum
template <typename InputT>
InputT open_arith_point(InputT a0, InputT a1) const
{
constexpr auto to_int = utils::to_integral_type<InputT>{};
using FromI = typename utils::make_from_integral_value<InputT>::integral_type;
using U = std::make_unsigned_t<FromI>;
const U u0 = static_cast<U>(to_int(a0));
const U u1 = static_cast<U>(to_int(a1));
U sum = 0;
uint8_t carry = 0;
constexpr std::size_t nbits = utils::bitlength_of_v<InputT>;
for (std::size_t i = 0; i < nbits; ++i)
{
const uint8_t b0 = static_cast<uint8_t>((u0 >> i) & U{1});
const uint8_t b1 = static_cast<uint8_t>((u1 >> i) & U{1});
const uint8_t s = open_sum_bit(b0, b1, carry, carry);
sum = static_cast<U>(sum | (static_cast<U>(s) << i));
}
InputT out = utils::make_from_integral_value<InputT>{}(
static_cast<FromI>(sum));
utils::flip_msb_if_signed_integral(out);
return out;
}
/// @brief Encode shares for the XOR-style CW walk. XOR mode flips party 0's MSB
/// (linear over XOR). Arithmetic mode opens the sum (carry + signed MSB)
/// and returns `(alpha, 0)` so the walk matches `make_dpf(alpha)`.
/// @tparam InputT input domain type
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param arith the `arith`
template <typename InputT>
void encode_walk_shares(InputT & x0, InputT & x1, bool arith) const
{
if (arith)
{
const InputT alpha = open_arith_point(x0, x1);
x0 = alpha;
x1 = InputT{};
}
else
{
utils::flip_msb_if_signed_integral(x0);
}
}
/// @brief Constrained comparison Π_CCMP: open `1{x0 < x1}` when `|x0−x1|=1`.
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @return Constrained comparison Π_CCMP: open `1{x0 < x1}` when `|x0−x1|=1`
HEDLEY_NO_THROW
uint8_t open_ccmp(std::uint64_t x0, std::uint64_t x1) noexcept
{
(void)pads; // MPC backend would consume an AND pad here
return dpf::local_ccmp(x0, x1);
}
/// @brief Open a public leaf CW for a shared payload.
/// @details Ring: `β = y0 + y1`; `g = CCMP(t0,t1)` selects `β − M` vs `M − β`
/// (matches `make_leaf` with `sign = t0`). Characteristic 2: `β = y0 ⊕ y1`
/// and CW = `β ⊕ M` (sign mux is a no-op under XOR).
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam I output index
/// @tparam OutputsTuple outputs tuple
/// @tparam InteriorBlock interior block
/// @tparam OutputT output type
/// @param seed0 the `seed0`
/// @param seed1 the `seed1`
/// @param t0 the `t0`
/// @param t1 the `t1`
/// @param y0 the `y0`
/// @param y1 the `y1`
/// @param pos_base the `pos_base`
/// @param lane_x lane of the shared payload
/// @return the opened leaf correction word
template <typename ExteriorPRG, std::size_t I = 0, typename OutputsTuple,
typename InteriorBlock, typename OutputT>
auto open_arith_leaf(const InteriorBlock & seed0, const InteriorBlock & seed1,
uint8_t t0, uint8_t t1, OutputT y0, OutputT y1, std::size_t pos_base,
std::size_t lane_x) -> dpf::leaf_node_t<typename ExteriorPRG::block_type,
OutputT>
{
using output_type = OutputT;
using node_type = typename ExteriorPRG::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using leaf_type = dpf::leaf_node_t<node_type, output_type>;
HEDLEY_PRAGMA(GCC diagnostic pop)
const auto M = dpf::make_leaf_mask<ExteriorPRG, I, OutputsTuple, InteriorBlock>(
seed0, seed1, pos_base);
output_type beta{};
if constexpr (utils::has_characteristic_two_v<output_type>)
{
(void)t0;
(void)t1;
beta = static_cast<output_type>(y0 ^ y1);
const leaf_type naked = dpf::make_naked_leaf<node_type>(lane_x, beta);
return dpf::subtract_leaf<output_type>(naked, M);
}
else
{
const uint8_t g = open_ccmp(t0, t1);
beta = static_cast<output_type>(y0 + y1);
const leaf_type naked = dpf::make_naked_leaf<node_type>(lane_x, beta);
// CW = (−1)^{t1}(β − M): g=0 → β−M; g=1 → M−β. Matches make_leaf(sign=t0).
if (g & 1u)
return dpf::subtract_leaf<output_type>(M, naked);
return dpf::subtract_leaf<output_type>(naked, M);
}
}
/// @brief 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`. After
/// `encode_walk_shares`, arithmetic inputs are already `(alpha, 0)`.
/// @tparam InputT input domain type
/// @tparam LeafFn leaf fn
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param leaf_fn the `leaf_fn`
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));
}
};
/// @brief Generation-side level state (seeds / home bits). Not an eval path memoizer.
/// @tparam NodeT GGM node type
template <typename NodeT>
struct ds_gen_state
{
NodeT inbox[2];
int home[2];
NodeT root0;
NodeT root1;
HEDLEY_NO_THROW
void init(NodeT r0, NodeT r1) noexcept
{
root0 = r0;
root1 = r1;
inbox[0] = r0;
inbox[1] = r1;
home[0] = 0;
home[1] = 1;
}
HEDLEY_NO_THROW
NodeT & seed0() noexcept { return inbox[home[0]]; }
HEDLEY_NO_THROW
NodeT & seed1() noexcept { return inbox[home[1]]; }
HEDLEY_NO_THROW
const NodeT & seed0() const noexcept { return inbox[home[0]]; }
HEDLEY_NO_THROW
const NodeT & seed1() const noexcept { return inbox[home[1]]; }
};
/// @brief One interior level: expand, protocol open, advance both party seeds.
/// @details When `cmp` is non-null and active for `level`, also opens `value_cw` via
/// the protocol (no second PRG expand outside).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam CwProtocol correction-word protocol
/// @tparam NodeT GGM node type
/// @tparam InputT input domain type
/// @tparam MaskT mask type
/// @tparam AdviceT advice type
/// @param st the `st`
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param mask the bit mask
/// @param level the tree level
/// @param depth the tree depth
/// @param proto the `proto`
/// @param cw_out the `cw_out`
/// @param advice_out the `advice_out`
/// @param value_cw_out the `value_cw_out`
/// @param cmp the comparison specification
template <typename InteriorPRG, typename CwProtocol, typename NodeT,
typename InputT, typename MaskT, typename AdviceT>
void ds_advance_level(ds_gen_state<NodeT> & st, InputT x0, InputT x1,
MaskT mask, std::size_t level, std::size_t depth, CwProtocol & proto,
NodeT & cw_out, AdviceT & advice_out, uint64_t * value_cw_out = nullptr,
ds_cmp_gen_state * cmp = nullptr)
{
using tree = dpf::tree_traits<InteriorPRG>;
// Integral bridge so bit extraction works for `keyword` / `modint` /
// signed / bitstring the same way dealer gen does via `mask & x`.
// `msb_mask` is the unsigned bit pattern; a signed input must not be
// required to have that same type.
constexpr auto to_int = utils::to_integral_type<InputT>{};
constexpr auto to_mask = utils::to_integral_type<MaskT>{};
const auto mi = to_mask(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)));
const bool is_last = tree::is_last_level(level, depth);
NodeT s0 = st.seed0();
NodeT s1 = st.seed1();
const uint8_t adv0 = static_cast<uint8_t>(dpf::get_lo_bit(s0));
const uint8_t adv1 = static_cast<uint8_t>(dpf::get_lo_bit(s1));
const auto c0 = tree::expand(s0, is_last);
const auto c1 = tree::expand(s1, is_last);
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)
{
// Convert uses expand_value (HT: always two-tweak); seed walk used expand.
const auto v0 = tree::expand_value(s0);
const auto v1 = tree::expand_value(s1);
auto vblinds = blinds;
vblinds.L0 = v0[0];
vblinds.R0 = v0[1];
vblinds.L1 = v1[0];
vblinds.R1 = v1[1];
const int ai = static_cast<int>(
(cmp->thresh >> (cmp->nbits - 1 - level)) & 1);
if (cmp->paint)
{
const uint64_t unit = dcf_impl::paint_unit(cmp->kind, level,
cmp->thresh, cmp->nbits, cmp->length_bits, false,
cmp->paint_cb, cmp->paint_ctx);
const uint64_t plant = dcf_impl::scale_plant(unit, cmp->beta,
cmp->mask);
*value_cw_out = proto.open_planted_cw(vblinds, adv0, adv1, ai,
cmp->Va, plant, cmp->mask);
if (cmp->track_coeff)
{
const uint64_t plant1 = dcf_impl::scale_plant(unit, 1ULL,
cmp->mask);
const uint64_t v1w = proto.open_planted_cw(vblinds, adv0, adv1,
ai, cmp->Va1, plant1, cmp->mask);
cmp->last_vcw_coeff =
(v1w + dcf_impl::neg_m(*value_cw_out, cmp->mask)) & cmp->mask;
}
}
else
{
*value_cw_out = proto.open_value_cw(vblinds, 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 v1w = proto.open_value_cw(vblinds, adv0, adv1, ai,
cmp->Va1, 1ULL, cmp->mask);
cmp->last_vcw_coeff =
(v1w + dcf_impl::neg_m(*value_cw_out, cmp->mask)) & cmp->mask;
}
}
}
auto [cw, tpack] = proto.open_cw(blinds);
// Half-Tree mid levels store no advice; last level keeps BGI packing.
if constexpr (tree::is_half_tree)
{
if (!is_last)
tpack = 0;
}
else
{
// BGI: opened advice stands.
}
// Dealer-equivalent CW for Half-Tree mid: off-path children XOR already
// matches H(s0)⊕H(s1)⊕ᾱΔ via the open. For last/BGI, open matches Gen.
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;
if constexpr (tree::is_half_tree)
{
if (!is_last)
{
// Mid: next = child[bit] ⊕ (t ? full_cw : 0).
const NodeT D0 = ds_xor(c0[0], c0[1]);
const NodeT D1 = ds_xor(c1[0], c1[1]);
M0 = D0;
base0 = (adv0 & 1u) ? ds_xor(c0[0], cw) : c0[0];
M1 = D1;
base1 = (adv1 & 1u) ? ds_xor(c1[0], cw) : c1[0];
}
else
{
ds_next_terms(c0[0], c0[1], adv0, cw, tpack, M0, base0);
ds_next_terms(c1[0], c1[1], adv1, cw, tpack, M1, base1);
}
}
else
{
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 RootSampler,
typename CwProtocol>
auto make_dpf_doerner_shelat_impl(bool arith, bool arith_out, InputT x0,
InputT x1, RootSampler & root_sampler, CwProtocol & proto, OutputT y0,
OutputT y1 = OutputT{})
{
static_assert(!dpf::is_wildcard_v<InputT>,
"Doerner–Shelat gen takes shares of a concrete point");
static_assert(!dpf::is_secret_share_v<InputT>,
"Doerner–Shelat: pass additive_share of xor_wrapper, or raw shares");
static_assert(!dpf::is_wildcard_v<OutputT>,
"arith_output / classic DS leaf expects a concrete payload");
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>;
using node = typename dpf_type::interior_node;
using input_type = typename dpf_type::input_type;
using leaf_tuple = typename dpf_type::leaf_tuple;
using beaver_tuple = typename dpf_type::beaver_tuple;
using outputs_tuple = std::tuple<OutputT>;
constexpr auto depth = dpf_type::depth;
proto.encode_walk_shares(x0, x1, arith);
using tree = dpf::tree_traits<InteriorPRG>;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
node roots[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(roots, [&]() -> node {
return static_cast<node>(root_sampler());
});
const node root0 = roots[0];
const node root1 = roots[1];
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_gen_state<node> st;
HEDLEY_PRAGMA(GCC diagnostic pop)
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, depth, 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);
const uint8_t t0 = static_cast<uint8_t>(sign0);
const uint8_t t1 = static_cast<uint8_t>(dpf::get_lo_bit(parent1));
input_type x = utils::xor_input_shares(x0, x1);
leaf_tuple leaves0{};
leaf_tuple leaves1{};
beaver_tuple beavers0{};
beaver_tuple beavers1{};
if (arith_out)
{
constexpr auto to_int = utils::to_integral_type<input_type>{};
const std::size_t lane = static_cast<std::size_t>(to_int(x));
auto cw = proto.template open_arith_leaf<ExteriorPRG, 0, outputs_tuple>(
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), t0, t1,
y0, y1, std::size_t{0}, lane);
std::get<0>(leaves0) = cw;
std::get<0>(leaves1) = cw;
}
else
{
auto built = dpf::make_leaves<ExteriorPRG>(x,
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), sign0,
std::size_t{0}, y0);
leaves0 = std::move(built.first.first);
beavers0 = std::move(built.first.second);
leaves1 = std::move(built.second.first);
beavers1 = std::move(built.second.second);
(void)y1;
}
input_type off0{};
input_type off1{};
return dpf::make_party_key_pair(
dpf_type{root0, correction_words, correction_advice,
leaves0, beavers0, off0},
dpf_type{root1, correction_words, correction_advice,
leaves1, beavers1, off1});
}
/// @brief Plaintext-β multi-output classic path (unchanged).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam OutputTs output ts
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam CwProtocol correction-word protocol
/// @param arith the `arith`
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param root_sampler the `root_sampler`
/// @param proto the `proto`
/// @param y the `y`
/// @param ys the `ys`
/// @return Plaintext-β multi-output classic path (unchanged)
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs,
typename RootSampler,
typename CwProtocol,
typename = std::enable_if_t<(sizeof...(OutputTs) > 0)>>
auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1,
RootSampler & root_sampler, CwProtocol & proto, OutputT && y,
OutputTs && ...ys)
{
static_assert(!dpf::is_wildcard_v<InputT>,
"Doerner–Shelat gen takes shares of a concrete point");
static_assert(!dpf::is_secret_share_v<InputT>,
"Doerner–Shelat: pass additive_share of xor_wrapper, or raw 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;
proto.encode_walk_shares(x0, x1, arith);
using tree = dpf::tree_traits<InteriorPRG>;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
node roots[2];
HEDLEY_PRAGMA(GCC diagnostic pop)
tree::root_init(roots, [&]() -> node {
return static_cast<node>(root_sampler());
});
const node root0 = roots[0];
const node root1 = roots[1];
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
ds_gen_state<node> st;
HEDLEY_PRAGMA(GCC diagnostic pop)
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, depth, 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});
}
/// @brief Single-output plaintext β (disambiguates from arith_out overload).
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
/// @tparam InputT input domain type
/// @tparam OutputT output type
/// @tparam RootSampler sampler for the Doerner–Shelat root seed
/// @tparam CwProtocol correction-word protocol
/// @param arith the `arith`
/// @param x0 the `x0`
/// @param x1 the `x1`
/// @param root_sampler the `root_sampler`
/// @param proto the `proto`
/// @param y the `y`
/// @return Single-output plaintext β (disambiguates from arith_out overload)
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename RootSampler,
typename CwProtocol>
auto make_dpf_doerner_shelat_impl(bool arith, InputT x0, InputT x1,
RootSampler & root_sampler, CwProtocol & proto, OutputT && y)
{
return make_dpf_doerner_shelat_impl<InteriorPRG, ExteriorPRG>(arith, false,
std::move(x0), std::move(x1), root_sampler, proto,
std::forward<OutputT>(y), OutputT{});
}
} // namespace detail
/// @brief 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__