libdpf/party/dist_ds.hpp

2727 lines
107 KiB
C++
Raw Permalink Normal View History

/// @file party/dist_ds.hpp
/// @brief Two-party Doerner–Shelat point/DCF generation over trio sockets.
/// @details p2 deals correlated pads and, for Half-Tree, the roots. p0 and p1
/// each keep only their own seed and input share. Correction words
/// and VDPF correction seeds are opened from blinds on the p0–p1
/// link. DCF value words use the same blinded-difference pattern in
/// their output ring while `Va` remains additively shared.
///
/// Additive point shares are converted to XOR shares by a beaver
/// ripple-carry before the walk. The sum is not opened. Callers that
/// pass `RevealPoint` (comparison: `Reveal`) exchange a share and
/// get that value back. Otherwise the walk keeps XOR shares: the
/// verifiable hash is an AES circuit, comparison value words and
/// blocked suffixes are a bit-serial mux, and a packed lane is a
/// shared-output mux whose only opened result is the public leaf
/// correction word. p2 never receives the point. Interval keys keep
/// the mask `r` as XOR shares: `correction_s` and `γ = r − 1` are
/// computed from those shares, and the inner comparison walk does not
/// take `PointIsClear`. Pass `Reveal` to reconstruct `r` for tests.
#ifndef LIBDPF_PARTY_DIST_DS_HPP__
#define LIBDPF_PARTY_DIST_DS_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <initializer_list>
#include <optional>
#include <stdexcept>
#include <type_traits>
#include <utility>
#include <vector>
#include "dpf.hpp"
#include "dpf/net/trio.hpp"
#include "dpf/net/mux_sink.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/net/sink_exchange.hpp"
#include "flow_util.hpp"
#include "aes_mmo_ref.hpp"
#include "key_io.hpp"
namespace dpf
{
namespace party
{
namespace dist
{
using net::role;
using net::trio;
/// @name Wire messages
/// @brief Trivially copyable frames for pads, blinds, and Beaver rounds.
/// @{
struct cw_pad_msg
{
simde__m128i rand{};
simde__m128i gamma{};
std::uint8_t bit = 0;
};
struct and_share_msg
{
std::uint8_t a = 0;
simde__m128i b{};
simde__m128i c{};
};
struct level_pad_msg
{
cw_pad_msg cw{};
and_share_msg mine{};
and_share_msg theirs{};
};
/// @brief Beaver shares for the packed-leaf mux. `lg == 0` still sends one
/// unused pad. Otherwise two products per mux node (each party's
/// delta), `2 * (2^lg - 1)` in total.
template <std::size_t N>
struct leaf_pad_msg
{
static_assert(N >= 1, "leaf mux sends at least one pad");
std::array<and_share_msg, N> lanes{};
};
/// @brief Key together with the tree prefix both holders reconstructed.
/// @details This is `α` shifted right by the packed-lane width: the bits
/// passed to `hash_node`. It is not the packed lane.
template <typename Key, typename Point>
struct opened_prefix
{
Key dpf_key;
Point opened_prefix{};
};
/// @brief Key together with a packed lane both holders reconstructed.
template <typename Key>
struct opened_lane
{
Key dpf_key;
unsigned opened_lane = 0;
};
/// @brief Verifiable wildcard: hash prefix and the lane the scale vector needs.
template <typename Key, typename Point>
struct opened_prefix_lane
{
Key dpf_key;
Point opened_prefix{};
unsigned opened_lane = 0;
};
/// @brief Key together with the encoded point both holders reconstructed.
template <typename Key, typename Point>
struct opened_point
{
Key dpf_key;
Point opened_point{};
};
template <typename OutputT, typename Leaf>
struct wildcard_leaf_pad_msg
{
OutputT output_blind{};
Leaf vector_blind{};
Leaf peer_vector_blind{};
Leaf cross{};
Leaf zero_leaf{};
};
struct ring_zero_share_msg
{
std::uint64_t share = 0;
};
struct blind_msg
{
simde__m128i msg{};
std::uint8_t bit = 0;
};
struct share_msg
{
simde__m128i share{};
};
struct advice_msg
{
std::uint8_t aL = 0;
std::uint8_t aR = 0;
};
struct and_round1_msg
{
std::uint8_t d_bit = 0;
simde__m128i b_bit{};
simde__m128i e_m{};
std::uint8_t a_m = 0;
};
struct and_round2_msg
{
simde__m128i z_bit{};
};
/// @brief One party's shares of a bit-Beaver triple `(α, β, α∧β)`.
struct bit_and_pad_msg
{
std::uint8_t a = 0;
std::uint8_t b = 0;
std::uint8_t c = 0;
};
/// @brief Masked bit shares `x⊕α` and `y⊕β` for one carry AND.
struct bit_and_mask_msg
{
std::uint8_t x = 0;
std::uint8_t y = 0;
};
/// @brief Mask for one bit of a XOR-share → additive-share conversion.
/// @details `r` is a XOR share of a random bit. `add` is an additive share of
/// that same bit in `uint64_t`. Opening `bit ⊕ r` does not open `bit`.
struct b2a_pad_msg
{
std::uint8_t r = 0;
std::uint64_t add = 0;
};
/// @brief One bit of a secret-bit × known-word product, then a B2A of that bit.
struct word_bit_pad
{
std::uint8_t a = 0;
std::uint8_t b = 0;
std::uint8_t c = 0;
std::uint8_t r = 0;
std::uint64_t add = 0;
};
/// @brief Pads for one oblivious comparison level: two 64-bit words and the path bit.
struct cmp_level_obliv_msg
{
std::array<word_bit_pad, 128> bits{};
b2a_pad_msg ai{};
};
/// @brief Pads for a 2-bit blocked suffix: `s0 ∧ s1`, then B2A of OR, AND, and `s1`.
struct suffix_obliv_msg
{
bit_and_pad_msg and_pad{};
b2a_pad_msg b2a[3]{};
};
/// @brief Products in the packed-leaf mux. `lg == 0` keeps a single unused pad.
inline constexpr std::size_t leaf_and_slots(std::size_t lg) noexcept
{
return lg == 0 ? std::size_t{1} : (std::size_t{2} * ((std::size_t{1} << lg) - 1));
}
inline b2a_pad_msg b2a_for_party(std::uint8_t r_share, std::uint64_t add) noexcept
{
return b2a_pad_msg{r_share, add};
}
template <typename Pad>
inline void sample_b2a(Pad & rng, b2a_pad_msg & p0, b2a_pad_msg & p1)
{
const auto shares = dpf::beavers::sample_bit_arith(rng);
p0 = b2a_for_party(shares.xor0, shares.add0);
p1 = b2a_for_party(shares.xor1, shares.add1);
}
template <typename Pad>
inline void send_b2a_vec(trio & net, std::size_t n, Pad & rng)
{
std::vector<b2a_pad_msg> a(n), b(n);
for (std::size_t i = 0; i < n; ++i)
sample_b2a(rng, a[i], b[i]);
net.send_vec_to(role::p0, net::msg::beaver_tape, a);
net.send_vec_to(role::p1, net::msg::beaver_tape, b);
}
inline word_bit_pad word_bit_for(const bit_and_pad_msg & beaver, const b2a_pad_msg & b2a)
{
word_bit_pad m;
m.a = beaver.a;
m.b = beaver.b;
m.c = beaver.c;
m.r = b2a.r;
m.add = b2a.add;
return m;
}
inline bit_and_pad_msg bit_pad_for(const detail::ds_bit_triple & t, int party)
{
if (party == 0)
return bit_and_pad_msg{t.a0, t.b0, t.c0};
return bit_and_pad_msg{t.a1, t.b1, t.c1};
}
template <typename Pad>
inline void send_cmp_level_pads(trio & net, Pad & rng)
{
cmp_level_obliv_msg m0{}, m1{};
for (std::size_t i = 0; i < 128; ++i)
{
const auto t = detail::ds_sample_bit_and(rng);
b2a_pad_msg a{}, b{};
sample_b2a(rng, a, b);
m0.bits[i] = word_bit_for(bit_pad_for(t, 0), a);
m1.bits[i] = word_bit_for(bit_pad_for(t, 1), b);
}
sample_b2a(rng, m0.ai, m1.ai);
net.send_to(role::p0, net::msg::beaver_tape, m0);
net.send_to(role::p1, net::msg::beaver_tape, m1);
}
template <typename Pad>
inline void send_suffix_pads(trio & net, Pad & rng)
{
suffix_obliv_msg m0{}, m1{};
const auto t = detail::ds_sample_bit_and(rng);
m0.and_pad = bit_pad_for(t, 0);
m1.and_pad = bit_pad_for(t, 1);
for (int i = 0; i < 3; ++i)
sample_b2a(rng, m0.b2a[i], m1.b2a[i]);
net.send_to(role::p0, net::msg::beaver_tape, m0);
net.send_to(role::p1, net::msg::beaver_tape, m1);
}
template <typename Pad>
inline void send_bit_and_vec(trio & net, std::size_t n, Pad & rng)
{
std::vector<bit_and_pad_msg> a(n), b(n);
for (std::size_t i = 0; i < n; ++i)
{
const auto t = detail::ds_sample_bit_and(rng);
a[i] = bit_pad_for(t, 0);
b[i] = bit_pad_for(t, 1);
}
net.send_vec_to(role::p0, net::msg::beaver_tape, a);
net.send_vec_to(role::p1, net::msg::beaver_tape, b);
}
/// @}
inline cw_pad_msg cw_pad_for(const detail::ds_cw_party & p)
{
cw_pad_msg m;
m.rand = p.rand;
m.gamma = p.gamma;
m.bit = p.bit;
return m;
}
inline and_share_msg and_share_for(const detail::ds_and_pads & p, int party)
{
and_share_msg m;
if (party == 0)
{
m.a = p.a0;
m.b = p.b0_share;
m.c = p.c0_share;
}
else
{
m.a = p.a1;
m.b = p.b1_share;
m.c = p.c1_share;
}
return m;
}
inline level_pad_msg level_pad_for(const detail::ds_cw_pads & cw,
const detail::ds_and_pads & my_and, const detail::ds_and_pads & their_and,
int party)
{
level_pad_msg m;
m.cw = cw_pad_for(party == 0 ? cw.p0 : cw.p1);
m.mine = and_share_for(my_and, party);
m.theirs = and_share_for(their_and, party);
return m;
}
template <typename Leaf>
Leaf random_leaf();
template <int Me>
cs_block open_cs(net::trio & net, const cs_block & mine);
#include "oblivious_hash.hpp"
template <typename Pad = detail::urandom_pad_rng>
inline void send_hash_tape(trio & net, Pad && htape = Pad{})
{
constexpr std::size_t nand = hash_level_and_count();
std::vector<bit_and_pad_msg> t0(nand), t1(nand);
for (std::size_t i = 0; i < nand; ++i)
{
const auto full = detail::ds_sample_bit_and(htape);
t0[i] = bit_pad_for(full, 0);
t1[i] = bit_pad_for(full, 1);
}
net.send_vec_to(role::p0, net::msg::beaver_tape, t0);
net.send_vec_to(role::p1, net::msg::beaver_tape, t1);
}
/// @brief Dealer samples Half-Tree roots (when needed) and one pad per level.
/// @tparam InteriorPRG PRG that expands interior nodes
/// @tparam InputT input domain type
/// @tparam OutputT payload type
/// @param net the connected trio. This process must be p2.
/// \complexity O(n) sampling. n is `depth`. The loop body is one CW pad and two AND pads per level.
/// \rounds n level-pad sends to each of p0 and p1, plus one leaf-pad send to each. A Half-Tree also sends one root to each. `ObliviousHash` adds one `bit_and_pad_msg` vector per level (`hash_level_and_count()` bit triples).
/// \communication Those messages. Each `level_pad_msg` carries one party's share of the CW pad and the two AND pads.
/// \preprocessing This function is the preprocessing. The online walk is `point_party`.
/// @param pads beaver-pad stream. `dpf::prg_pad_rng` reads that stream from a PRG seed.
/// @param root_draw Half-Tree roots. `dpf::pseudorandom_root_sampler` is the PRG form.
template <typename InteriorPRG, typename InputT, typename OutputT,
bool ObliviousHash = false,
typename Pad = detail::urandom_pad_rng,
typename RootDraw = dpf::uniform_node_sampler<
typename dpf::tree_traits<InteriorPRG>::node>>
void deal_point(trio & net, Pad && pads = Pad{}, RootDraw && root_draw = RootDraw{})
{
using dpf_type =
utils::dpf_type_t<InteriorPRG, InteriorPRG, InputT, OutputT, verifiable>;
using node = typename dpf_type::interior_node;
using tree = dpf::tree_traits<InteriorPRG>;
constexpr std::size_t depth = dpf_type::depth;
if constexpr (tree::is_half_tree)
{
node roots[2];
tree::root_init(roots, [&]() -> node { return root_draw(); });
net.send_to(role::p0, net::msg::dpf_key, roots[0]);
net.send_to(role::p1, net::msg::dpf_key, roots[1]);
}
for (std::size_t level = 0; level < depth; ++level)
{
const auto cw = detail::ds_sample_cw(pads);
const auto and0 = detail::ds_sample_and(pads);
const auto and1 = detail::ds_sample_and(pads);
net.send_to(role::p0, net::msg::beaver_tape, level_pad_for(cw, and0, and1, 0));
net.send_to(role::p1, net::msg::beaver_tape, level_pad_for(cw, and1, and0, 1));
if constexpr (ObliviousHash)
{
constexpr std::size_t nand = hash_level_and_count();
std::vector<bit_and_pad_msg> t0(nand), t1(nand);
for (std::size_t i = 0; i < nand; ++i)
{
const auto full = ::dpf::detail::ds_sample_bit_and(pads);
t0[i] = bit_pad_for(full, 0);
t1[i] = bit_pad_for(full, 1);
}
net.send_vec_to(role::p0, net::msg::beaver_tape, t0);
net.send_vec_to(role::p1, net::msg::beaver_tape, t1);
}
}
constexpr std::size_t lg = dpf_type::lg_outputs_per_leaf;
constexpr std::size_t n_leaf_ands = leaf_and_slots(lg);
leaf_pad_msg<n_leaf_ands> pads0{};
leaf_pad_msg<n_leaf_ands> pads1{};
for (std::size_t i = 0; i < n_leaf_ands; ++i)
{
const auto lane = detail::ds_sample_and(pads);
pads0.lanes[i] = and_share_for(lane, 0);
pads1.lanes[i] = and_share_for(lane, 1);
}
net.send_to(role::p0, net::msg::beaver_tape, pads0);
net.send_to(role::p1, net::msg::beaver_tape, pads1);
// Concrete lanes always mux in shares. A wildcard lane is muxed only when
// the caller did not ask for it (ObliviousHash is that request's inverse).
constexpr bool leaf_b2a = lg > 0
&& !utils::has_characteristic_two_v<dpf::concrete_type_t<OutputT>>
&& (!dpf::is_wildcard_v<OutputT> || ObliviousHash);
if constexpr (leaf_b2a)
send_b2a_vec(net, n_leaf_ands * 128, pads);
if constexpr (dpf::is_wildcard_v<OutputT>)
{
using concrete = dpf::concrete_type_t<OutputT>;
using exterior_node = typename dpf_type::exterior_node;
using leaf_type = dpf::leaf_node_t<exterior_node, concrete>;
concrete out0{}, out1{};
leaf_type vec0{}, vec1{};
constexpr std::size_t nlanes =
dpf::outputs_per_leaf_v<concrete, exterior_node>;
auto lane_rng = [&pads]() -> concrete {
concrete value{};
auto * bytes = reinterpret_cast<unsigned char *>(std::addressof(value));
std::size_t left = sizeof(concrete);
std::size_t off = 0;
while (left != 0)
{
const auto block = pads.block();
const std::size_t n = left < sizeof(block) ? left : sizeof(block);
std::memcpy(bytes + off, &block, n);
off += n;
left -= n;
}
return value;
};
dpf::beavers::fill_wildcard_scale_blinds(out0, out1, vec0, vec1,
nlanes > 0 ? nlanes : std::size_t{1}, lane_rng);
const leaf_type zero0 = random_leaf<leaf_type>();
const leaf_type zero1 =
dpf::subtract_leaf<concrete>(leaf_type{}, zero0);
const leaf_type cross0 = dpf::multiply_leaf(vec0, out1);
const leaf_type cross1 = dpf::multiply_leaf(vec1, out0);
net.send_to(role::p0, net::msg::beaver_tape, wildcard_leaf_pad_msg<concrete, leaf_type>{
out0, vec0, vec1, cross0, zero0});
net.send_to(role::p1, net::msg::beaver_tape, wildcard_leaf_pad_msg<concrete, leaf_type>{
out1, vec1, vec0, cross1, zero1});
}
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
typename Spec, typename ...Tags>
using comparison_pair_t = decltype(dpf::make_dpf<InteriorPRG, ExteriorPRG>(
std::declval<InputT>(), std::declval<Spec>(), std::declval<Tags>()...));
inline std::uint64_t random_ring_word(std::uint64_t mask);
inline void deal_ring_zero(trio & net, std::uint64_t mask);
/// @brief Dealer samples roots, DS level pads, and a comparison-ring zero share.
/// @tparam Oblivious When set, also deal the bit tapes the value-word, suffix,
/// and verifiable-hash circuits consume. The target itself is not opened.
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
bool Oblivious = false, typename Spec, typename ...Tags>
void deal_comparison(trio & net, const Spec & spec, const Tags & ...tags);
template <typename Pad, typename RootDraw,
typename InteriorPRG, typename ExteriorPRG, typename InputT,
bool Oblivious = false,
typename Spec, typename ...Tags>
void deal_comparison(Pad && pads, RootDraw && root_draw, trio & net,
const Spec & spec, const Tags & ...tags)
{
(void)std::initializer_list<int>{((void)tags, 0)...};
using pair_type =
comparison_pair_t<InteriorPRG, ExteriorPRG, InputT, Spec, Tags...>;
using key_type = typename pair_type::first_type::key_type;
using node = typename key_type::interior_node;
using tree = dpf::tree_traits<InteriorPRG>;
static_assert(key_type::num_outputs == 0,
"dist comparison keygen currently builds comparison-only keys");
dcf_runtime_spec runtime{};
bool found = false;
detail::incr::collect_cmp_one(runtime, found, spec);
if (!found)
throw std::invalid_argument("dist comparison keygen needs one comparison");
if (runtime.prefix == 0)
runtime.prefix = utils::bitlength_of_v<InputT>;
if constexpr (Oblivious)
{
if (is_paint_kind(runtime.kind))
throw std::invalid_argument(
"oblivious comparison does not support paint kinds");
}
if constexpr (tree::is_half_tree)
{
node roots[2];
tree::root_init(roots, [&]() -> node { return root_draw(); });
net.send_to(role::p0, net::msg::dpf_key, roots[0]);
net.send_to(role::p1, net::msg::dpf_key, roots[1]);
}
if constexpr (Oblivious)
{
// leq/gt: deal the `∧ path bits` tape. Holders keep that product
// shared (no open_bit) so the domain-edge bit stays hidden.
const bool edge = runtime.kind == cmp_kind::leq
|| runtime.kind == cmp_kind::gt;
if (edge && runtime.prefix > 1)
send_bit_and_vec(net, runtime.prefix - 1, pads);
}
for (std::size_t level = 0; level < key_type::depth; ++level)
{
const auto cw = detail::ds_sample_cw(pads);
const auto and0 = detail::ds_sample_and(pads);
const auto and1 = detail::ds_sample_and(pads);
net.send_to(role::p0, net::msg::beaver_tape, level_pad_for(cw, and0, and1, 0));
net.send_to(role::p1, net::msg::beaver_tape, level_pad_for(cw, and1, and0, 1));
if constexpr (Oblivious)
{
if (key_type::cmp_block == 0 && level < runtime.prefix)
send_cmp_level_pads(net, pads);
if constexpr (key_type::cmp_block > 0 && key_type::cmp_q > 0)
{
if (level + 1 == key_type::cmp_h)
send_suffix_pads(net, pads);
}
if constexpr (key_type::is_verifiable)
send_hash_tape(net, pads);
}
}
deal_ring_zero(net, runtime.mask);
}
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
bool Oblivious, typename Spec, typename ...Tags>
void deal_comparison(trio & net, const Spec & spec, const Tags & ...tags)
{
using pair_type =
comparison_pair_t<InteriorPRG, ExteriorPRG, InputT, Spec, Tags...>;
using node = typename pair_type::first_type::key_type::interior_node;
detail::urandom_pad_rng pads;
dpf::uniform_node_sampler<node> roots;
deal_comparison<detail::urandom_pad_rng &, dpf::uniform_node_sampler<node> &,
InteriorPRG, ExteriorPRG, InputT, Oblivious, Spec, Tags...>(
pads, roots, net, spec, tags...);
}
template <typename Leaf>
Leaf random_leaf()
{
Leaf out{};
auto * bytes = reinterpret_cast<unsigned char *>(&out);
std::size_t left = sizeof(Leaf);
while (left > 0)
{
const auto block = dpf::uniform_sample<simde__m128i>();
const std::size_t n = left < sizeof(block) ? left : sizeof(block);
std::memcpy(bytes, &block, n);
bytes += n;
left -= n;
}
return out;
}
inline std::uint64_t random_ring_word(std::uint64_t mask)
{
return detail::dcf_impl::convert_node(
dpf::uniform_sample<simde__m128i>(), mask);
}
/// @brief Dealer sends additive shares of zero in the comparison ring.
inline void deal_ring_zero(trio & net, std::uint64_t mask)
{
const std::uint64_t r = random_ring_word(mask);
net.send_to(role::p0, net::msg::beaver_tape, ring_zero_share_msg{r});
net.send_to(role::p1, net::msg::beaver_tape, ring_zero_share_msg{detail::dcf_impl::neg_m(r, mask)});
}
/// @brief Open `party1 - party0` without revealing either local ring word.
/// @details This is the scalar analogue of the blinded leaf-mask difference:
/// p0 blinds its word, p1 subtracts it, and p0 removes the blind before
/// returning the public difference.
template <int Me>
HEDLEY_WARN_UNUSED_RESULT
std::uint64_t open_ring_difference(trio & net, std::uint64_t mine,
std::uint64_t mask)
{
static_assert(Me == 0 || Me == 1, "ring opener party is 0 or 1");
const role peer = Me == 0 ? role::p1 : role::p0;
mine &= mask;
if constexpr (Me == 0)
{
const std::uint64_t pad = random_ring_word(mask);
const std::uint64_t blinded = (mine + pad) & mask;
net.send_to(peer, net::msg::ring_vector, blinded);
const std::uint64_t back = net.recv_from<std::uint64_t>(peer, net::msg::ring_vector);
const std::uint64_t opened = (back + pad) & mask;
net.send_to(peer, net::msg::ring_vector, opened);
return opened;
}
else
{
const std::uint64_t blinded =
net.recv_from<std::uint64_t>(peer, net::msg::ring_vector);
const std::uint64_t back =
(mine + detail::dcf_impl::neg_m(blinded, mask)) & mask;
net.send_to(peer, net::msg::ring_vector, back);
return net.recv_from<std::uint64_t>(peer, net::msg::ring_vector) & mask;
}
}
template <typename OutputT, typename Leaf>
Leaf leaf_add(Leaf a, Leaf b)
{
const Leaf neg = dpf::subtract_leaf<OutputT>(Leaf{}, b);
return dpf::subtract_leaf<OutputT>(a, neg);
}
#include "oblivious_select.hpp"
inline cs_block xor_cs(const cs_block & a, const cs_block & b)
{
return cs_block{
simde_mm_xor_si128(a[0], b[0]),
simde_mm_xor_si128(a[1], b[1]),
simde_mm_xor_si128(a[2], b[2]),
simde_mm_xor_si128(a[3], b[3])};
}
inline cs_block random_cs()
{
return cs_block{
dpf::uniform_sample<simde__m128i>(),
dpf::uniform_sample<simde__m128i>(),
dpf::uniform_sample<simde__m128i>(),
dpf::uniform_sample<simde__m128i>()};
}
/// @brief Open `H0 XOR H1` with the same pad pattern as the leaf mask.
/// @tparam Me `0` for p0, `1` for p1
/// @param net the connected trio
/// @param mine this party's `hash_node` digest
/// @return the public correction seed
template <int Me>
HEDLEY_WARN_UNUSED_RESULT
cs_block open_cs(trio & net, const cs_block & mine)
{
const role peer = Me == 0 ? role::p1 : role::p0;
if constexpr (Me == 0)
{
const cs_block pad = random_cs();
const cs_block blinded = xor_cs(mine, pad);
net.send_to(peer, net::msg::delta, blinded);
const cs_block back = net.recv_from<cs_block>(peer, net::msg::delta);
const cs_block cs = xor_cs(back, pad);
net.send_to(peer, net::msg::delta, cs);
return cs;
}
else
{
const cs_block blinded = net.recv_from<cs_block>(peer, net::msg::delta);
const cs_block back = xor_cs(mine, blinded);
net.send_to(peer, net::msg::delta, back);
return net.recv_from<cs_block>(peer, net::msg::delta);
}
}
/// @brief Product of a XOR-shared bit and a block known to one party.
/// @details Each party contributes `bit_share XOR a_share`. Only the party
/// holding the block receives the product. When that block is a
/// public difference of two leaves, the holder can read the bit
/// off a nonzero product; the peer receives zero and does not.
/// @tparam Exchange callable `peer_msg = exch(my_msg)` on the p0–p1 link
/// @param exch the exchange
/// @param i_hold_block `true` when this party knows `block`
/// @param bit_share this party's XOR share of the bit
/// @param block the block, meaningful only when `i_hold_block`
/// @param mine this party's Beaver triple share
/// @return the product when this party holds the block, otherwise zero
template <typename Exchange>
HEDLEY_WARN_UNUSED_RESULT
simde__m128i beaver_shared_bit(Exchange && exch, bool i_hold_block,
std::uint8_t bit_share, simde__m128i block, const and_share_msg & mine)
{
and_round1_msg msg{};
msg.d_bit = static_cast<std::uint8_t>(bit_share ^ mine.a);
msg.b_bit = mine.b;
msg.a_m = mine.a;
msg.e_m = i_hold_block ? detail::ds_xor(block, mine.b) : mine.b;
const and_round1_msg peer = exch(msg);
const std::uint8_t d = static_cast<std::uint8_t>(msg.d_bit ^ peer.d_bit);
const simde__m128i e = i_hold_block
? detail::ds_xor(msg.e_m, peer.b_bit)
: detail::ds_xor(peer.e_m, mine.b);
and_round2_msg zmsg{};
zmsg.z_bit = detail::ds_xor(
detail::ds_xor(detail::ds_gate(d, mine.b), detail::ds_gate(mine.a, e)),
mine.c);
const and_round2_msg zpeer = exch(zmsg);
if (!i_hold_block)
return simde_mm_setzero_si128();
const simde__m128i z_me = detail::ds_xor(
detail::ds_xor(
detail::ds_xor(detail::ds_gate(d, e), detail::ds_gate(d, mine.b)),
detail::ds_gate(mine.a, e)),
mine.c);
return detail::ds_xor(z_me, zpeer.z_bit);
}
/// @brief Beaver product of a bit known to the other party and our block.
/// @details Both parties call this. Each is the block holder for its own
/// block and the bit holder for the peer's block, in one exchange.
/// @tparam Exchange callable `peer_msg = exch(my_msg)` on the p0–p1 link
/// @param exch the exchange
/// @param i_hold_m `true` when this party knows `M`
/// @param my_bit this party's share of the selector bit
/// @param M this party's block
/// @param my_triple Beaver pads for the product this party reconstructs
/// @param their_triple Beaver pads for the product the peer reconstructs
/// @return the product when this party holds `M`, otherwise zero
template <typename Exchange>
HEDLEY_WARN_UNUSED_RESULT
simde__m128i beaver_bit_block(Exchange && exch, bool i_hold_m,
std::uint8_t my_bit, simde__m128i M, const and_share_msg & my_triple,
const and_share_msg & their_triple)
{
and_round1_msg mine{};
mine.d_bit = static_cast<std::uint8_t>(my_bit ^ their_triple.a);
mine.b_bit = their_triple.b;
mine.e_m = detail::ds_xor(M, my_triple.b);
mine.a_m = my_triple.a;
const and_round1_msg peer = exch(mine);
const std::uint8_t d_as_bit = static_cast<std::uint8_t>(
(my_bit ^ their_triple.a) ^ peer.a_m);
const simde__m128i e_as_bit = detail::ds_xor(peer.e_m, their_triple.b);
and_round2_msg zmsg{};
zmsg.z_bit = detail::ds_xor(
detail::ds_xor(detail::ds_gate(d_as_bit, their_triple.b),
detail::ds_gate(their_triple.a, e_as_bit)),
their_triple.c);
const and_round2_msg zpeer = exch(zmsg);
if (!i_hold_m)
return simde_mm_setzero_si128();
const std::uint8_t d = static_cast<std::uint8_t>(peer.d_bit ^ my_triple.a);
const simde__m128i e = detail::ds_xor(detail::ds_xor(M, my_triple.b), peer.b_bit);
const simde__m128i z_me = detail::ds_xor(
detail::ds_xor(
detail::ds_xor(detail::ds_gate(d, e), detail::ds_gate(d, my_triple.b)),
detail::ds_gate(my_triple.a, e)),
my_triple.c);
return detail::ds_xor(z_me, zpeer.z_bit);
}
template <typename InputT>
HEDLEY_WARN_UNUSED_RESULT
InputT input_from_prefix(psnip_uint64_t prefix)
{
constexpr auto to_int = utils::to_integral_type<InputT>{};
using U = std::make_unsigned_t<std::decay_t<decltype(to_int(std::declval<InputT>()))>>;
using FromI = typename utils::make_from_integral_value<InputT>::integral_type;
return utils::make_from_integral_value<InputT>{}(
static_cast<FromI>(static_cast<U>(prefix)));
}
template <bool Verifiable, bool WildLane, typename Held, typename InputT>
HEDLEY_WARN_UNUSED_RESULT
auto pack_point_held(Held held, psnip_uint64_t prefix, unsigned lane)
{
if constexpr (Verifiable && WildLane)
{
return opened_prefix_lane<Held, InputT>{
std::move(held), input_from_prefix<InputT>(prefix), lane};
}
else if constexpr (Verifiable)
{
return opened_prefix<Held, InputT>{
std::move(held), input_from_prefix<InputT>(prefix)};
}
else if constexpr (WildLane)
{
return opened_lane<Held>{std::move(held), lane};
}
else
{
return held;
}
}
/// @brief One computing party builds its point key from dealt pads.
/// @tparam InteriorPRG PRG that expands interior nodes
/// @tparam ExteriorPRG PRG that expands the root
/// @tparam InputT input domain type
/// @tparam OutputT payload type
/// @tparam Self `role::p0` or `role::p1`
/// @param net the connected trio
/// @param x_share this party's XOR share of the target, or the opened target
/// on p0 when `already_encoded` is set
/// @param beta the payload
/// @param already_encoded p0's share is the opened target and p1's share is 0
/// @return Verifiable keys return `opened_prefix` (and `opened_lane` as well
/// when the payload is a packed wildcard). Other keys return the key,
/// or `opened_lane` when a wildcard scale vector needs the clear lane.
/// \complexity O(n) local PRG expansions (one `expand` per level) plus the exchanges below. n is `depth`.
/// \rounds n dealer pad deliveries, and per level three `exchange_with` calls (blind, CW share, advice) plus one `beaver_bit_block` exchange. A verifiable `RevealPoint` walk adds a prefix-bit exchange and `open_cs` (two further block exchanges). Without reveal, a verifiable walk instead receives a `bit_and_pad_msg` vector from p2. Counted in the `for (level)` loop of `point_party`.
/// \communication Per level, one `level_pad_msg` from p2 and the peer messages in those exchanges (each exchange is one `blind_msg`, `share_msg`, or `advice_msg`, 16 bytes plus a bit, or a block). Half-tree keys also receive the root from p2.
/// \preprocessing `deal_point` samples the roots and one `level_pad_msg` per level before this walk.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
role Self,
bool Extractable = false,
bool RevealPoint = false,
typename RootDraw = dpf::uniform_node_sampler<
typename dpf::tree_traits<InteriorPRG>::node>>
HEDLEY_WARN_UNUSED_RESULT
auto point_party(trio & net, InputT x_share, OutputT beta,
bool already_encoded = false,
typename dpf::tree_traits<InteriorPRG>::node * leaf_seed_out = nullptr,
RootDraw && root_draw = RootDraw{},
net::RoundSink * batch_sink = nullptr,
std::size_t batch_index = 0)
{
static_assert(Self == role::p0 || Self == role::p1,
"dist DS parties are p0 and p1");
using auth_tag = std::conditional_t<Extractable, extractable, verifiable>;
using dpf_type =
utils::dpf_type_t<InteriorPRG, ExteriorPRG, InputT, OutputT, auth_tag>;
using node = typename dpf_type::interior_node;
using tree = dpf::tree_traits<InteriorPRG>;
constexpr int me = Self == role::p0 ? 0 : 1;
constexpr std::size_t depth = dpf_type::depth;
constexpr auto to_int = utils::to_integral_type<InputT>{};
const role peer = Self == role::p0 ? role::p1 : role::p0;
std::optional<net::sink_exchange> sink_ex;
if (batch_sink)
sink_ex.emplace(*batch_sink, batch_index);
auto exch = [&](const auto & mine) {
using T = std::decay_t<decltype(mine)>;
if (sink_ex)
return (*sink_ex)(static_cast<const T &>(mine));
return net.exchange_with(peer, static_cast<const T &>(mine),
net::msg::delta);
};
if (me == 0 && !already_encoded)
utils::flip_msb_if_signed_integral(x_share);
node root{};
if constexpr (tree::is_half_tree)
root = net.recv_from<node>(role::p2, net::msg::dpf_key);
else if constexpr (me == 0)
root = dpf::unset_lo_bit(root_draw());
else
root = dpf::set_lo_bit(root_draw());
node seed = root;
int home = me;
typename dpf_type::correction_words_array cws{};
typename dpf_type::correction_advice_array advice{};
typename dpf_type::correction_seeds_array seeds{};
[[maybe_unused]] psnip_uint64_t prefix = 0;
[[maybe_unused]] std::uint64_t prefix_share = 0;
auto mask = dpf_type::msb_mask;
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{
const auto pads = net.recv_from<level_pad_msg>(role::p2, net::msg::beaver_tape);
const std::uint8_t my_bit = static_cast<std::uint8_t>(
!!(to_int(mask) & to_int(x_share)));
const bool is_last = tree::is_last_level(level, depth);
const std::uint8_t adv = static_cast<std::uint8_t>(dpf::get_lo_bit(seed));
const auto child = tree::expand(seed, is_last);
const node L = child[0];
const node R = child[1];
blind_msg blind{};
blind.bit = static_cast<std::uint8_t>(my_bit ^ pads.cw.bit);
blind.msg = detail::ds_xor(detail::ds_xor(L, R), pads.cw.rand);
const blind_msg their_blind = exch(blind);
detail::ds_cw_party mine_pad{pads.cw.rand, pads.cw.gamma, pads.cw.bit};
detail::ds_blind peer_blind{their_blind.msg, their_blind.bit};
share_msg mine_share{};
mine_share.share = detail::ds_cw_share(L, R, my_bit, mine_pad, peer_blind);
const share_msg their_share = exch(mine_share);
node cw = detail::ds_xor(mine_share.share, their_share.share);
advice_msg am{};
am.aL = static_cast<std::uint8_t>(dpf::get_lo_bit(L) ^ my_bit);
am.aR = static_cast<std::uint8_t>(dpf::get_lo_bit(R) ^ my_bit);
const advice_msg at = exch(am);
std::uint8_t tpack = static_cast<std::uint8_t>(
((am.aR ^ at.aR) << 1) | ((am.aL ^ at.aL ^ 1u) & 1u));
if constexpr (tree::is_half_tree)
{
if (!is_last)
tpack = 0;
}
const bool exp_is_mine = home == me;
node M{};
node base{};
if constexpr (tree::is_half_tree)
{
if (!is_last)
{
M = detail::ds_xor(L, R);
base = (adv & 1u) ? detail::ds_xor(L, cw) : L;
}
else
{
detail::ds_next_terms(L, R, adv, cw, tpack, M, base);
}
}
else
{
detail::ds_next_terms(L, R, adv, cw, tpack, M, base);
}
const simde__m128i cross = beaver_bit_block(exch, true, my_bit, M,
pads.mine, pads.theirs);
const simde__m128i local = detail::ds_gate(my_bit, M);
const simde__m128i exp_gate = exp_is_mine ? local : cross;
const simde__m128i rec_prod = exp_is_mine ? cross : local;
seed = detail::ds_xor(base, detail::ds_xor(exp_gate, rec_prod));
home ^= 1;
cws[level] = cw;
advice[level] = tpack;
// The proof hash tweaks AES with this prefix. Reveal it only when
// the caller asked; otherwise hash the shared prefix.
if constexpr (dpf_type::is_verifiable && RevealPoint)
{
const std::uint8_t peer_bit = exch(my_bit);
prefix = (prefix << 1) | static_cast<psnip_uint64_t>(my_bit ^ peer_bit);
const cs_block mine_h =
detail::vdpf::hash_node(level, prefix, seed);
seeds[level] = open_cs<me>(net, mine_h);
}
else if constexpr (dpf_type::is_verifiable)
{
prefix_share = (prefix_share << 1) | my_bit;
const auto tape = net.recv_vec_from<bit_and_pad_msg>(role::p2, net::msg::beaver_tape);
if (tape.size() != hash_level_and_count())
throw std::runtime_error("oblivious hash tape size");
if (sink_ex)
{
seeds[level] = oblivious_cs_on_sink<me>(net, *batch_sink,
sink_ex->round_ref(), batch_index, level, prefix_share,
seed, tape.data());
}
else
{
seeds[level] = oblivious_cs<me>(net, level, prefix_share, seed,
tape.data());
}
}
}
constexpr std::size_t lg = dpf_type::lg_outputs_per_leaf;
constexpr std::size_t n_leaf_ands = leaf_and_slots(lg);
constexpr bool wild = dpf::is_wildcard_v<OutputT>;
const auto lane_pads =
net.recv_from<leaf_pad_msg<n_leaf_ands>>(role::p2, net::msg::beaver_tape);
using concrete_output = dpf::concrete_type_t<OutputT>;
using exterior_node = typename dpf_type::exterior_node;
using leaf_type = dpf::leaf_node_t<exterior_node, concrete_output>;
constexpr bool need_b2a = lg > 0
&& !utils::has_characteristic_two_v<concrete_output>
&& (!wild || !RevealPoint);
std::vector<b2a_pad_msg> b2a_pads;
if constexpr (need_b2a)
{
b2a_pads = net.recv_vec_from<b2a_pad_msg>(role::p2, net::msg::beaver_tape);
if (b2a_pads.size() != n_leaf_ands * 128)
throw std::runtime_error("leaf b2a tape size");
}
// The clear lane is returned only when the caller asked for it.
[[maybe_unused]] unsigned opened_lane_bits = 0;
if constexpr (RevealPoint && wild && lg > 0)
{
for (std::size_t b = 0; b < lg; ++b)
{
const std::uint8_t mine_bit = static_cast<std::uint8_t>(
(to_int(x_share) >> b) & 1u);
const std::uint8_t other = exch(mine_bit);
opened_lane_bits |= static_cast<unsigned>(mine_bit ^ other) << b;
}
}
using leaf_prg = std::conditional_t<Extractable,
detail::vdpf::extractable_leaf_prg<ExteriorPRG>, ExteriorPRG>;
const leaf_type mask_i = dpf::make_leaf_mask_inner<leaf_prg, 0,
std::tuple<concrete_output>>(dpf::unset_lo_2bits(seed));
leaf_type cw_leaf{};
bool sign0 = me == 0
? static_cast<bool>(dpf::get_lo_bit(seed))
: !static_cast<bool>(dpf::get_lo_bit(seed));
auto bit_at = [&](std::size_t b) {
return static_cast<std::uint8_t>((to_int(x_share) >> b) & 1u);
};
auto mux_made = [&](auto && make) {
return mux_leaf_share<me, concrete_output, leaf_type>(net, exch, lg,
std::forward<decltype(make)>(make), bit_at, lane_pads.lanes.data(),
b2a_pads.empty() ? nullptr : b2a_pads.data());
};
if constexpr (!wild && lg >= 1)
{
const leaf_type selected = mux_made([&](unsigned i) {
if constexpr (me == 0)
return dpf::make_naked_leaf<exterior_node>(
static_cast<InputT>(i), beta);
else
return leaf_type{};
});
cw_leaf = cw_from_selected<me, concrete_output>(
net, selected, mask_i, sign0);
}
else if constexpr (me == 0)
{
const leaf_type pad = random_leaf<leaf_type>();
const leaf_type blinded = leaf_add<concrete_output>(mask_i, pad);
net.send_to(role::p1, net::msg::ring_vector, blinded);
const leaf_type back = net.recv_from<leaf_type>(role::p1, net::msg::ring_vector);
const leaf_type D = leaf_add<concrete_output>(back, pad);
if constexpr (wild)
{
cw_leaf = sign0
? dpf::subtract_leaf<concrete_output>(leaf_type{}, D)
: D;
}
else
{
const auto naked = dpf::make_naked_leaf<exterior_node>(InputT{}, beta);
cw_leaf = sign0
? dpf::subtract_leaf<concrete_output>(naked, D)
: dpf::subtract_leaf<concrete_output>(D, naked);
}
net.send_to(role::p1, net::msg::dpf_key, cw_leaf);
}
else
{
const leaf_type blinded = net.recv_from<leaf_type>(role::p0, net::msg::ring_vector);
const leaf_type back =
dpf::subtract_leaf<concrete_output>(mask_i, blinded);
net.send_to(role::p0, net::msg::ring_vector, back);
cw_leaf = net.recv_from<leaf_type>(role::p0, net::msg::dpf_key);
}
if (leaf_seed_out)
*leaf_seed_out = seed;
using wrap_t = std::tuple_element_t<0, typename dpf_type::leaf_wrapper_tuple>;
using held_type = party_key<me, dpf_type>;
constexpr bool wild_lane = dpf::is_wildcard_v<OutputT> && lg > 0;
InputT off{};
if constexpr (dpf::is_wildcard_v<OutputT>)
{
using pad_type = wildcard_leaf_pad_msg<concrete_output, leaf_type>;
const pad_type wild =
net.recv_from<pad_type>(role::p2, net::msg::beaver_tape);
concrete_output coeff{};
if constexpr (utils::has_characteristic_two_v<concrete_output>)
coeff = static_cast<concrete_output>(~std::uint64_t{0});
else
coeff = static_cast<concrete_output>(sign0 ? 1 : -1);
leaf_type blinded_vector{};
if constexpr (lg == 0 || RevealPoint)
{
const leaf_type vector = dpf::make_naked_leaf<exterior_node>(
static_cast<InputT>(opened_lane_bits), coeff);
blinded_vector = leaf_add<concrete_output>(
vector, wild.peer_vector_blind);
}
else
{
const leaf_type share = mux_made([&](unsigned i) {
if constexpr (me == 0)
return dpf::make_naked_leaf<exterior_node>(
static_cast<InputT>(i), coeff);
else
return leaf_type{};
});
const leaf_type mine_msg = leaf_add<concrete_output>(
share, wild.vector_blind);
const leaf_type peer_msg = exch(mine_msg);
blinded_vector = leaf_add<concrete_output>(share, peer_msg);
}
typename wrap_t::beaver_type beaver{wild.output_blind,
wild.vector_blind, blinded_vector};
leaf_type leaf_share =
leaf_add<concrete_output>(wild.zero_leaf, wild.cross);
if constexpr (me == 0)
leaf_share = leaf_add<concrete_output>(leaf_share, cw_leaf);
typename dpf_type::leaf_wrapper_tuple leaves{
wrap_t{leaf_share, beaver}};
return pack_point_held<dpf_type::is_verifiable && RevealPoint,
RevealPoint && wild_lane,
held_type, InputT>(party_key<me, dpf_type>(
dpf_type{root, cws, advice, std::move(leaves), off, {}, {}, 0, 0,
{}, {}, 0, {}, {}, {}, {}, seeds}),
prefix, opened_lane_bits);
}
else
{
typename dpf_type::leaf_wrapper_tuple leaves{wrap_t{cw_leaf}};
return pack_point_held<dpf_type::is_verifiable && RevealPoint,
RevealPoint && wild_lane,
held_type, InputT>(party_key<me, dpf_type>(
dpf_type{root, cws, advice, std::move(leaves), off, {}, {}, 0, 0,
{}, {}, 0, {}, {}, {}, {}, seeds}),
prefix, opened_lane_bits);
}
}
/// @brief One computing party builds its comparison key over the socket DS walk.
/// @details Each party expands and converts only its own seed. Per-level value
/// correction words are opened as masked ring differences; the running `Va`
/// residual remains additively shared between the key holders.
///
/// `PointIsClear` means both callers already hold the same encoded
/// point; this function does not exchange it and returns only the key.
/// `RevealPoint` exchanges the shares and returns the encoded point.
/// Otherwise the value words are a bit-serial mux of the two children
/// and the return value is the key alone. `leq`/`gt` still compute the
/// shared product `1{α = 2^n-1}` from the edge tape, but do not open
/// it or write `cmp.trivial` (that would publish the bit). The full
/// walk already matches the trivial predicate at the domain edge.
template <typename RootDraw, role Self,
typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool PointIsClear = false,
bool RevealPoint = false,
typename InputT,
typename Spec,
typename ...Tags>
HEDLEY_WARN_UNUSED_RESULT
auto comparison_party_impl(RootDraw && root_draw, trio & net, InputT x_share,
const Spec & spec, net::RoundSink * batch_sink, std::size_t batch_index,
const Tags & ...tags)
{
static_assert(Self == role::p0 || Self == role::p1,
"dist DS parties are p0 and p1");
using pair_type =
comparison_pair_t<InteriorPRG, ExteriorPRG, InputT, Spec, Tags...>;
using key_type = typename pair_type::first_type::key_type;
using node = typename key_type::interior_node;
using tree = dpf::tree_traits<InteriorPRG>;
constexpr int me = Self == role::p0 ? 0 : 1;
constexpr std::size_t depth = key_type::depth;
constexpr auto to_int = utils::to_integral_type<InputT>{};
constexpr std::size_t input_bits = utils::bitlength_of_v<InputT>;
static_assert(key_type::num_outputs == 0,
"dist comparison keygen currently builds comparison-only keys");
static_assert(!key_type::cmp_idcf,
"dist comparison keygen does not yet expose incremental prefix CWs");
(void)std::initializer_list<int>{((void)tags, 0)...};
const role peer = Self == role::p0 ? role::p1 : role::p0;
std::optional<net::sink_exchange> sink_ex;
if (batch_sink)
sink_ex.emplace(*batch_sink, batch_index);
auto exch = [&](const auto & mine) {
using T = std::decay_t<decltype(mine)>;
if (sink_ex)
return (*sink_ex)(static_cast<const T &>(mine));
return net.exchange_with(peer, static_cast<const T &>(mine),
net::msg::delta);
};
constexpr bool clear_point = PointIsClear || RevealPoint;
dcf_runtime_spec runtime{};
bool found = false;
detail::incr::collect_cmp_one(runtime, found, spec);
if (!found)
throw std::invalid_argument("dist comparison keygen needs one comparison");
if (runtime.prefix == 0)
runtime.prefix = input_bits;
if (runtime.custom)
throw std::invalid_argument(
"dist comparison keygen currently supports <=64-bit ring payloads");
if constexpr (!clear_point)
{
if (is_paint_kind(runtime.kind))
throw std::invalid_argument(
"oblivious comparison does not support paint kinds");
}
InputT encoded_x{};
if constexpr (PointIsClear)
{
// Both callers pass the same encoded point. The seed walk still
// needs a sharing: p0 holds the point, p1 holds zero.
encoded_x = x_share;
if constexpr (me == 1)
x_share = InputT{};
}
else if constexpr (RevealPoint)
{
if constexpr (me == 0)
utils::flip_msb_if_signed_integral(x_share);
const InputT peer_share = exch(x_share);
encoded_x = utils::xor_input_shares(x_share, peer_share);
}
else if constexpr (me == 0)
{
utils::flip_msb_if_signed_integral(x_share);
}
detail::cmp_meta cmp{};
cmp.nbits = static_cast<int>(runtime.prefix);
cmp.mask = runtime.mask;
cmp.kind = runtime.kind;
cmp.active = true;
cmp.incremental = runtime.incremental;
cmp.block_width = static_cast<int>(key_type::cmp_block);
cmp.tail_bits = static_cast<int>(key_type::cmp_q);
const std::size_t cmp_nbits = runtime.prefix;
const std::uint64_t delta = runtime.beta & cmp.mask;
const std::uint64_t false_value = runtime.false_value & cmp.mask;
unsigned __int128 thresh = 0;
if constexpr (clear_point)
{
const InputT lane =
detail::incr::lane_input(encoded_x, cmp_nbits, input_bits);
thresh = static_cast<unsigned __int128>(to_int(lane));
}
detail::incr::adjust_cmp_threshold(cmp, thresh, cmp_nbits);
node root{};
if constexpr (tree::is_half_tree)
root = net.recv_from<node>(role::p2, net::msg::dpf_key);
else if constexpr (me == 0)
root = dpf::unset_lo_bit(root_draw());
else
root = dpf::set_lo_bit(root_draw());
// leq/gt domain edge: `∧` of the path bits selects the trivial absorb.
// Keep that product shared — opening it (or writing `cmp.trivial`) would
// publish `1{α = 2^n-1}`. The full oblivious walk is already correct at
// the edge; the clear-point path still uses `adjust_cmp_threshold`.
std::uint8_t edge_bit_share = 0;
bool have_edge_bit = false;
if constexpr (!clear_point)
{
const bool edge = cmp.kind == cmp_kind::leq || cmp.kind == cmp_kind::gt;
if (edge && cmp_nbits > 1)
{
const auto pads = net.recv_vec_from<bit_and_pad_msg>(role::p2, net::msg::beaver_tape);
if (pads.size() + 1 != cmp_nbits)
throw std::runtime_error("comparison edge tape size");
std::vector<std::uint8_t> bits(cmp_nbits);
for (std::size_t i = 0; i < cmp_nbits; ++i)
bits[i] = static_cast<std::uint8_t>((to_int(x_share) >> i) & 1u);
edge_bit_share = and_tree<me>(net, bits.data(), cmp_nbits,
pads.data());
have_edge_bit = true;
}
}
node seed = root;
int home = me;
typename key_type::correction_words_array cws{};
typename key_type::correction_advice_array advice{};
typename key_type::correction_seeds_array seeds{};
typename key_type::value_cw_array value_cws{};
typename key_type::value_cw_array value_cw_coeff{};
typename key_type::tail_array tail{};
typename key_type::tail_array tail_coeff{};
typename key_type::prefix_cw_array prefix_cw{};
typename key_type::prefix_cw_array prefix_coeff{};
std::uint64_t cw_last = 0;
std::uint64_t cw_last_coeff = 0;
std::uint64_t va_share = 0;
std::uint64_t va1_share = 0;
psnip_uint64_t prefix = 0;
std::uint64_t prefix_share = 0;
const auto on_path_for = [&](std::uint64_t scale) -> std::uint64_t {
if (!is_paint_kind(cmp.kind))
return (cmp.include_eq ? scale : std::uint64_t{0}) & cmp.mask;
const std::uint64_t unit = detail::dcf_impl::paint_unit(cmp.kind,
cmp_nbits, thresh, cmp_nbits, runtime.length_bits, true,
runtime.paint ? &paint_fn_adapter : nullptr,
runtime.paint ? &runtime.paint : nullptr);
return detail::dcf_impl::scale_plant(unit, scale, cmp.mask);
};
const std::uint64_t on_path = on_path_for(delta);
const std::uint64_t on_path_unit = on_path_for(1ULL);
auto mask = key_type::msb_mask;
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
{
const auto pads = net.recv_from<level_pad_msg>(role::p2, net::msg::beaver_tape);
const std::uint8_t my_bit = static_cast<std::uint8_t>(
!!(to_int(mask) & to_int(x_share)));
std::uint8_t ai = 0;
if constexpr (clear_point)
ai = static_cast<std::uint8_t>(!!(to_int(mask) & to_int(encoded_x)));
const bool is_last = tree::is_last_level(level, depth);
const std::uint8_t adv =
static_cast<std::uint8_t>(dpf::get_lo_bit(seed));
const std::uint8_t t1 =
static_cast<std::uint8_t>(me == 1 ? adv : (adv ^ 1u));
if constexpr (key_type::cmp_block == 0)
{
cmp_level_obliv_msg word_pads{};
if constexpr (!clear_point)
{
if (level < cmp_nbits)
word_pads = net.recv_from<cmp_level_obliv_msg>(role::p2, net::msg::beaver_tape);
}
if (cmp.trivial == cmp_trivial::none && level < cmp_nbits)
{
if constexpr (clear_point)
{
const auto converted = tree::expand_value(seed);
const node & keep = converted[ai];
const node & lose = converted[ai ^ 1u];
const std::uint64_t keep_word =
detail::dcf_impl::convert_node(keep, cmp.mask);
const std::uint64_t lose_word =
detail::dcf_impl::convert_node(lose, cmp.mask);
const auto open_value = [&](std::uint64_t scale,
std::uint64_t & va) {
std::uint64_t plant = 0;
if (is_paint_kind(cmp.kind))
{
const std::uint64_t unit =
detail::dcf_impl::paint_unit(cmp.kind, level,
thresh, cmp_nbits, runtime.length_bits, false,
runtime.paint
? &paint_fn_adapter : nullptr,
runtime.paint ? &runtime.paint : nullptr);
plant = detail::dcf_impl::scale_plant(
unit, scale, cmp.mask);
}
else if (ai != 0)
{
plant = scale & cmp.mask;
}
std::uint64_t local = 0;
if constexpr (me == 0)
{
local = (lose_word + va) & cmp.mask;
}
else
{
local = (lose_word
+ detail::dcf_impl::neg_m(va, cmp.mask)
+ plant) & cmp.mask;
}
const std::uint64_t raw =
open_ring_difference<me>(net, local, cmp.mask);
const std::uint64_t word =
detail::dcf_impl::sgn_m(t1, raw, cmp.mask);
if constexpr (me == 0)
{
va = (va + keep_word
+ detail::dcf_impl::sgn_m(t1, word, cmp.mask))
& cmp.mask;
}
else
{
va = (va
+ detail::dcf_impl::neg_m(keep_word, cmp.mask))
& cmp.mask;
}
return word;
};
const std::uint64_t base = open_value(delta, va_share);
value_cws[level] =
static_cast<typename key_type::value_cw_word>(base);
if constexpr (key_type::cmp_is_wildcard)
{
const std::uint64_t one = open_value(1ULL, va1_share);
value_cw_coeff[level] =
static_cast<typename key_type::value_cw_word>(
(one + detail::dcf_impl::neg_m(base, cmp.mask))
& cmp.mask);
}
}
else
{
const auto converted = tree::expand_value(seed);
const std::uint64_t left =
detail::dcf_impl::convert_node(converted[0], cmp.mask);
const std::uint64_t right =
detail::dcf_impl::convert_node(converted[1], cmp.mask);
const std::uint64_t diff = left + (~right + 1);
const std::uint64_t prod0 = mul_known_word<me>(net, me == 0,
me == 0 ? diff : 0ULL, my_bit, word_pads.bits.data());
const std::uint64_t prod1 = mul_known_word<me>(net, me == 1,
me == 1 ? diff : 0ULL, my_bit, word_pads.bits.data() + 64);
std::uint64_t lose0 = 0;
std::uint64_t keep0 = 0;
std::uint64_t lose1 = 0;
std::uint64_t keep1 = 0;
if constexpr (me == 0)
{
lose0 = right + prod0;
keep0 = left + right - lose0;
lose1 = prod1;
keep1 = 0ULL - prod1;
}
else
{
lose0 = prod0;
keep0 = 0ULL - prod0;
lose1 = right + prod1;
keep1 = left + right - lose1;
}
const std::uint64_t ai_add = b2a_bit<me>(net, my_bit, word_pads.ai);
const auto open_shared = [&](std::uint64_t scale,
std::uint64_t & va) {
const std::uint64_t plant = ai_add * scale;
const std::uint64_t local = me == 0
? (lose0 + va - plant - lose1)
: (lose1 - va + plant - lose0);
const std::uint64_t raw =
open_ring_difference<me>(net, local, cmp.mask);
const std::uint64_t word =
detail::dcf_impl::sgn_m(t1, raw, cmp.mask);
const std::uint64_t signed_word = me == 0
? detail::dcf_impl::sgn_m(t1, word, cmp.mask) : 0ULL;
va = (va + keep0 - keep1 + signed_word) & cmp.mask;
return word;
};
const std::uint64_t base = open_shared(delta, va_share);
value_cws[level] =
static_cast<typename key_type::value_cw_word>(base);
if constexpr (key_type::cmp_is_wildcard)
{
const std::uint64_t one = open_shared(1ULL, va1_share);
value_cw_coeff[level] =
static_cast<typename key_type::value_cw_word>(
(one + detail::dcf_impl::neg_m(base, cmp.mask))
& cmp.mask);
}
}
}
}
const auto child = tree::expand(seed, is_last);
const node L = child[0];
const node R = child[1];
blind_msg blind{};
blind.bit = static_cast<std::uint8_t>(my_bit ^ pads.cw.bit);
blind.msg = detail::ds_xor(detail::ds_xor(L, R), pads.cw.rand);
const blind_msg their_blind = exch(blind);
detail::ds_cw_party mine_pad{pads.cw.rand, pads.cw.gamma, pads.cw.bit};
detail::ds_blind peer_blind{their_blind.msg, their_blind.bit};
share_msg mine_share{};
mine_share.share =
detail::ds_cw_share(L, R, my_bit, mine_pad, peer_blind);
const share_msg their_share = exch(mine_share);
node cw = detail::ds_xor(mine_share.share, their_share.share);
advice_msg am{};
am.aL = static_cast<std::uint8_t>(dpf::get_lo_bit(L) ^ my_bit);
am.aR = static_cast<std::uint8_t>(dpf::get_lo_bit(R) ^ my_bit);
const advice_msg at = exch(am);
std::uint8_t tpack = static_cast<std::uint8_t>(
((am.aR ^ at.aR) << 1) | ((am.aL ^ at.aL ^ 1u) & 1u));
if constexpr (tree::is_half_tree)
{
if (!is_last)
tpack = 0;
}
const bool exp_is_mine = home == me;
node M{};
node base{};
if constexpr (tree::is_half_tree)
{
if (!is_last)
{
M = detail::ds_xor(L, R);
base = (adv & 1u) ? detail::ds_xor(L, cw) : L;
}
else
{
detail::ds_next_terms(L, R, adv, cw, tpack, M, base);
}
}
else
{
detail::ds_next_terms(L, R, adv, cw, tpack, M, base);
}
const simde__m128i cross = beaver_bit_block(exch, true, my_bit, M,
pads.mine, pads.theirs);
const simde__m128i local = detail::ds_gate(my_bit, M);
const simde__m128i exp_gate = exp_is_mine ? local : cross;
const simde__m128i rec_prod = exp_is_mine ? cross : local;
seed = detail::ds_xor(base, detail::ds_xor(exp_gate, rec_prod));
home ^= 1;
cws[level] = cw;
advice[level] = tpack;
if constexpr (clear_point)
prefix = (prefix << 1) | static_cast<psnip_uint64_t>(ai);
else
prefix_share = (prefix_share << 1) | my_bit;
suffix_obliv_msg suffix_pad{};
[[maybe_unused]] bool have_suffix = false;
if constexpr (!clear_point && key_type::cmp_block > 0 && key_type::cmp_q > 0)
{
if (level + 1 == key_type::cmp_h)
{
suffix_pad = net.recv_from<suffix_obliv_msg>(role::p2, net::msg::beaver_tape);
have_suffix = true;
}
}
if constexpr (key_type::cmp_block > 0)
{
if (cmp.trivial == cmp_trivial::none)
{
using sched =
detail::blocked::schedule<key_type::cmp_h,
key_type::cmp_block>;
const std::size_t c = level + 1;
const std::uint8_t next_adv =
static_cast<std::uint8_t>(dpf::get_lo_bit(seed));
const std::uint8_t next_t1 = static_cast<std::uint8_t>(
me == 1 ? next_adv : (next_adv ^ 1u));
if (c <= key_type::cmp_h && sched::contains(c))
{
const std::size_t wi = sched::index(c);
const std::uint64_t rho =
detail::blocked::rho_of<InteriorPRG>(seed, cmp.mask);
const std::uint64_t local_word = me == 0
? rho
: (rho + delta) & cmp.mask;
const std::uint64_t raw =
open_ring_difference<me>(net, local_word, cmp.mask);
value_cws[wi] =
static_cast<typename key_type::value_cw_word>(
detail::dcf_impl::sgn_m(
next_t1, raw, cmp.mask));
if constexpr (key_type::cmp_is_wildcard)
{
value_cw_coeff[wi] =
static_cast<typename key_type::value_cw_word>(
next_t1
? detail::dcf_impl::neg_m(1ULL, cmp.mask)
: (1ULL & cmp.mask));
}
}
if (c == key_type::cmp_h && key_type::cmp_q > 0)
{
std::uint64_t local_masks[4]{};
detail::blocked::suffix_masks<InteriorPRG>(seed,
key_type::cmp_q, cmp.mask, local_masks);
if constexpr (clear_point)
{
const std::uint64_t suffix =
static_cast<std::uint64_t>(thresh)
& ((1ULL << key_type::cmp_q) - 1ULL);
for (std::size_t z = 0; z < key_type::cmp_tail; ++z)
{
const bool pred = cmp.include_eq ? (z <= suffix)
: (z < suffix);
const std::uint64_t local_word = me == 0
? local_masks[z]
: (local_masks[z] + (pred ? delta : 0ULL))
& cmp.mask;
const std::uint64_t raw = open_ring_difference<me>(
net, local_word, cmp.mask);
tail[z] =
static_cast<typename key_type::value_cw_word>(
detail::dcf_impl::sgn_m(
next_t1, raw, cmp.mask));
if constexpr (key_type::cmp_is_wildcard)
{
tail_coeff[z] =
static_cast<typename key_type::value_cw_word>(
pred
? (next_t1
? detail::dcf_impl::neg_m(
1ULL, cmp.mask)
: (1ULL & cmp.mask))
: 0ULL);
}
}
}
else if constexpr (key_type::cmp_q == 2)
{
const suffix_obliv_msg & sp = suffix_pad;
const std::uint8_t s0 = static_cast<std::uint8_t>(
to_int(x_share) & 1u);
const std::uint8_t s1 = static_cast<std::uint8_t>(
(to_int(x_share) >> 1) & 1u);
const std::uint8_t sand = and_bit_share<me>(
net, s0, s1, sp.and_pad);
const std::uint8_t sor = static_cast<std::uint8_t>(
s0 ^ s1 ^ sand);
const std::uint64_t a_or = b2a_bit<me>(net, sor, sp.b2a[0]);
const std::uint64_t a_and = b2a_bit<me>(net, sand, sp.b2a[1]);
const std::uint64_t a_s1 = b2a_bit<me>(net, s1, sp.b2a[2]);
const std::uint64_t one_share = me == 0 ? 1ULL : 0ULL;
for (std::size_t z = 0; z < key_type::cmp_tail; ++z)
{
std::uint64_t pred_share = 0;
std::uint8_t pred_bit = 0;
if (cmp.include_eq)
{
if (z == 0) { pred_share = one_share; pred_bit = me == 0 ? 1 : 0; }
else if (z == 1) { pred_share = a_or; pred_bit = sor; }
else if (z == 2) { pred_share = a_s1; pred_bit = s1; }
else { pred_share = a_and; pred_bit = sand; }
}
else if (z == 0) { pred_share = a_or; pred_bit = sor; }
else if (z == 1) { pred_share = a_s1; pred_bit = s1; }
else if (z == 2) { pred_share = a_and; pred_bit = sand; }
const std::uint64_t plant = pred_share * delta;
const std::uint64_t local_word = me == 0
? (local_masks[z] - plant)
: (local_masks[z] + plant);
const std::uint64_t raw = open_ring_difference<me>(
net, local_word, cmp.mask);
tail[z] =
static_cast<typename key_type::value_cw_word>(
detail::dcf_impl::sgn_m(
next_t1, raw, cmp.mask));
if constexpr (key_type::cmp_is_wildcard)
{
const std::uint8_t pred = open_bit<me>(net, pred_bit);
const std::uint64_t unit = next_t1
? detail::dcf_impl::neg_m(1ULL, cmp.mask)
: (1ULL & cmp.mask);
tail_coeff[z] =
static_cast<typename key_type::value_cw_word>(
pred ? unit : 0ULL);
}
}
}
}
}
}
if constexpr (key_type::is_verifiable)
{
const std::size_t tag = key_type::cmp_block > 0
? (detail::blocked::fold_spine_tag | level) : level;
if constexpr (clear_point)
{
const cs_block mine_h =
detail::vdpf::hash_node(tag, prefix, seed);
seeds[level] = open_cs<me>(net, mine_h);
}
else
{
const auto tape = net.recv_vec_from<bit_and_pad_msg>(role::p2, net::msg::beaver_tape);
if (tape.size() != hash_level_and_count())
throw std::runtime_error("oblivious hash tape size");
if (sink_ex)
{
seeds[level] = oblivious_cs_on_sink<me>(net, *batch_sink,
sink_ex->round_ref(), batch_index, tag, prefix_share,
seed, tape.data());
}
else
{
seeds[level] = oblivious_cs<me>(net, tag, prefix_share, seed,
tape.data());
}
}
}
if constexpr (key_type::cmp_block == 0)
{
if (cmp.trivial == cmp_trivial::none
&& level + 1 == cmp_nbits)
{
const std::uint8_t next_adv =
static_cast<std::uint8_t>(dpf::get_lo_bit(seed));
const std::uint8_t next_t1 = static_cast<std::uint8_t>(
me == 1 ? next_adv : (next_adv ^ 1u));
const std::uint64_t converted =
detail::dcf_impl::convert_node(seed, cmp.mask);
const auto open_final = [&](std::uint64_t path_value,
std::uint64_t va) {
const std::uint64_t local_word = me == 0
? (converted + va) & cmp.mask
: (converted + detail::dcf_impl::neg_m(va, cmp.mask)
+ path_value) & cmp.mask;
const std::uint64_t raw = open_ring_difference<me>(
net, local_word, cmp.mask);
return detail::dcf_impl::sgn_m(
next_t1, raw, cmp.mask);
};
cw_last = open_final(on_path, va_share);
if constexpr (key_type::cmp_is_wildcard)
{
const std::uint64_t one =
open_final(on_path_unit, va1_share);
cw_last_coeff =
(one + detail::dcf_impl::neg_m(cw_last, cmp.mask))
& cmp.mask;
}
}
}
}
const auto zero = net.recv_from<ring_zero_share_msg>(role::p2, net::msg::beaver_tape);
std::uint64_t target = false_value;
if (cmp.trivial == cmp_trivial::always_true)
target = (delta + false_value) & cmp.mask;
else if (cmp.trivial == cmp_trivial::always_false)
target = false_value;
else if (cmp.eval_as_ge)
target = (delta + false_value) & cmp.mask;
// `edge_bit_share` is the shared product `∧ path bits`. Leaving
// `cmp.trivial` unset avoids publishing that bit; the planted value
// words already match the trivial predicate when every path bit is 1.
if (have_edge_bit)
(void)edge_bit_share;
const std::uint64_t cmp_addend =
(zero.share + (me == 1 ? target : 0ULL)) & cmp.mask;
typename key_type::leaf_wrapper_tuple leaves{};
typename key_type::addend_tuple addends{};
InputT off{};
key_type key{root, cws, advice, std::move(leaves), off, cmp, value_cws,
cw_last, cmp_addend, addends, value_cw_coeff, cw_last_coeff, tail,
tail_coeff, prefix_cw, prefix_coeff, seeds};
party_key<me, key_type> held{std::move(key)};
if constexpr (!RevealPoint || PointIsClear)
return held;
else
return opened_point<party_key<me, key_type>, InputT>{
std::move(held), encoded_x};
}
template <typename RootDraw, role Self,
typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool PointIsClear = false,
bool RevealPoint = false,
typename InputT,
typename Spec,
typename ...Tags,
typename = std::enable_if_t<
(sizeof...(Tags) == 0)
|| !std::is_convertible_v<
std::tuple_element_t<0, std::tuple<Tags..., void *>>,
net::RoundSink *>>>
HEDLEY_WARN_UNUSED_RESULT
auto comparison_party(RootDraw && root_draw, trio & net, InputT x_share,
const Spec & spec, const Tags & ...tags)
{
return comparison_party_impl<RootDraw, Self, InteriorPRG, ExteriorPRG,
PointIsClear, RevealPoint, InputT, Spec, Tags...>(
std::forward<RootDraw>(root_draw), net, x_share, spec,
nullptr, 0, tags...);
}
template <typename RootDraw, role Self,
typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool PointIsClear = false,
bool RevealPoint = false,
typename InputT,
typename Spec,
typename ...Tags>
HEDLEY_WARN_UNUSED_RESULT
auto comparison_party(RootDraw && root_draw, trio & net, InputT x_share,
const Spec & spec, net::RoundSink * batch_sink, std::size_t batch_index,
const Tags & ...tags)
{
return comparison_party_impl<RootDraw, Self, InteriorPRG, ExteriorPRG,
PointIsClear, RevealPoint, InputT, Spec, Tags...>(
std::forward<RootDraw>(root_draw), net, x_share, spec,
batch_sink, batch_index, tags...);
}
template <role Self,
typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool PointIsClear = false,
bool RevealPoint = false,
typename InputT,
typename Spec,
typename ...Tags,
typename = std::enable_if_t<
(sizeof...(Tags) == 0)
|| !std::is_convertible_v<
std::tuple_element_t<0, std::tuple<Tags..., void *>>,
net::RoundSink *>>>
HEDLEY_WARN_UNUSED_RESULT
auto comparison_party(trio & net, InputT x_share, const Spec & spec,
const Tags & ...tags)
{
using pair_type =
comparison_pair_t<InteriorPRG, ExteriorPRG, InputT, Spec, Tags...>;
using node = typename pair_type::first_type::key_type::interior_node;
dpf::uniform_node_sampler<node> roots;
return comparison_party_impl<dpf::uniform_node_sampler<node> &, Self,
InteriorPRG, ExteriorPRG, PointIsClear, RevealPoint, InputT, Spec,
Tags...>(roots, net, x_share, spec, nullptr, 0, tags...);
}
template <role Self,
typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool PointIsClear = false,
bool RevealPoint = false,
typename InputT,
typename Spec,
typename ...Tags>
HEDLEY_WARN_UNUSED_RESULT
auto comparison_party(trio & net, InputT x_share, const Spec & spec,
net::RoundSink * batch_sink, std::size_t batch_index, const Tags & ...tags)
{
using pair_type =
comparison_pair_t<InteriorPRG, ExteriorPRG, InputT, Spec, Tags...>;
using node = typename pair_type::first_type::key_type::interior_node;
dpf::uniform_node_sampler<node> roots;
return comparison_party_impl<dpf::uniform_node_sampler<node> &, Self,
InteriorPRG, ExteriorPRG, PointIsClear, RevealPoint, InputT, Spec,
Tags...>(roots, net, x_share, spec, batch_sink, batch_index, tags...);
}
/// @brief p2 deals one bit-Beaver per carry. The final bit has no carry-out.
/// @tparam InputT input domain type
/// @param net the connected trio. This process must be p2.
template <typename InputT, typename Pad = detail::urandom_pad_rng>
void deal_additive_carry(trio & net, Pad && pads = Pad{})
{
constexpr std::size_t nbits = utils::bitlength_of_v<InputT>;
for (std::size_t i = 0; i + 1 < nbits; ++i)
{
const detail::ds_bit_triple t = detail::ds_sample_bit_and(pads);
net.send_to(role::p0, net::msg::beaver_tape, bit_and_pad_msg{t.a0, t.b0, t.c0});
net.send_to(role::p1, net::msg::beaver_tape, bit_and_pad_msg{t.a1, t.b1, t.c1});
}
}
/// @brief This party's XOR share of the sum of two additive shares.
/// @details Matches `split_additive_to_xor`. Each carry is one beaver bit-AND.
/// The raw additive bits are not sent. The sum is not opened.
/// @tparam Self `role::p0` or `role::p1`
/// @tparam InputT input domain type
/// @param net the connected trio
/// @param mine this party's additive share
/// @return XOR share of the unsigned sum, before the signed-MSB flip
template <role Self, typename InputT>
HEDLEY_WARN_UNUSED_RESULT
InputT additive_to_xor_share(trio & net, InputT mine)
{
static_assert(Self == role::p0 || Self == role::p1,
"additive carry parties are p0 and p1");
constexpr int me = Self == role::p0 ? 0 : 1;
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>;
constexpr std::size_t nbits = utils::bitlength_of_v<InputT>;
const role peer = Self == role::p0 ? role::p1 : role::p0;
const U u = static_cast<U>(to_int(mine));
U share = 0;
std::uint8_t carry = 0;
for (std::size_t i = 0; i < nbits; ++i)
{
const std::uint8_t bit = static_cast<std::uint8_t>((u >> i) & U{1});
const std::uint8_t sum_bit = static_cast<std::uint8_t>(bit ^ carry);
share = static_cast<U>(share | (static_cast<U>(sum_bit) << i));
if (i + 1 == nbits)
break;
const bit_and_pad_msg pad =
net.recv_from<bit_and_pad_msg>(role::p2, net::msg::beaver_tape);
const std::uint8_t left = me == 0
? static_cast<std::uint8_t>(bit ^ carry) : carry;
const std::uint8_t right = me == 0
? carry : static_cast<std::uint8_t>(bit ^ carry);
const bit_and_mask_msg mine_mask{
static_cast<std::uint8_t>(left ^ pad.a),
static_cast<std::uint8_t>(right ^ pad.b)};
const bit_and_mask_msg peer_mask =
net.exchange_with(peer, mine_mask, net::msg::delta);
const std::uint8_t d = static_cast<std::uint8_t>(mine_mask.x ^ peer_mask.x);
const std::uint8_t e = static_cast<std::uint8_t>(mine_mask.y ^ peer_mask.y);
const std::uint8_t prod = detail::ds_bit_and_party(
d, e, pad.a, pad.b, pad.c, me == 0);
carry = static_cast<std::uint8_t>(prod ^ carry);
}
return utils::make_from_integral_value<InputT>{}(static_cast<FromI>(share));
}
/// @brief Beaver pads to XOR-add a public constant into an n-bit XOR share.
inline constexpr std::size_t ic_add_and_count(std::size_t nbits) noexcept
{
return nbits > 0 ? nbits - 1 : 0;
}
/// @brief Beaver pads for one XOR-shared greater-than (four ANDs per bit).
inline constexpr std::size_t ic_gt_and_count(std::size_t nbits) noexcept
{
return 4 * nbits;
}
/// @brief Pads for `γ = r − 1` on XOR shares (borrow chain).
inline constexpr std::size_t ic_gamma_and_count(std::size_t nbits) noexcept
{
return ic_add_and_count(nbits);
}
/// @brief Total bit-AND pads for shared `correction_s` (three adds, three
/// comparisons, and `aq == nmask`).
inline constexpr std::size_t ic_correction_and_count(std::size_t nbits) noexcept
{
return 3 * ic_add_and_count(nbits) + 3 * ic_gt_and_count(nbits)
+ ic_add_and_count(nbits);
}
/// @brief XOR-add public `c` into bit shares of `x` (in place). Uses `nbits-1`
/// Beaver ANDs for the carry chain.
template <int Me>
void xor_add_public(trio & net, std::uint8_t * bits, std::size_t nbits,
std::uint64_t c, const bit_and_pad_msg * pads)
{
std::uint8_t carry = 0;
for (std::size_t i = 0; i < nbits; ++i)
{
const std::uint8_t ci = static_cast<std::uint8_t>((c >> i) & 1u);
const std::uint8_t ri = bits[i];
if (ci == 0)
{
bits[i] = static_cast<std::uint8_t>(ri ^ carry);
if (i + 1 < nbits)
carry = and_bit_share<Me>(net, ri, carry, pads[i]);
}
else
{
bits[i] = static_cast<std::uint8_t>(
ri ^ carry ^ (Me == 0 ? 1u : 0u));
if (i + 1 < nbits)
{
const std::uint8_t prod =
and_bit_share<Me>(net, ri, carry, pads[i]);
carry = static_cast<std::uint8_t>(ri ^ carry ^ prod);
}
}
}
}
/// @brief XOR share of `a > b` for two XOR-shared n-bit values.
template <int Me>
std::uint8_t xor_gt(trio & net, const std::uint8_t * a, const std::uint8_t * b,
std::size_t nbits, const bit_and_pad_msg * pads)
{
std::uint8_t gt = 0;
std::uint8_t eq = Me == 0 ? std::uint8_t{1} : std::uint8_t{0};
std::size_t pi = 0;
for (std::size_t i = nbits; i-- > 0;)
{
const std::uint8_t diff = static_cast<std::uint8_t>(a[i] ^ b[i]);
const std::uint8_t not_diff =
static_cast<std::uint8_t>(diff ^ (Me == 0 ? 1u : 0u));
const std::uint8_t not_b =
static_cast<std::uint8_t>(b[i] ^ (Me == 0 ? 1u : 0u));
const std::uint8_t eq_a =
and_bit_share<Me>(net, eq, a[i], pads[pi++]);
const std::uint8_t term =
and_bit_share<Me>(net, eq_a, not_b, pads[pi++]);
const std::uint8_t gt_and =
and_bit_share<Me>(net, gt, term, pads[pi++]);
gt = static_cast<std::uint8_t>(gt ^ term ^ gt_and);
eq = and_bit_share<Me>(net, eq, not_diff, pads[pi++]);
}
return gt;
}
/// @brief XOR share of `(x − 1) & nmask` from XOR bit shares of `x`.
template <int Me>
void xor_sub_one(trio & net, std::uint8_t * bits, std::size_t nbits,
const bit_and_pad_msg * pads)
{
// borrow starts as the constant 1.
std::uint8_t borrow = Me == 0 ? std::uint8_t{1} : std::uint8_t{0};
for (std::size_t i = 0; i < nbits; ++i)
{
const std::uint8_t ri = bits[i];
bits[i] = static_cast<std::uint8_t>(ri ^ borrow);
if (i + 1 < nbits)
{
const std::uint8_t not_r =
static_cast<std::uint8_t>(ri ^ (Me == 0 ? 1u : 0u));
borrow = and_bit_share<Me>(net, not_r, borrow, pads[i]);
}
}
}
/// @brief Pack XOR bit shares into an input word.
template <typename InputT>
InputT bits_to_input(const std::uint8_t * bits, std::size_t nbits)
{
using FromI = typename utils::make_from_integral_value<InputT>::integral_type;
using U = std::make_unsigned_t<FromI>;
U v = 0;
for (std::size_t i = 0; i < nbits; ++i)
v = static_cast<U>(v | (static_cast<U>(bits[i] & 1u) << i));
return utils::make_from_integral_value<InputT>{}(static_cast<FromI>(v));
}
/// @brief Dealer tape for shared `γ = r − 1` and shared `correction_s`.
template <typename Pad = detail::urandom_pad_rng>
inline void deal_ic_mask_pads(trio & net, std::size_t nbits, Pad && pads = Pad{})
{
send_bit_and_vec(net, ic_gamma_and_count(nbits), pads);
send_bit_and_vec(net, ic_correction_and_count(nbits), pads);
send_b2a_vec(net, 4, pads);
}
/// @brief This party's additive share of `correction_s(r,p,q)` in `gmask`, and
/// XOR bit shares of `γ = r − 1`, without opening `r`.
template <int Me, typename InputT>
std::pair<InputT, std::uint64_t> ic_gamma_and_correction(trio & net,
InputT r_share, std::uint64_t p, std::uint64_t q, std::uint64_t nmask,
std::uint64_t gmask)
{
constexpr auto to_int = utils::to_integral_type<InputT>{};
using U = std::make_unsigned_t<
std::decay_t<decltype(to_int(std::declval<InputT>()))>>;
const std::size_t nbits = utils::bitlength_of_v<InputT>;
const U ru = static_cast<U>(to_int(r_share)) & static_cast<U>(nmask);
const auto gamma_pads =
net.recv_vec_from<bit_and_pad_msg>(role::p2, net::msg::beaver_tape);
if (gamma_pads.size() != ic_gamma_and_count(nbits))
throw std::runtime_error("ic gamma tape size");
std::vector<std::uint8_t> r_bits(nbits);
for (std::size_t i = 0; i < nbits; ++i)
r_bits[i] = static_cast<std::uint8_t>((ru >> i) & 1u);
std::vector<std::uint8_t> g_bits = r_bits;
xor_sub_one<Me>(net, g_bits.data(), nbits, gamma_pads.data());
const InputT gamma_share = bits_to_input<InputT>(g_bits.data(), nbits);
const auto corr_pads =
net.recv_vec_from<bit_and_pad_msg>(role::p2, net::msg::beaver_tape);
if (corr_pads.size() != ic_correction_and_count(nbits))
throw std::runtime_error("ic correction tape size");
const auto b2a_pads =
net.recv_vec_from<b2a_pad_msg>(role::p2, net::msg::beaver_tape);
if (b2a_pads.size() != 4)
throw std::runtime_error("ic correction b2a size");
const std::size_t add_n = ic_add_and_count(nbits);
const std::size_t gt_n = ic_gt_and_count(nbits);
const bit_and_pad_msg * pad = corr_pads.data();
auto make_bits = [&](std::uint64_t pub) {
std::vector<std::uint8_t> out = r_bits;
xor_add_public<Me>(net, out.data(), nbits, pub, pad);
pad += add_n;
return out;
};
const auto aq = make_bits(q);
const auto ap = make_bits(p);
const std::uint64_t q0 = (q + 1ULL) & nmask;
const auto aq0 = make_bits(q0);
auto pub_bits = [&](std::uint64_t pub) {
std::vector<std::uint8_t> out(nbits);
for (std::size_t i = 0; i < nbits; ++i)
out[i] = Me == 0
? static_cast<std::uint8_t>((pub >> i) & 1u)
: std::uint8_t{0};
return out;
};
const auto p_bits = pub_bits(p);
const auto q0_bits = pub_bits(q0);
const std::uint8_t b_ap_aq = xor_gt<Me>(net, ap.data(), aq.data(), nbits, pad);
pad += gt_n;
const std::uint8_t b_ap_p = xor_gt<Me>(net, ap.data(), p_bits.data(), nbits, pad);
pad += gt_n;
const std::uint8_t b_aq0_q0 =
xor_gt<Me>(net, aq0.data(), q0_bits.data(), nbits, pad);
pad += gt_n;
const std::uint8_t b_aq_max =
and_tree<Me>(net, aq.data(), nbits, pad);
const std::uint64_t s0 = b2a_bit<Me>(net, b_ap_aq, b2a_pads[0]);
const std::uint64_t s1 = b2a_bit<Me>(net, b_ap_p, b2a_pads[1]);
const std::uint64_t s2 = b2a_bit<Me>(net, b_aq0_q0, b2a_pads[2]);
const std::uint64_t s3 = b2a_bit<Me>(net, b_aq_max, b2a_pads[3]);
// correction_s = s0 - s1 + s2 + s3, embedded in the payload group.
const std::uint64_t cr_share =
(s0 + detail::dcf_impl::neg_m(s1, gmask) + s2 + s3) & gmask;
return {gamma_share, cr_share};
}
} // namespace dist
using net::role;
using net::trio;
/// @brief Run distributed point-key generation and hand each party its key.
/// @tparam InteriorPRG PRG that expands interior nodes
/// @tparam ExteriorPRG PRG that expands the root
/// @tparam InputT input domain type
/// @tparam OutputT payload type
/// @tparam Fn0 callable invoked on p0 with its key
/// @tparam Fn1 callable invoked on p1 with its key
/// @param net the connected trio
/// @param self this process's role
/// @param x0 p0's share of the target. XOR, unless `additive` is set, in
/// which case it is p0's additive share. When `already_encoded` is
/// set this is the opened target.
/// @param x1 p1's share of the target, or zero when `already_encoded` is set
/// @param beta the payload
/// @param on0 p0's continuation
/// @param on1 p1's continuation
/// @param already_encoded the target was opened before keygen
/// @param additive `x0` and `x1` are additive shares. They are converted to
/// XOR shares of the sum by a beaver carry. The sum is not opened.
/// @brief Values a verifiable point-keygen reconstructed.
/// @details `opened_lane` is set only when the payload is a packed wildcard,
/// which has to build its scale vector at the clear lane.
template <typename Point>
struct point_key_opened
{
Point opened_prefix{};
std::optional<unsigned> opened_lane{};
};
/// @return With `RevealPoint`, the tree prefix (`encoded >> lane_bits`) and,
/// for a packed wildcard, the lane. Empty on p2. With the default,
/// nothing: the prefix stays shared and the hash is an AES circuit.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool RevealPoint = false,
typename InputT,
typename OutputT,
typename Fn0,
typename Fn1>
auto dist_with_point_key(
trio & net, role self, InputT x0, InputT x1, OutputT beta, Fn0 && on0,
Fn1 && on1, bool already_encoded = false, bool additive = false)
{
using opened_type = point_key_opened<InputT>;
using dpf_type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, InputT,
OutputT, verifiable>;
constexpr bool wild_lane = RevealPoint && dpf::is_wildcard_v<OutputT>
&& dpf_type::lg_outputs_per_leaf > 0;
if (additive && already_encoded)
throw std::invalid_argument(
"additive DS input is converted to XOR shares; it is not pre-opened");
if (self == role::p2)
{
if (additive)
dist::deal_additive_carry<InputT>(net);
dist::deal_point<InteriorPRG, InputT, OutputT, !RevealPoint>(net);
if constexpr (RevealPoint)
return std::optional<opened_type>(std::nullopt);
else
return;
}
if (additive)
{
if (self == role::p0)
x0 = dist::additive_to_xor_share<role::p0>(net, x0);
else
x1 = dist::additive_to_xor_share<role::p1>(net, x1);
}
if constexpr (RevealPoint)
{
auto finish = [](auto result) {
opened_type out;
out.opened_prefix = result.opened_prefix;
if constexpr (wild_lane)
out.opened_lane = result.opened_lane;
return out;
};
auto slots = net::ds_walk_slot_bytes(dpf_type::depth, false,
dpf_type::lg_outputs_per_leaf);
auto sink = net.batch(1, std::move(slots));
if (self == role::p0)
{
auto result = dist::point_party<InteriorPRG, ExteriorPRG, InputT,
OutputT, role::p0, false, true>(net, x0, beta, already_encoded,
nullptr, dpf::uniform_node_sampler<typename dpf_type::interior_node>{}, sink.get(), 0);
std::forward<Fn0>(on0)(result.dpf_key);
return std::optional<opened_type>(finish(std::move(result)));
}
auto result = dist::point_party<InteriorPRG, ExteriorPRG, InputT,
OutputT, role::p1, false, true>(net, x1, beta, already_encoded,
nullptr, dpf::uniform_node_sampler<typename dpf_type::interior_node>{}, sink.get(), 0);
std::forward<Fn1>(on1)(result.dpf_key);
return std::optional<opened_type>(finish(std::move(result)));
}
else
{
auto slots = net::ds_walk_slot_bytes(dpf_type::depth, true,
dpf_type::lg_outputs_per_leaf);
auto sink = net.batch(1, std::move(slots));
if (self == role::p0)
std::forward<Fn0>(on0)(dist::point_party<InteriorPRG, ExteriorPRG,
InputT, OutputT, role::p0, false, false>(
net, x0, beta, already_encoded, nullptr,
dpf::uniform_node_sampler<typename dpf_type::interior_node>{}, sink.get(), 0));
else
std::forward<Fn1>(on1)(dist::point_party<InteriorPRG, ExteriorPRG,
InputT, OutputT, role::p1, false, false>(
net, x1, beta, already_encoded, nullptr,
dpf::uniform_node_sampler<typename dpf_type::interior_node>{}, sink.get(), 0));
}
}
/// @brief Socket point-key generation using the extractable leaf stretch.
/// @details Not verifiable, so path bits stay shared. Packed lanes are muxed.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename InputT,
typename OutputT,
typename Fn0,
typename Fn1>
void dist_with_extractable_point_key(trio & net, role self, InputT x0,
InputT x1, OutputT beta, Fn0 && on0, Fn1 && on1)
{
static_assert(!dpf::is_wildcard_v<OutputT>,
"extractable socket gen does not reconstruct a wildcard lane");
if (self == role::p2)
{
dist::deal_point<InteriorPRG, InputT, OutputT>(net);
return;
}
if (self == role::p0)
{
on0(dist::point_party<InteriorPRG, ExteriorPRG, InputT, OutputT,
role::p0, true>(net, x0, beta));
}
else
{
on1(dist::point_party<InteriorPRG, ExteriorPRG, InputT, OutputT,
role::p1, true>(net, x1, beta));
}
}
/// @brief Comparison keygen reconstructed nothing. The point stayed shared.
struct suppressed_point {};
/// @brief Run distributed comparison-key generation and hand each holder its key.
/// @tparam Reveal When set, exchange the target and return it. Otherwise the
/// value words are computed from XOR shares and the return is empty.
/// @return Encoded point when `Reveal` is set. Empty on p2. `suppressed_point`
/// when the point stays shared.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool Reveal = false,
typename InputT,
typename Spec,
typename Fn0,
typename Fn1,
typename ...Tags>
auto dist_with_cmp_key(trio & net, role self, InputT x0, InputT x1,
const Spec & spec, Fn0 && on0, Fn1 && on1, const Tags & ...tags)
-> std::conditional_t<Reveal, std::optional<InputT>, suppressed_point>
{
if (self == role::p2)
{
dist::deal_comparison<InteriorPRG, ExteriorPRG, InputT, !Reveal>(
net, spec, tags...);
if constexpr (Reveal)
return std::nullopt;
else
return {};
}
if constexpr (Reveal)
{
using pair_type = dist::comparison_pair_t<InteriorPRG, ExteriorPRG,
InputT, Spec, Tags...>;
constexpr std::size_t depth =
pair_type::first_type::key_type::depth;
auto slots = net::ds_walk_slot_bytes(depth, false);
auto sink = net.batch(1, std::move(slots));
if (self == role::p0)
{
auto result = dist::comparison_party<role::p0, InteriorPRG,
ExteriorPRG, false, true>(net, x0, spec, sink.get(), 0, tags...);
std::forward<Fn0>(on0)(result.dpf_key);
return result.opened_point;
}
auto result = dist::comparison_party<role::p1, InteriorPRG,
ExteriorPRG, false, true>(net, x1, spec, sink.get(), 0, tags...);
std::forward<Fn1>(on1)(result.dpf_key);
return result.opened_point;
}
else
{
using pair_type = dist::comparison_pair_t<InteriorPRG, ExteriorPRG,
InputT, Spec, Tags...>;
constexpr std::size_t depth =
pair_type::first_type::key_type::depth;
auto slots = net::ds_walk_slot_bytes(depth, true);
auto sink = net.batch(1, std::move(slots));
if (self == role::p0)
{
std::forward<Fn0>(on0)(dist::comparison_party<role::p0, InteriorPRG,
ExteriorPRG, false, false>(net, x0, spec, sink.get(), 0, tags...));
}
else
{
std::forward<Fn1>(on1)(dist::comparison_party<role::p1, InteriorPRG,
ExteriorPRG, false, false>(net, x1, spec, sink.get(), 0, tags...));
}
return {};
}
}
/// @brief Build an interval-containment key around a socket-generated DCF spine.
/// @tparam Reveal When set, reconstruct mask `r` and return it (test path).
/// Otherwise `r` stays XOR-shared: shared `correction_s` / absorb and
/// shared `γ = r − 1` feed an oblivious comparison walk.
/// @return Opened `r` when `Reveal` is set. Empty on p2. `suppressed_point`
/// when the mask stays shared.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
bool Reveal = false,
typename InputT,
typename Beta,
typename Fn0,
typename Fn1>
auto dist_with_ic_key(trio & net, role self, InputT r0, InputT r1,
const ic_pack<Beta> & spec, Fn0 && on0, Fn1 && on1)
-> std::conditional_t<Reveal, std::optional<std::decay_t<InputT>>,
suppressed_point>
{
using input_type = std::decay_t<InputT>;
using out_beta = concrete_type_t<Beta>;
detail::ic_impl::check_input<input_type>();
detail::ic_impl::check_bounds<input_type>(spec);
if constexpr (detail::cmp_group_info<Beta>::custom)
{
throw std::invalid_argument(
"dist interval keygen currently supports <=64-bit ring payloads");
if constexpr (Reveal)
return std::nullopt;
else
return {};
}
else
{
const auto inner_spec = detail::ic_impl::inner_lt(spec);
const std::uint64_t gmask =
detail::ic_impl::group_mask_of<Beta>();
const std::uint64_t nmask =
detail::ic_impl::input_mask_of<input_type>();
constexpr std::size_t nbits = utils::bitlength_of_v<input_type>;
if (self == role::p2)
{
if constexpr (!Reveal)
dist::deal_ic_mask_pads(net, nbits);
dist::deal_comparison<InteriorPRG, ExteriorPRG, input_type,
!Reveal>(net, inner_spec);
dist::deal_ring_zero(net, gmask);
if constexpr (Reveal)
dist::deal_ring_zero(net, gmask);
if constexpr (Reveal)
return std::nullopt;
else
return {};
}
const input_type mine = self == role::p0 ? r0 : r1;
std::uint64_t delta = 0;
std::uint64_t fval = 0;
if constexpr (!is_wildcard_v<Beta>)
{
delta = detail::dcf_impl::beta_delta_u64(
spec.if_true, spec.if_false, gmask);
fval = detail::dcf_impl::beta_to_u64_simple(
spec.if_false, gmask);
}
if constexpr (Reveal)
{
const role peer = self == role::p0 ? role::p1 : role::p0;
const input_type theirs =
net.exchange_with(peer, mine, net::msg::delta);
const input_type r = utils::xor_input_shares(mine, theirs);
const std::uint64_t r_bits = detail::ic_impl::bits_of(r);
input_type gamma = detail::ic_impl::gamma_of(r);
utils::flip_msb_if_signed_integral(gamma);
const std::uint64_t cr = detail::ic_impl::correction(
r_bits, spec.lo, spec.hi, nmask, gmask);
const std::uint64_t absorb =
(detail::ic_impl::mul_mask(delta, cr, gmask) + fval) & gmask;
const auto recv_side_shares = [&]() {
const auto z0 =
net.recv_from<dist::ring_zero_share_msg>(role::p2, net::msg::beaver_tape);
const auto z1 =
net.recv_from<dist::ring_zero_share_msg>(role::p2, net::msg::beaver_tape);
const std::uint64_t side = self == role::p1 ? 1ULL : 0ULL;
const std::uint64_t dshare = is_wildcard_v<Beta>
? 0ULL : (z0.share + side * delta) & gmask;
const std::uint64_t cshare = is_wildcard_v<Beta>
? 0ULL : (z1.share + side * absorb) & gmask;
const std::uint64_t dcoeff = is_wildcard_v<Beta>
? (z0.share + side) & gmask : 0ULL;
const std::uint64_t ccoeff = is_wildcard_v<Beta>
? (z1.share + side * cr) & gmask : 0ULL;
return std::array<std::uint64_t, 4>{
dshare, cshare, dcoeff, ccoeff};
};
if (self == role::p0)
{
auto inner = dist::comparison_party<role::p0, InteriorPRG,
ExteriorPRG, true>(net, gamma, inner_spec);
const auto shares = recv_side_shares();
using raw_key = typename decltype(inner)::key_type;
auto key = detail::ic_impl::make_side<0, raw_key, input_type,
out_beta>(std::move(inner), spec.lo, spec.hi, nmask, gmask,
shares[0], shares[1], shares[2], shares[3]);
on0(std::move(key));
}
else
{
auto inner = dist::comparison_party<role::p1, InteriorPRG,
ExteriorPRG, true>(net, gamma, inner_spec);
const auto shares = recv_side_shares();
using raw_key = typename decltype(inner)::key_type;
auto key = detail::ic_impl::make_side<1, raw_key, input_type,
out_beta>(std::move(inner), spec.lo, spec.hi, nmask, gmask,
shares[0], shares[1], shares[2], shares[3]);
on1(std::move(key));
}
return r;
}
else
{
input_type gamma_share{};
std::uint64_t cr_share = 0;
if (self == role::p0)
{
auto got = dist::ic_gamma_and_correction<0, input_type>(net,
mine, spec.lo, spec.hi, nmask, gmask);
gamma_share = got.first;
cr_share = got.second;
}
else
{
auto got = dist::ic_gamma_and_correction<1, input_type>(net,
mine, spec.lo, spec.hi, nmask, gmask);
gamma_share = got.first;
cr_share = got.second;
}
const std::uint64_t absorb_share = is_wildcard_v<Beta>
? 0ULL
: (detail::ic_impl::mul_mask(delta, cr_share, gmask)
+ (self == role::p1 ? fval : 0ULL)) & gmask;
const std::uint64_t cr_coeff_share =
is_wildcard_v<Beta> ? cr_share : 0ULL;
const auto recv_delta = [&]() {
const auto z0 =
net.recv_from<dist::ring_zero_share_msg>(role::p2, net::msg::beaver_tape);
const std::uint64_t side = self == role::p1 ? 1ULL : 0ULL;
const std::uint64_t dshare = is_wildcard_v<Beta>
? 0ULL : (z0.share + side * delta) & gmask;
const std::uint64_t dcoeff = is_wildcard_v<Beta>
? (z0.share + side) & gmask : 0ULL;
return std::array<std::uint64_t, 2>{dshare, dcoeff};
};
if (self == role::p0)
{
auto inner = dist::comparison_party<role::p0, InteriorPRG,
ExteriorPRG, false, false>(net, gamma_share, inner_spec);
const auto d = recv_delta();
using raw_key = typename decltype(inner)::key_type;
auto key = detail::ic_impl::make_side<0, raw_key, input_type,
out_beta>(std::move(inner), spec.lo, spec.hi, nmask, gmask,
d[0], absorb_share, d[1], cr_coeff_share);
on0(std::move(key));
}
else
{
auto inner = dist::comparison_party<role::p1, InteriorPRG,
ExteriorPRG, false, false>(net, gamma_share, inner_spec);
const auto d = recv_delta();
using raw_key = typename decltype(inner)::key_type;
auto key = detail::ic_impl::make_side<1, raw_key, input_type,
out_beta>(std::move(inner), spec.lo, spec.hi, nmask, gmask,
d[0], absorb_share, d[1], cr_coeff_share);
on1(std::move(key));
}
return {};
}
}
}
/// @brief Encoded XOR point the comparison walk reconstructs (MSB flipped on share 0).
template <typename T>
HEDLEY_WARN_UNUSED_RESULT
T encoded_xor_point(T x0, T x1)
{
utils::flip_msb_if_signed_integral(x0);
return utils::xor_input_shares(x0, x1);
}
/// @brief Tree prefix a verifiable point walk reconstructs: `encoded >> lg`.
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename OutputT,
typename T>
HEDLEY_WARN_UNUSED_RESULT
T verifiable_tree_prefix(T x0, T x1, bool already_encoded = false)
{
using dpf_type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, T, OutputT,
verifiable>;
constexpr auto lg = dpf_type::lg_outputs_per_leaf;
T encoded = already_encoded ? x0 : encoded_xor_point(x0, x1);
if constexpr (lg == 0)
return encoded;
constexpr auto to_int = utils::to_integral_type<T>{};
using U = std::make_unsigned_t<std::decay_t<decltype(to_int(std::declval<T>()))>>;
using FromI = typename utils::make_from_integral_value<T>::integral_type;
const U bits = static_cast<U>(to_int(encoded)) >> lg;
return utils::make_from_integral_value<T>{}(static_cast<FromI>(bits));
}
/// @brief p2 reconstructed nothing. p0 and p1 reconstructed `expected`.
template <typename Point>
void require_opened(const std::optional<Point> & got, role self, Point expected,
const char * what = "reconstructed point")
{
if (self == role::p2)
{
if (got.has_value())
throw std::runtime_error(what);
return;
}
if (!got.has_value() || !(*got == expected))
throw std::runtime_error(what);
}
/// @brief Check a verifiable point-keygen's returned prefix and optional lane.
template <typename Point>
void require_tree_prefix(const std::optional<point_key_opened<Point>> & got,
role self, Point expected_prefix,
std::optional<unsigned> expected_lane = std::nullopt,
const char * what = "tree prefix")
{
if (self == role::p2)
{
if (got.has_value())
throw std::runtime_error(what);
return;
}
if (!got.has_value() || !(got->opened_prefix == expected_prefix)
|| got->opened_lane != expected_lane)
throw std::runtime_error(what);
}
/// @brief Low `lg` bits of an encoded point (the packed lane).
template <typename InteriorPRG = dpf::prg::aes128,
typename ExteriorPRG = InteriorPRG,
typename OutputT,
typename T>
HEDLEY_WARN_UNUSED_RESULT
unsigned verifiable_lane(T encoded)
{
using dpf_type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, T, OutputT,
verifiable>;
constexpr auto lg = dpf_type::lg_outputs_per_leaf;
if constexpr (lg == 0)
return 0;
constexpr auto to_int = utils::to_integral_type<T>{};
using U = std::make_unsigned_t<std::decay_t<decltype(to_int(std::declval<T>()))>>;
return static_cast<unsigned>(static_cast<U>(to_int(encoded))
& ((U{1} << lg) - 1));
}
/// @brief (2+1) dealer-side grow: p2 builds `at<1>` then `extend_ds` to `at<2>`
/// and sends each party its key. Exercises library DS grow over the trio.
template <typename InputT, typename OutputT, typename Fn0, typename Fn1>
void dist_with_grow_extend(trio & net, role self, InputT alpha, OutputT beta_new,
Fn0 && on0, Fn1 && on1)
{
using short_keys =
decltype(dpf::make_dpf(alpha, dpf::at<1>(std::uint64_t{1})));
using grown_keys = decltype(dpf::extend(
std::declval<typename short_keys::first_type>(),
std::declval<typename short_keys::second_type>(), alpha,
dpf::at<2>(beta_new)));
using grown0 = typename grown_keys::first_type;
using grown1 = typename grown_keys::second_type;
if (self == role::p2)
{
auto [a0, a1] = dpf::make_dpf(alpha, dpf::at<1>(std::uint64_t{1}));
auto m0 = dpf::make_basic_path_memoizer(a0);
auto m1 = dpf::make_basic_path_memoizer(a1);
(void)dpf::eval_point(dpf::out<0>, a0, alpha, m0);
(void)dpf::eval_point(dpf::out<0>, a1, alpha, m1);
struct pad_t
{
simde__m128i block() { return dpf::uniform_sample<simde__m128i>(); }
std::uint8_t bit()
{
return static_cast<std::uint8_t>(
dpf::uniform_sample<unsigned char>() & 1u);
}
} pads{};
dpf::local_cw_protocol<pad_t> proto{pads};
auto [n0, n1] = dpf::extend_ds(a0, a1, m0, m1, alpha, InputT{}, proto,
dpf::at<2>(beta_new));
send_key(net, role::p0, n0);
send_key(net, role::p1, n1);
return;
}
if (self == role::p0)
std::forward<Fn0>(on0)(recv_key<grown0>(net, role::p2));
else
std::forward<Fn1>(on1)(recv_key<grown1>(net, role::p2));
}
/// @brief Two-party grow without p2: p0 runs joint `extend_ds` and delivers p1's key.
template <typename InputT, typename OutputT, typename Fn0, typename Fn1>
void dist_with_grow_extend_iknp(trio & net, role self, InputT alpha,
OutputT beta_new, Fn0 && on0, Fn1 && on1)
{
if (self == role::p2)
throw std::invalid_argument("iknp grow has no dealer");
using short_keys =
decltype(dpf::make_dpf(alpha, dpf::at<1>(std::uint64_t{1})));
using grown_keys = decltype(dpf::extend(
std::declval<typename short_keys::first_type>(),
std::declval<typename short_keys::second_type>(), alpha,
dpf::at<2>(beta_new)));
using grown0 = typename grown_keys::first_type;
using grown1 = typename grown_keys::second_type;
if (self == role::p0)
{
auto [a0, a1] = dpf::make_dpf(alpha, dpf::at<1>(std::uint64_t{1}));
auto m0 = dpf::make_basic_path_memoizer(a0);
auto m1 = dpf::make_basic_path_memoizer(a1);
(void)dpf::eval_point(dpf::out<0>, a0, alpha, m0);
(void)dpf::eval_point(dpf::out<0>, a1, alpha, m1);
struct pad_t
{
simde__m128i block() { return dpf::uniform_sample<simde__m128i>(); }
std::uint8_t bit()
{
return static_cast<std::uint8_t>(
dpf::uniform_sample<unsigned char>() & 1u);
}
} pads{};
dpf::local_cw_protocol<pad_t> proto{pads};
auto [n0, n1] = dpf::extend_ds(a0, a1, m0, m1, alpha, InputT{}, proto,
dpf::at<2>(beta_new));
send_key(net, role::p1, n1);
std::forward<Fn0>(on0)(std::move(n0));
}
else
{
std::forward<Fn1>(on1)(recv_key<grown1>(net, role::p0));
}
}
} // namespace party
} // namespace dpf
#endif // LIBDPF_PARTY_DIST_DS_HPP__