Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
2726 lines
107 KiB
C++
2726 lines
107 KiB
C++
/// @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__
|