libdpf/include/dpf/compose.hpp

4294 lines
162 KiB
C++
Raw Permalink Normal View History

/// @file dpf/compose.hpp
/// @brief Protocol composition: domain-tagged strands on a RoundSink.
/// @details Values carry a share domain (`fss`, `a`, `b`, `rss`, `y`) matching
/// the letters in `secret_share.hpp`. Recording a composition interns
/// identical blinds, expansions, computes, and local conversions so
/// sub-strands share work. Independent opens pack into one RoundSink
/// round; dependency chains become successive waves. ABY2.0 sessions
/// contribute their δ barriers to those waves. Classic Beaver triples
/// reuse an ABY wire's λ when the factor is already in the session.
/// RSS products use `rss_mul` then `rss_from_y` for the neighbor
/// refresh. Express/Sabre audits use `*_fused`; BGI early-stop uses
/// `fss_point_early_stop`; Poplar checkpoints use `level_walk_prefixes`.
/// Doerner–Shelat walks use `level_walk_ds`; idpf_agg uses
/// `step_adaptive_prefix` / `retain_adaptive_prefix`; keyword PIR uses
/// `multipoint_fan`. Multi-lane ABY is `aby_lane`. Prepaid expands are
/// `defer_expand` / `leaf_later_walk`. Party-count changes are never
/// implied: use an explicit `reshare`.
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details.
#ifndef LIBDPF_INCLUDE_DPF_COMPOSE_HPP__
#define LIBDPF_INCLUDE_DPF_COMPOSE_HPP__
#include <algorithm>
#include <atomic>
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <functional>
#include <map>
#include <memory>
#include <mutex>
#include <stdexcept>
#include <string>
#include <thread>
#include <typeindex>
#include <unordered_map>
#include <utility>
#include <vector>
#include "simde/simde/x86/avx.h"
#include "hedley/hedley.h"
#include "dpf/net/asio_ns.hpp"
#include "dpf/beaver.hpp"
#include "dpf/bit_inject.hpp"
#include "dpf/edabit.hpp"
#include "dpf/net/edge_mesh.hpp"
#include "dpf/net/io_pool.hpp"
#include "dpf/net/stream_array.hpp"
#include "dpf/net/stream_edge_mesh.hpp"
#include "dpf/trunc.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/net/sink_exchange.hpp"
#include "dpf/prg_aes.hpp"
#include "dpf/protocol.hpp"
#include "dpf/rss_seed.hpp"
#include "dpf/secret_share.hpp"
namespace dpf
{
namespace protocol
{
struct drive_options;
struct beaver_host;
enum class domain : unsigned char
{
fss = 0,
a = 1,
b = 2,
rss = 3,
y = 4,
bin = 5, ///< (2,2) XOR bit shares (packed)
bin_rss = 6 ///< (2,3) replicated bits
};
enum class effect : unsigned char
{
input = 0,
blind = 1,
expand = 2,
compute = 3,
exchange = 4,
convert = 5
};
enum class op_flags : unsigned
{
none = 0,
commutative = 1u
};
inline constexpr op_flags operator|(op_flags a, op_flags b) noexcept
{
return static_cast<op_flags>(
static_cast<unsigned>(a) | static_cast<unsigned>(b));
}
HEDLEY_CONST
HEDLEY_NO_THROW
inline constexpr bool has_flag(op_flags f, op_flags bit) noexcept
{
return (static_cast<unsigned>(f) & static_cast<unsigned>(bit)) != 0;
}
namespace opcodes
{
constexpr std::uint32_t conv_a2b = 1;
constexpr std::uint32_t conv_b2a = 2;
constexpr std::uint32_t conv_a2fss = 3;
constexpr std::uint32_t conv_fss2a = 4;
constexpr std::uint32_t conv_b2fss = 5;
constexpr std::uint32_t conv_fss2b = 6;
constexpr std::uint32_t conv_rss2y = 7;
constexpr std::uint32_t conv_y2rss = 8;
constexpr std::uint32_t reshare_mask = 9; ///< Local masked reshare (no open)
constexpr std::uint32_t rss_mul = 100;
constexpr std::uint32_t fss_expand = 200;
constexpr std::uint32_t fss_expand_pair = 202; ///< Both children in one PRG stretch
constexpr std::uint32_t fss_step = 201;
constexpr std::uint32_t fss_fuse_trailer = 203; ///< Pack step‖sketch into one slot
constexpr std::uint32_t fss_early_stop_pack = 204; ///< BGI Remark 3.4 leaf pack (local)
constexpr std::uint32_t fss_prefix_emit = 205; ///< Prefix share at a walk checkpoint
constexpr std::uint32_t fss_defer_expand = 206; ///< Prepaid full expand (no CW open)
constexpr std::uint32_t fss_rotate = 207; ///< Local rotate of a share buffer
constexpr std::uint32_t fss_leaf_later = 208; ///< Path walk; leaf CW applied later
constexpr std::uint32_t fss_apply_leaf = 209; ///< `buf += F * control` (local)
constexpr std::uint32_t fss_prefix_pair = 210; ///< Pack L‖R child opens (idpf_agg)
constexpr std::uint32_t fss_prefix_retain = 211; ///< Keep one child after open
constexpr std::uint32_t fss_answer_pack = 212; ///< Pack multipoint / bucket answers
constexpr std::uint32_t fss_fuse_segment = 213; ///< Slice one fused open (level = offset)
constexpr std::uint32_t ds_blind = 220; ///< Doerner–Shelat blind open
constexpr std::uint32_t ds_cw = 221; ///< Correction-word share open
constexpr std::uint32_t ds_advice = 222; ///< Advice-bit open
constexpr std::uint32_t ds_and = 223; ///< Bit-block AND round 1
constexpr std::uint32_t ds_and2 = 227; ///< Bit-block AND round 2
constexpr std::uint32_t ds_oh = 224; ///< Oblivious-hash AND layer open
constexpr std::uint32_t ds_leaf_mux = 225; ///< Packed-leaf mux open
constexpr std::uint32_t dealer_pad = 226; ///< p2 → parties pad / tape delivery
constexpr std::uint32_t dealer_zero = 228; ///< p2 samples a sharing of zero
constexpr std::uint32_t client_upload = 229; ///< Client sends one query to each server
constexpr std::uint32_t client_answer = 230; ///< Each server returns one answer share
constexpr std::uint32_t beaver_delta = 300;
constexpr std::uint32_t bin_a2b = 400;
constexpr std::uint32_t bin_b2a = 401;
constexpr std::uint32_t bin_inject = 402;
constexpr std::uint32_t bin_and = 403;
constexpr std::uint32_t bin_rss_and = 404;
constexpr std::uint32_t bin_rss_inject = 405;
constexpr std::uint32_t rss_zero = 406; ///< Pairwise-seed zero mask (local)
constexpr std::uint32_t trunc_prob = 410;
constexpr std::uint32_t trunc_exact = 411;
constexpr std::uint32_t mul_trunc = 412;
constexpr std::uint32_t share_cmp = 420;
constexpr std::uint32_t share_mux = 421;
constexpr std::uint32_t declassify = 430;
constexpr std::uint32_t share_input = 431;
constexpr std::uint32_t matmul_open = 440;
constexpr std::uint32_t shuffle_send = 441; ///< One hidden-shuffle pass (`shuffle::shuffle_hidden_pass`)
constexpr std::uint32_t bench_emit = 450; ///< Local half of a bench cell (aux = cell id)
constexpr std::uint32_t bench_apply = 451; ///< Evaluator half, after the payload arrives
constexpr std::uint32_t bench_wire = 452; ///< 2-party payload on the peer edge
constexpr std::uint32_t bench_ring = 453; ///< 3-party payload on the RSS ring
constexpr std::uint32_t user_base = 1000;
} // namespace opcodes
struct node
{
std::uint32_t id = 0;
friend bool operator==(node a, node b) noexcept { return a.id == b.id; }
friend bool operator!=(node a, node b) noexcept { return !(a == b); }
};
struct block_span
{
std::uint8_t * data = nullptr;
std::size_t lanes = 0;
std::size_t value_bytes = 0;
std::uint8_t * at(std::size_t lane) const
{
return data + lane * value_bytes;
}
simde__m128i load128(std::size_t lane) const
{
if (value_bytes != 16)
throw std::logic_error("block_span::load128 requires 16-byte values");
simde__m128i v;
std::memcpy(&v, at(lane), 16);
return v;
}
void store128(std::size_t lane, simde__m128i v) const
{
if (value_bytes != 16)
throw std::logic_error("block_span::store128 requires 16-byte values");
std::memcpy(at(lane), &v, 16);
}
};
using kernel_fn = std::function<void(std::uint32_t opcode,
const std::vector<node> & nodes, const std::vector<block_span> & inputs,
block_span output, std::size_t sink_lanes)>;
struct simd_group
{
std::uint32_t opcode = 0;
effect kind = effect::compute;
std::vector<node> nodes;
};
struct wave_info
{
std::size_t index = 0;
std::vector<simd_group> groups;
std::vector<node> exchanges;
std::size_t slot_bytes = 0;
};
class plan
{
public:
/// @brief Full dependency waves (includes trailing compute-only waves).
HEDLEY_NO_THROW
std::size_t waves() const noexcept { return waves_.size(); }
/// @brief RoundSink flushes — equal to `exchange_waves()`, not `waves()`.
/// @details The online plan and the sink share this count: no empty rounds.
HEDLEY_NO_THROW
std::size_t rounds() const noexcept { return exchange_waves(); }
HEDLEY_WARN_UNUSED_RESULT
std::size_t slot_bytes(std::uint16_t round) const
{
std::size_t seen = 0;
for (const auto & w : waves_)
{
if (w.exchanges.empty())
continue;
if (seen == round)
return w.slot_bytes;
++seen;
}
throw std::out_of_range("plan::slot_bytes");
}
/// @brief Slot widths for each RoundSink round (exchange-bearing only).
HEDLEY_WARN_UNUSED_RESULT
std::vector<std::size_t> slot_bytes_all() const
{
std::vector<std::size_t> out;
out.reserve(waves_.size());
for (const auto & w : waves_)
if (!w.exchanges.empty())
out.push_back(w.slot_bytes);
return out;
}
/// @brief Tail of `slot_bytes_all` starting at exchange-wave `from`.
HEDLEY_WARN_UNUSED_RESULT
std::vector<std::size_t> slot_bytes_from(std::size_t from_exchange_wave) const
{
auto all = slot_bytes_all();
if (from_exchange_wave >= all.size())
return {};
return {all.begin()
+ static_cast<std::ptrdiff_t>(from_exchange_wave),
all.end()};
}
HEDLEY_WARN_UNUSED_RESULT
const wave_info & wave(std::size_t i) const
{
if (i >= waves_.size())
throw std::out_of_range("plan::wave");
return waves_[i];
}
HEDLEY_WARN_UNUSED_RESULT
std::size_t wave_of(node n) const
{
return node_wave_.at(n.id);
}
HEDLEY_WARN_UNUSED_RESULT
std::size_t effect_count(effect e) const
{
auto it = effect_counts_.find(e);
return it == effect_counts_.end() ? 0 : it->second;
}
/// @brief Waves that actually flush on the wire (alias of `rounds()`).
HEDLEY_NO_THROW
HEDLEY_WARN_UNUSED_RESULT
std::size_t exchange_waves() const noexcept
{
std::size_t n = 0;
for (const auto & w : waves_)
if (!w.exchanges.empty())
++n;
return n;
}
HEDLEY_WARN_UNUSED_RESULT
std::size_t conversion_count(domain from, domain to) const
{
auto it = conversion_counts_.find({from, to});
return it == conversion_counts_.end() ? 0 : it->second;
}
HEDLEY_NO_THROW
HEDLEY_WARN_UNUSED_RESULT
const std::vector<node> & nodes() const noexcept { return all_nodes_; }
HEDLEY_WARN_UNUSED_RESULT
std::size_t value_bytes_of(std::uint32_t id) const
{
return value_bytes_.at(id);
}
HEDLEY_WARN_UNUSED_RESULT
domain domain_of(std::uint32_t id) const
{
return domains_.at(id);
}
HEDLEY_WARN_UNUSED_RESULT
const std::vector<std::uint32_t> & inputs_of(std::uint32_t id) const
{
return inputs_.at(id);
}
HEDLEY_WARN_UNUSED_RESULT
std::uint32_t opcode_of(std::uint32_t id) const
{
return opcodes_.at(id);
}
HEDLEY_WARN_UNUSED_RESULT
bool alias_of(std::uint32_t id) const
{
return alias_bytes_.at(id);
}
/// @brief Intern key `level` (ABY session index for beaver barriers).
HEDLEY_WARN_UNUSED_RESULT
std::uint16_t level_of(std::uint32_t id) const
{
return levels_.at(id);
}
/// @brief Byte offset (fuse segments) or other per-node aux.
HEDLEY_WARN_UNUSED_RESULT
std::uint32_t aux_of(std::uint32_t id) const
{
return aux_.at(id);
}
/// @brief Online exchanges plus a setup cost (IKNP `setup_rounds`, dealer tape, …).
HEDLEY_WARN_UNUSED_RESULT
std::size_t rounds_including(std::size_t setup_rounds) const
{
return rounds() + setup_rounds;
}
HEDLEY_WARN_UNUSED_RESULT
effect effect_of(std::uint32_t id) const
{
return effects_.at(id);
}
/// @brief Sum of exchange slot bytes tagged setup (pad / IKNP).
HEDLEY_NO_THROW
std::size_t setup_bytes() const noexcept { return setup_bytes_; }
/// @brief Sum of exchange slot bytes tagged online.
HEDLEY_NO_THROW
std::size_t online_bytes() const noexcept { return online_bytes_; }
private:
friend class composer;
friend void drive(const plan &, net::RoundSink &,
std::vector<std::vector<std::uint8_t>> &,
const std::map<std::uint32_t, kernel_fn> &, std::size_t,
const struct drive_options &);
std::vector<wave_info> waves_;
std::map<effect, std::size_t> effect_counts_;
std::map<std::pair<domain, domain>, std::size_t> conversion_counts_;
std::vector<node> all_nodes_;
std::vector<std::size_t> node_wave_;
std::vector<std::size_t> value_bytes_;
std::vector<domain> domains_;
std::vector<effect> effects_;
std::vector<std::uint32_t> opcodes_;
std::vector<std::uint16_t> levels_;
std::vector<std::uint32_t> aux_;
std::vector<std::vector<std::uint32_t>> inputs_;
std::vector<bool> alias_bytes_;
std::size_t setup_bytes_ = 0;
std::size_t online_bytes_ = 0;
};
namespace detail
{
HEDLEY_CONST
HEDLEY_NO_THROW
inline bool is_two_party(domain d) noexcept
{
return d == domain::fss || d == domain::a || d == domain::b
|| d == domain::bin;
}
HEDLEY_CONST
HEDLEY_NO_THROW
inline bool is_boolean(domain d) noexcept
{
return d == domain::bin || d == domain::bin_rss;
}
HEDLEY_CONST
HEDLEY_NO_THROW
inline bool local_convertible(domain from, domain to) noexcept
{
if (from == to)
return true;
// Boolean domains never locally cast to/from arithmetic or FSS.
if (is_boolean(from) || is_boolean(to))
return false;
if (is_two_party(from) && is_two_party(to)
&& from != domain::bin && to != domain::bin)
return true;
if ((from == domain::rss && to == domain::y)
|| (from == domain::y && to == domain::rss))
return true;
return false;
}
HEDLEY_WARN_UNUSED_RESULT
inline std::uint32_t conversion_opcode(domain from, domain to)
{
if (from == domain::a && to == domain::b)
return opcodes::conv_a2b;
if (from == domain::b && to == domain::a)
return opcodes::conv_b2a;
if (from == domain::a && to == domain::fss)
return opcodes::conv_a2fss;
if (from == domain::fss && to == domain::a)
return opcodes::conv_fss2a;
if (from == domain::b && to == domain::fss)
return opcodes::conv_b2fss;
if (from == domain::fss && to == domain::b)
return opcodes::conv_fss2b;
if (from == domain::rss && to == domain::y)
return opcodes::conv_rss2y;
if (from == domain::y && to == domain::rss)
return opcodes::conv_y2rss;
throw std::invalid_argument("compose: no local conversion between domains");
}
HEDLEY_CONST
HEDLEY_NO_THROW
inline bool bit_preserving_cast(domain from, domain to, std::size_t party) noexcept
{
if (from == to)
return true;
if ((from == domain::b && to == domain::fss)
|| (from == domain::fss && to == domain::b))
return true;
if (party == 0
&& ((from == domain::a && to == domain::b)
|| (from == domain::b && to == domain::a)
|| (from == domain::a && to == domain::fss)
|| (from == domain::fss && to == domain::a)))
return true;
return false;
}
struct intern_key
{
effect kind{};
domain dom{};
std::uint32_t opcode = 0;
std::uint16_t level = 0;
std::uint8_t child = 0;
std::uint32_t aux = 0;
std::vector<std::uint32_t> inputs;
friend bool operator==(const intern_key & a, const intern_key & b) noexcept
{
return a.kind == b.kind && a.dom == b.dom && a.opcode == b.opcode
&& a.level == b.level && a.child == b.child && a.aux == b.aux
&& a.inputs == b.inputs;
}
};
struct intern_key_hash
{
std::size_t operator()(const intern_key & k) const noexcept
{
std::size_t h = static_cast<std::size_t>(k.kind);
h = h * 131u + static_cast<std::size_t>(k.dom);
h = h * 131u + k.opcode;
h = h * 131u + k.level;
h = h * 131u + k.child;
h = h * 131u + k.aux;
for (auto id : k.inputs)
h = h * 131u + id;
return h;
}
};
struct aby_slot_key
{
std::type_index type;
std::size_t lane = 0;
bool operator<(const aby_slot_key & o) const noexcept
{
if (type < o.type)
return true;
if (o.type < type)
return false;
return lane < o.lane;
}
bool operator==(const aby_slot_key & o) const noexcept
{
return type == o.type && lane == o.lane;
}
};
struct aby_holder_base
{
virtual ~aby_holder_base() = default;
virtual std::type_index type() const noexcept = 0;
virtual std::size_t wire_count() const = 0;
virtual std::size_t barrier_count() const = 0;
virtual std::size_t ring_bytes() const = 0;
};
template <typename Ring>
struct aby_holder final : aby_holder_base
{
beavers::session<Ring> session;
std::type_index type() const noexcept override
{
return std::type_index(typeid(Ring));
}
std::size_t wire_count() const override { return session.wire_count(); }
std::size_t barrier_count() const override
{
return session.exchange_barriers().size();
}
std::size_t ring_bytes() const override { return sizeof(Ring); }
};
inline void run_builtin_convert(std::uint32_t opcode, std::size_t party,
const std::uint8_t * in, std::size_t in_bytes, std::uint8_t * out,
std::size_t out_bytes)
{
if (opcode == opcodes::conv_b2fss || opcode == opcodes::conv_fss2b
|| (party == 0
&& (opcode == opcodes::conv_a2b || opcode == opcodes::conv_b2a
|| opcode == opcodes::conv_a2fss
|| opcode == opcodes::conv_fss2a)))
{
if (in_bytes != out_bytes)
throw std::logic_error("compose convert size");
std::memcpy(out, in, out_bytes);
return;
}
if (opcode == opcodes::conv_a2b || opcode == opcodes::conv_b2a
|| opcode == opcodes::conv_a2fss || opcode == opcodes::conv_fss2a)
{
if (in_bytes != out_bytes)
throw std::logic_error("compose convert size");
std::vector<std::uint8_t> tmp(in, in + in_bytes);
unsigned borrow = 1;
for (std::size_t i = 0; i < tmp.size(); ++i)
{
const unsigned x =
static_cast<unsigned>(static_cast<std::uint8_t>(~tmp[i])) + borrow;
tmp[i] = static_cast<std::uint8_t>(x);
borrow = x >> 8;
}
std::memcpy(out, tmp.data(), out_bytes);
return;
}
if (opcode == opcodes::conv_rss2y)
{
const std::size_t n = std::min(out_bytes, in_bytes);
std::memcpy(out, in, n);
return;
}
if (opcode == opcodes::conv_y2rss)
{
if (in_bytes != out_bytes)
throw std::logic_error("compose y2rss size");
std::memcpy(out, in, out_bytes);
return;
}
throw std::invalid_argument("compose: unknown convert opcode");
}
} // namespace detail
class composer
{
public:
explicit composer(std::size_t party = 0) : party_(party)
{
// Party id is free-form: 2-party, 3-party RSS, and N-party meshes all
// use the same composer; topology is chosen by the drive/mesh binding.
}
HEDLEY_NO_THROW
HEDLEY_WARN_UNUSED_RESULT
std::size_t party() const noexcept { return party_; }
HEDLEY_WARN_UNUSED_RESULT
node input(domain dom, std::size_t value_bytes)
{
// Inputs are never interned: each call is a distinct free value.
node_info ni;
ni.kind = effect::input;
ni.dom = dom;
ni.value_bytes = value_bytes;
const auto id = static_cast<std::uint32_t>(info_.size());
info_.push_back(std::move(ni));
return node{id};
}
node blind(node src, std::uint32_t opcode)
{
check(src);
return emplace(effect::blind, info_[src.id].dom, opcode, {src.id},
info_[src.id].value_bytes, 0, 0, op_flags::none, false, domain::a);
}
node expand(node src, std::uint16_t level, std::uint8_t child,
std::uint32_t opcode, std::size_t value_bytes)
{
check(src);
return emplace(effect::expand, info_[src.id].dom, opcode, {src.id},
value_bytes, level, child, op_flags::none, false, domain::a);
}
node compute(std::uint32_t opcode, std::vector<node> inputs, domain dom,
std::size_t value_bytes, op_flags flags = op_flags::none)
{
std::vector<std::uint32_t> ids;
ids.reserve(inputs.size());
for (auto n : inputs)
{
check(n);
ids.push_back(n.id);
}
if (has_flag(flags, op_flags::commutative))
std::sort(ids.begin(), ids.end());
return emplace(effect::compute, dom, opcode, std::move(ids),
value_bytes, 0, 0, flags, false, domain::a);
}
node exchange(node payload)
{
check(payload);
return emplace(effect::exchange, info_[payload.id].dom, 0, {payload.id},
info_[payload.id].value_bytes, 0, 0, op_flags::none, false,
domain::a);
}
node as(node src, domain to)
{
check(src);
const domain from = info_[src.id].dom;
if (from == to)
return src;
if (!detail::local_convertible(from, to))
{
throw std::invalid_argument(
"compose::as cannot change party count; use reshare");
}
if (from == domain::y && to == domain::rss)
{
throw std::invalid_argument(
"compose::as(y, rss) needs own and next; use y2rss(own, next)");
}
if (info_[src.id].kind == effect::convert)
{
if (info_[src.id].from_domain == to && !info_[src.id].inputs.empty())
return node{info_[src.id].inputs[0]};
}
const auto opcode = detail::conversion_opcode(from, to);
const bool alias = detail::bit_preserving_cast(from, to, party_);
std::size_t bytes = info_[src.id].value_bytes;
if (from == domain::rss && to == domain::y && bytes >= 2)
bytes = bytes / 2;
return emplace(effect::convert, to, opcode, {src.id}, bytes, 0, 0,
op_flags::none, alias, from);
}
node y2rss(node own, node next)
{
check(own);
check(next);
if (info_[own.id].dom != domain::y || info_[next.id].dom != domain::y)
throw std::invalid_argument("y2rss requires y-domain nodes");
return emplace(effect::convert, domain::rss, opcodes::conv_y2rss,
{own.id, next.id}, info_[own.id].value_bytes * 2, 0, 0,
op_flags::none, false, domain::y);
}
/// @brief Re-share a local RSS product factor (`y`) back to replicated form.
/// @details Hand 3PC schedules exchange each party's `y` component with a
/// neighbor in one round, then `y2rss(own, next)`. Independent
/// refreshes pack into the same wave. Prefer this over `reshare`:
/// a reconstructing `exchange` would open the secret.
HEDLEY_WARN_UNUSED_RESULT
node rss_from_y(node y_own)
{
check(y_own);
if (info_[y_own.id].dom != domain::y)
throw std::invalid_argument("rss_from_y requires a y-domain node");
node peer = exchange(y_own);
return y2rss(y_own, peer);
}
/// @brief Local `rss_mul` then one-round neighbor refresh to RSS.
HEDLEY_WARN_UNUSED_RESULT
node rss_product_replicated(node x, node y)
{
return rss_from_y(rss_product(x, y));
}
node reshare(node src, domain to)
{
check(src);
if (info_[src.id].dom == domain::y && to == domain::rss)
return rss_from_y(src);
if (to == domain::rss && detail::is_two_party(info_[src.id].dom))
{
// Local y-component (p2's share stays 0) then one neighbor refresh.
// Does not reconstruct the secret.
node y = emplace(effect::convert, domain::y, opcodes::conv_rss2y,
{src.id}, info_[src.id].value_bytes, 0, 0, op_flags::none,
true, info_[src.id].dom);
return rss_from_y(y);
}
if (detail::local_convertible(info_[src.id].dom, to)
&& !(info_[src.id].dom == domain::y && to == domain::rss))
return as(src, to);
throw std::invalid_argument(
"reshare would open the secret; use rss_from_y or reshare_fresh");
}
/// @brief p2 samples a sharing of zero and delivers one share to each party.
HEDLEY_WARN_UNUSED_RESULT
node dealer_zero_mask(domain to, std::size_t bytes)
{
if (bytes == 0)
throw std::invalid_argument("dealer_zero_mask bytes must be > 0");
node payload = input(to, bytes);
return emplace(effect::exchange, to, opcodes::dealer_zero, {payload.id},
bytes, 0, 0, op_flags::none, false, to);
}
/// @brief Pairwise-seed zero mask: local compute, no dealer edge.
/// @details Fills `bytes` from RSS pairwise seeds at `seed_index`. Online
/// drive kernels call `rss::zero_share` when this opcode runs.
HEDLEY_WARN_UNUSED_RESULT
node rss_zero_mask(domain to, std::size_t bytes, std::uint32_t seed_index = 0)
{
if (bytes == 0)
throw std::invalid_argument("rss_zero_mask bytes must be > 0");
if (to != domain::rss && to != domain::a && to != domain::y
&& to != domain::bin_rss)
throw std::invalid_argument("rss_zero_mask domain");
return emplace(effect::compute, to, opcodes::rss_zero, {}, bytes, 0, 0,
op_flags::none, false, to, seed_index);
}
/// @brief One bench cell on the peer edge: emit, send, then the evaluator applies.
/// @details `aux` is the cell id. `nbytes` is the payload. The local work is
/// `drive_options::cell`.
HEDLEY_WARN_UNUSED_RESULT
node bench_peer(std::uint32_t cell, std::size_t nbytes)
{
if (nbytes == 0)
throw std::invalid_argument("bench_peer");
node in = input(domain::a, 8);
node produced = emplace(effect::compute, domain::a, opcodes::bench_emit,
{in.id}, nbytes, 0, 0, op_flags::none, false, domain::a, cell);
node sent = emplace(effect::exchange, domain::a, opcodes::bench_wire,
{produced.id}, nbytes, 0, 0, op_flags::none, false, domain::a, cell);
return emplace(effect::compute, domain::a, opcodes::bench_apply,
{sent.id}, 8, 0, 0, op_flags::none, false, domain::a, cell);
}
/// @brief Three bench passes on the RSS ring. `aux` packs the cell and the pass.
HEDLEY_WARN_UNUSED_RESULT
node bench_ring(std::uint32_t cell, std::size_t nbytes)
{
if (nbytes == 0)
throw std::invalid_argument("bench_ring");
node cur = input(domain::rss, nbytes);
for (unsigned pass = 0; pass < 3; ++pass)
{
const auto aux = cell | (pass << 16);
node produced = emplace(effect::compute, domain::rss, opcodes::bench_emit,
{cur.id}, nbytes, 0, 0, op_flags::none, false, domain::rss, aux);
cur = emplace(effect::exchange, domain::rss, opcodes::bench_ring,
{produced.id}, nbytes, 0, 0, op_flags::none, false, domain::rss, aux);
}
return cur;
}
/// @brief Three hidden-shuffle passes on one replicated array.
/// @details Each pass is one `shuffle_send` exchange. `aux` is the party
/// left out of that pass: 2, then 0, then 1. Messages travel to
/// the sender's RSS neighbor. The bytes are the whole array.
/// Party code is `shuffle::shuffle_hidden_pass`.
HEDLEY_WARN_UNUSED_RESULT
node shuffle_hidden(std::size_t n, std::size_t value_bytes)
{
if (n == 0 || value_bytes == 0)
throw std::invalid_argument("shuffle_hidden");
const std::size_t bytes = n * value_bytes;
node cur = input(domain::rss, bytes);
for (unsigned left : {2u, 0u, 1u})
{
cur = emplace(effect::exchange, domain::rss, opcodes::shuffle_send,
{cur.id}, bytes, 0, 0, op_flags::none, false, domain::rss, left);
}
return cur;
}
/// @brief Reshare under a dealer-sampled zero mask (one delivery, no open).
HEDLEY_WARN_UNUSED_RESULT
node reshare_fresh(node src, domain to)
{
check(src);
node mask = dealer_zero_mask(to, info_[src.id].value_bytes);
return reshare_with_mask(src, mask, to);
}
/// @brief Privacy-preserving reshare: `out = src + mask` locally.
/// @details `mask` is a fresh sharing of zero already in domain `to`
/// (typically from `dealer_deliver`). No reconstructing open.
HEDLEY_WARN_UNUSED_RESULT
node reshare_with_mask(node src, node mask, domain to)
{
check(src);
check(mask);
if (info_[src.id].value_bytes != info_[mask.id].value_bytes)
throw std::invalid_argument("reshare_with_mask size mismatch");
return compute(opcodes::reshare_mask, {src, mask}, to,
info_[src.id].value_bytes);
}
/// @brief One client and `servers` parallel servers. Two waves, no server-server open.
/// @details Wave 1 is the upload (`servers * query_bytes`). Wave 2 is the
/// answers (`servers * answer_bytes`). Servers do not talk to
/// each other; a two-party drive copies the peer message.
struct client_query
{
node upload{};
node answer{};
};
HEDLEY_WARN_UNUSED_RESULT
client_query client_servers(std::size_t servers, std::size_t query_bytes,
std::size_t answer_bytes)
{
if (servers < 2)
throw std::invalid_argument("client_servers needs at least two servers");
if (query_bytes == 0 || answer_bytes == 0)
throw std::invalid_argument("client_servers sizes must be > 0");
node query = input(domain::a, servers * query_bytes);
client_query q;
q.upload = emplace(effect::exchange, domain::a, opcodes::client_upload,
{query.id}, servers * query_bytes, 0, 0, op_flags::none, false,
domain::a);
node formed = compute(opcodes::client_answer, {q.upload}, domain::a,
servers * answer_bytes);
q.answer = emplace(effect::exchange, domain::a, opcodes::client_answer,
{formed.id}, servers * answer_bytes, 0, 0, op_flags::none, false,
domain::a);
return q;
}
/// @brief p2 delivers `payload` to p0 and p1 (one exchange, no reconstruct).
/// @details Drive with `drive_composed_trio`. RoundSink `drive` rejects it.
HEDLEY_WARN_UNUSED_RESULT
node dealer_deliver(node payload)
{
check(payload);
return emplace(effect::exchange, info_[payload.id].dom, opcodes::dealer_pad,
{payload.id}, info_[payload.id].value_bytes, 0, 0, op_flags::none,
false, info_[payload.id].dom);
}
/// @brief Fuse several payloads into one open, then local segment nodes.
struct fuse_result
{
node opened{};
std::vector<node> segments;
};
HEDLEY_WARN_UNUSED_RESULT
fuse_result exchange_fuse(const std::vector<node> & parts)
{
if (parts.size() < 2)
throw std::invalid_argument("exchange_fuse needs at least two parts");
std::size_t total = 0;
for (auto n : parts)
{
check(n);
total += info_[n.id].value_bytes;
}
node packed = compute(opcodes::fss_answer_pack, parts,
info_[parts[0].id].dom, total);
fuse_result fr;
fr.opened = exchange(packed);
std::size_t off = 0;
fr.segments.reserve(parts.size());
for (auto n : parts)
{
const auto width = info_[n.id].value_bytes;
node seg = emplace(effect::compute, info_[n.id].dom,
opcodes::fss_fuse_segment, {fr.opened.id}, width,
0, 0, op_flags::none, false, info_[n.id].dom,
static_cast<std::uint32_t>(off));
fr.segments.push_back(seg);
off += width;
}
return fr;
}
template <typename Fn>
std::vector<node> fan(std::size_t n, Fn && fn)
{
std::vector<node> out;
out.reserve(n);
for (std::size_t i = 0; i < n; ++i)
out.push_back(fn(i));
return out;
}
/// @brief One PRG stretch producing both children (hand DPF schedules do
/// this; separate `expand(..., 0/1)` doubles AES on PIR/Duoram walks).
node expand_pair(node src, std::uint16_t level, std::uint32_t opcode,
std::size_t child_bytes)
{
check(src);
return emplace(effect::expand, info_[src.id].dom, opcode, {src.id},
child_bytes * 2, level, /*both*/ 2, op_flags::none, false,
domain::a);
}
/// @brief Result of a walk whose last CW slot carries an auxiliary trailer
/// (Express sketch / Sabre proof share) — one RoundSink round, not two.
struct walk_result
{
node leaf{};
node trailer_open{}; ///< Reconstructed trailer from the fused exchange
node fused_exchange{};
};
node level_walk(node seed, std::size_t depth, std::size_t slot_bytes,
std::uint32_t expand_op = opcodes::fss_expand_pair,
std::uint32_t step_op = opcodes::fss_step)
{
return level_walk_impl(seed, depth, slot_bytes, node{}, 0, expand_op,
step_op).leaf;
}
/// @brief Like `level_walk`, but packs `trailer` into the last exchange.
/// @details Hand Express/Sabre schedules fold the audit/proof share into the
/// final CW flush so the sketch does not add a round. `trailer_open`
/// is the opened trailer payload from that same wave.
walk_result level_walk_fused(node seed, std::size_t depth,
std::size_t slot_bytes, node trailer,
std::uint32_t expand_op = opcodes::fss_expand_pair,
std::uint32_t step_op = opcodes::fss_step)
{
check(trailer);
return level_walk_impl(seed, depth, slot_bytes, trailer,
info_[trailer.id].value_bytes, expand_op, step_op);
}
node fss_point(node seed, std::size_t depth, std::size_t slot_bytes)
{
node leaf = level_walk(seed, depth, slot_bytes);
info_[leaf.id].dom = domain::b;
return leaf;
}
walk_result fss_point_fused(node seed, std::size_t depth,
std::size_t slot_bytes, node trailer)
{
auto wr = level_walk_fused(seed, depth, slot_bytes, trailer);
info_[wr.leaf.id].dom = domain::b;
return wr;
}
node fss_cmp(node seed, std::size_t depth, std::size_t slot_bytes)
{
node leaf = level_walk(seed, depth, slot_bytes);
info_[leaf.id].dom = domain::a;
return leaf;
}
walk_result fss_cmp_fused(node seed, std::size_t depth,
std::size_t slot_bytes, node trailer)
{
auto wr = level_walk_fused(seed, depth, slot_bytes, trailer);
info_[wr.leaf.id].dom = domain::a;
return wr;
}
/// @brief BGI Remark 3.4 early-stop: only `full_depth - early_stop` CW rounds.
/// @details The remaining `early_stop` domain bits select a lane inside one
/// λ-bit leaf locally — Pika bit leaves and small-output PIR keys.
/// Hand schedules drop those interactive levels; a full `fss_point`
/// would pay them.
node fss_point_early_stop(node seed, std::size_t full_depth,
std::size_t early_stop, std::size_t slot_bytes)
{
if (early_stop >= full_depth)
throw std::invalid_argument("early_stop must be < full_depth");
const std::size_t interactive = full_depth - early_stop;
node tip = level_walk(seed, interactive, slot_bytes);
node leaf = compute(opcodes::fss_early_stop_pack, {tip}, domain::b,
slot_bytes);
info_[leaf.id].level = static_cast<std::uint16_t>(early_stop);
return leaf;
}
node fss_cmp_early_stop(node seed, std::size_t full_depth,
std::size_t early_stop, std::size_t slot_bytes)
{
if (early_stop >= full_depth)
throw std::invalid_argument("early_stop must be < full_depth");
const std::size_t interactive = full_depth - early_stop;
node tip = level_walk(seed, interactive, slot_bytes);
node leaf = compute(opcodes::fss_early_stop_pack, {tip}, domain::a,
slot_bytes);
info_[leaf.id].level = static_cast<std::uint16_t>(early_stop);
return leaf;
}
/// @brief Poplar / idpf_agg: emit a prefix share after each CW — no extra rounds.
struct prefix_walk_result
{
std::vector<node> at_level; ///< Prefix output after level i's exchange
node tip{};
};
prefix_walk_result level_walk_prefixes(node seed, std::size_t depth,
std::size_t slot_bytes, std::size_t prefix_bytes,
std::uint32_t expand_op = opcodes::fss_expand_pair,
std::uint32_t step_op = opcodes::fss_step)
{
check(seed);
if (depth == 0)
throw std::invalid_argument("level_walk_prefixes depth must be positive");
prefix_walk_result wr;
wr.at_level.reserve(depth);
node state = seed;
for (std::size_t level = 0; level < depth; ++level)
{
node kids = expand_pair(state, static_cast<std::uint16_t>(level),
expand_op, slot_bytes);
node step = compute(step_op, {kids}, info_[state.id].dom, slot_bytes);
node msg = exchange(step);
state = compute(step_op + 1, {kids, msg}, info_[state.id].dom,
slot_bytes);
wr.at_level.push_back(compute(opcodes::fss_prefix_emit, {state},
domain::a, prefix_bytes));
}
wr.tip = state;
return wr;
}
/// @brief DCF `block_width` / variable CW size: one slot width per level.
node level_walk_sized(node seed,
const std::vector<std::size_t> & slot_bytes_per_level,
std::uint32_t expand_op = opcodes::fss_expand_pair,
std::uint32_t step_op = opcodes::fss_step)
{
return level_walk_sized_impl(seed, slot_bytes_per_level, node{}, 0,
expand_op, step_op).leaf;
}
walk_result level_walk_sized_fused(node seed,
const std::vector<std::size_t> & slot_bytes_per_level, node trailer,
std::uint32_t expand_op = opcodes::fss_expand_pair,
std::uint32_t step_op = opcodes::fss_step)
{
check(trailer);
return level_walk_sized_impl(seed, slot_bytes_per_level, trailer,
info_[trailer.id].value_bytes, expand_op, step_op);
}
/// @brief Doerner–Shelat interactive walk: blind‖CW‖advice‖AND per level.
/// @details Matches the classic 4-open/level shape. Oblivious-hash mode
/// emits `net::ds_oh_exchanges_per_level` AND-layer opens (Boyar–
/// Peralta SubBytes depth) at `net::ds_oh_slot_bytes`. Prefer
/// `level_walk_ds_sized` with `net::compose_ds_slot_bytes` when
/// RoundSink widths must match the schedule exactly.
HEDLEY_WARN_UNUSED_RESULT
node level_walk_ds(node seed, std::size_t depth, std::size_t slot_bytes,
bool with_oblivious_hash = false, std::size_t lg_outputs = 0,
std::uint32_t expand_op = opcodes::fss_expand_pair)
{
auto slots = net::compose_ds_slot_bytes(depth, slot_bytes,
with_oblivious_hash, lg_outputs);
return level_walk_ds_sized(seed, slots, with_oblivious_hash, lg_outputs,
expand_op);
}
/// @brief DS walk with per-exchange slot widths (see `compose_ds_slot_bytes`).
HEDLEY_WARN_UNUSED_RESULT
node level_walk_ds_sized(node seed, const std::vector<std::size_t> & slots,
bool with_oblivious_hash = false, std::size_t lg_outputs = 0,
std::uint32_t expand_op = opcodes::fss_expand_pair)
{
check(seed);
if (slots.empty())
throw std::invalid_argument("level_walk_ds_sized needs slot widths");
std::size_t si = 0;
auto next_slot = [&]() -> std::size_t {
if (si >= slots.size())
throw std::out_of_range("level_walk_ds_sized: slot vector short");
const auto s = slots[si++];
if (s == 0)
throw std::invalid_argument("level_walk_ds_sized slot must be > 0");
return s;
};
// Infer depth from slot layout when possible; otherwise walk until
// slots are exhausted for the 4(+OH) pattern.
const std::size_t per_level =
5 + (with_oblivious_hash ? net::ds_oh_exchanges_per_level : 0);
std::size_t mux_nodes = 0;
if (lg_outputs > 0 && lg_outputs < 32)
mux_nodes = 2 * ((std::size_t{1} << lg_outputs) - 1);
if (slots.size() < per_level + mux_nodes)
throw std::invalid_argument("level_walk_ds_sized: not enough slots");
const std::size_t depth = (slots.size() - mux_nodes) / per_level;
if (depth == 0 || depth * per_level + mux_nodes != slots.size())
throw std::invalid_argument(
"level_walk_ds_sized: slots must be depth×(4[+OH])+mux");
node state = seed;
for (std::size_t level = 0; level < depth; ++level)
{
const auto expand_slot = slots[si]; // peek; expand uses first open width
node kids = expand_pair(state, static_cast<std::uint16_t>(level),
expand_op, expand_slot);
node blind = compute(opcodes::ds_blind, {kids},
info_[state.id].dom, next_slot());
node blind_ex = exchange(blind);
node cw = compute(opcodes::ds_cw, {kids, blind_ex},
info_[state.id].dom, next_slot());
node cw_ex = exchange(cw);
node advice = compute(opcodes::ds_advice, {kids, cw_ex},
info_[state.id].dom, next_slot());
node advice_ex = exchange(advice);
node aand = compute(opcodes::ds_and, {kids, advice_ex},
info_[state.id].dom, next_slot());
node and1_ex = exchange(aand);
node aand2 = compute(opcodes::ds_and2, {kids, and1_ex},
info_[state.id].dom, next_slot());
node and_ex = exchange(aand2);
if (with_oblivious_hash)
{
for (std::size_t layer = 0; layer < net::ds_oh_exchanges_per_level;
++layer)
{
const auto oh_slot = next_slot();
// level/child distinguish BP AND-layers for interning.
node oh = emplace(effect::compute, info_[state.id].dom,
opcodes::ds_oh, {kids.id, and_ex.id}, oh_slot,
static_cast<std::uint16_t>(level),
static_cast<std::uint8_t>(layer & 0xffu),
op_flags::none, false, domain::a);
and_ex = exchange(oh);
}
}
state = compute(opcodes::fss_step + 1, {kids, and_ex},
info_[state.id].dom, expand_slot);
}
for (std::size_t i = 0; i < mux_nodes; ++i)
{
const auto mux_slot = next_slot();
node mux = compute(opcodes::ds_leaf_mux, {state},
info_[state.id].dom, mux_slot);
node mux_ex = exchange(mux);
state = compute(opcodes::ds_leaf_mux + 1, {state, mux_ex},
info_[state.id].dom, mux_slot);
}
return state;
}
/// @brief Adaptive idpf / eval_until frontier: open L‖R, then retain one.
/// @details Systemic staging — the schedule never learns the child bit:
/// 1. `step` records expand + packed L‖R open only.
/// 2. Drive **new** exchange waves only (`drive_options::from_exchange_wave`
/// = `exchanges_flushed`; size the sink with `slot_bytes_from`).
/// 3. Choose from opened counts; `retain` adds one local tip.
/// 4. Bump `exchanges_flushed` to `plan.exchange_waves()` and step again.
struct adaptive_prefix
{
node tip{};
node last_open{};
node left{}; ///< Opened prefix payload for child 0
node right{}; ///< Opened prefix payload for child 1
node kids{}; ///< Expand pair pending retain (valid while awaiting)
std::size_t slot_bytes = 0;
std::size_t depth_done = 0;
std::size_t exchanges_flushed = 0; ///< Sink rounds already driven
bool awaiting_retain = false;
};
HEDLEY_WARN_UNUSED_RESULT
adaptive_prefix begin_adaptive_prefix(node seed)
{
check(seed);
adaptive_prefix f;
f.tip = seed;
return f;
}
/// @brief One EvaluateUntil step: expand and pack both child opens.
HEDLEY_WARN_UNUSED_RESULT
adaptive_prefix step_adaptive_prefix(adaptive_prefix f,
std::size_t slot_bytes, std::size_t prefix_bytes,
std::uint32_t expand_op = opcodes::fss_expand_pair)
{
check(f.tip);
if (f.awaiting_retain)
throw std::logic_error(
"step_adaptive_prefix: call retain_adaptive_prefix first");
f.kids = expand_pair(f.tip, static_cast<std::uint16_t>(f.depth_done),
expand_op, slot_bytes);
f.slot_bytes = slot_bytes;
node left = compute(opcodes::fss_prefix_emit, {f.kids}, domain::a,
prefix_bytes);
node right = compute(opcodes::fss_prefix_emit + 1, {f.kids}, domain::a,
prefix_bytes);
node packed = compute(opcodes::fss_prefix_pair, {left, right}, domain::a,
prefix_bytes * 2);
f.last_open = exchange(packed);
f.left = compute(opcodes::fss_prefix_emit, {f.last_open}, domain::a,
prefix_bytes);
f.right = compute(opcodes::fss_prefix_emit + 1, {f.last_open}, domain::a,
prefix_bytes);
++f.depth_done;
f.awaiting_retain = true;
return f;
}
/// @brief Local retain after drive: one tip node for the chosen child.
HEDLEY_WARN_UNUSED_RESULT
adaptive_prefix retain_adaptive_prefix(adaptive_prefix f, std::uint8_t child)
{
if (!f.awaiting_retain)
throw std::logic_error(
"retain_adaptive_prefix: no pending step open");
if (child > 1)
throw std::invalid_argument("retain_adaptive_prefix child must be 0 or 1");
check(f.kids);
check(f.last_open);
const auto op = opcodes::fss_prefix_retain + child;
f.tip = compute(op, {f.kids, f.last_open}, info_[f.kids.id].dom,
f.slot_bytes);
info_[f.tip.id].child = child;
f.awaiting_retain = false;
return f;
}
/// @brief The plan RoundSink / `drive` consume — same object, no dual sketch.
HEDLEY_WARN_UNUSED_RESULT
plan default_plan() const { return schedule(); }
/// @brief Keyword PIR / PSI: N independent bucket walks, answers in one open.
struct multipoint_result
{
std::vector<node> leaves;
node answers_open{};
};
template <typename SeedFn>
HEDLEY_WARN_UNUSED_RESULT
multipoint_result multipoint_fan(std::size_t buckets, SeedFn && seed_at,
std::size_t depth, std::size_t slot_bytes, std::size_t answer_bytes)
{
if (buckets == 0)
throw std::invalid_argument("multipoint_fan needs buckets > 0");
multipoint_result mr;
mr.leaves.reserve(buckets);
for (std::size_t i = 0; i < buckets; ++i)
mr.leaves.push_back(fss_point(seed_at(i), depth, slot_bytes));
mr.answers_open = exchange_pack(mr.leaves, answer_bytes);
return mr;
}
/// @brief Occupied cuckoo buckets → one `multipoint_fan` (probes deduped).
template <typename SeedAt>
HEDLEY_WARN_UNUSED_RESULT
multipoint_result schedule_cuckoo_probes(
const std::vector<std::size_t> & probes, SeedAt && seed_at,
std::size_t depth, std::size_t slot_bytes, std::size_t answer_bytes)
{
std::vector<std::size_t> uniq = probes;
std::sort(uniq.begin(), uniq.end());
uniq.erase(std::unique(uniq.begin(), uniq.end()), uniq.end());
if (uniq.empty())
throw std::invalid_argument("schedule_cuckoo_probes: no probes");
return multipoint_fan(uniq.size(),
[&](std::size_t i) { return seed_at(uniq[i]); },
depth, slot_bytes, answer_bytes);
}
/// @brief Pack many payloads into one RoundSink exchange (bucket answers).
HEDLEY_WARN_UNUSED_RESULT
node exchange_pack(const std::vector<node> & payloads,
std::size_t per_payload_bytes = 0)
{
if (payloads.empty())
throw std::invalid_argument("exchange_pack needs payloads");
std::vector<node> ins = payloads;
std::size_t bytes = 0;
for (auto n : payloads)
{
check(n);
bytes += (per_payload_bytes != 0) ? per_payload_bytes
: info_[n.id].value_bytes;
}
node packed = compute(opcodes::fss_answer_pack, std::move(ins),
info_[payloads[0].id].dom, bytes);
return exchange(packed);
}
/// @brief Prepaid full-domain expand — zero online CW rounds (SUBLEQ offline).
HEDLEY_WARN_UNUSED_RESULT
node defer_expand(node seed, std::size_t depth, std::size_t slot_bytes,
std::uint32_t expand_op = opcodes::fss_expand_pair)
{
check(seed);
if (depth == 0)
throw std::invalid_argument("defer_expand depth must be positive");
node state = seed;
for (std::size_t level = 0; level < depth; ++level)
{
node kids = expand_pair(state, static_cast<std::uint16_t>(level),
expand_op, slot_bytes);
state = compute(opcodes::fss_defer_expand, {kids},
info_[state.id].dom, slot_bytes);
}
return state;
}
/// @brief Local rotate of a share buffer after the offset is public.
HEDLEY_WARN_UNUSED_RESULT
node rotate_share(node buf, std::uint32_t public_shift)
{
check(buf);
node out = compute(opcodes::fss_rotate, {buf}, info_[buf.id].dom,
info_[buf.id].value_bytes);
info_[out.id].level = static_cast<std::uint16_t>(public_shift & 0xffffu);
return out;
}
/// @brief Path CW walk with leaf correction deferred (Duoram write shape).
struct leaf_later_result
{
node values{};
node control{};
};
HEDLEY_WARN_UNUSED_RESULT
leaf_later_result leaf_later_walk(node seed, std::size_t depth,
std::size_t slot_bytes,
std::uint32_t expand_op = opcodes::fss_expand_pair,
std::uint32_t step_op = opcodes::fss_step)
{
node tip = level_walk(seed, depth, slot_bytes, expand_op, step_op);
leaf_later_result r;
r.values = compute(opcodes::fss_leaf_later, {tip}, info_[tip.id].dom,
slot_bytes);
r.control = compute(opcodes::fss_leaf_later + 1, {tip}, domain::b, 1);
return r;
}
HEDLEY_WARN_UNUSED_RESULT
node apply_leaf_correction(node values, node control, node leaf_cw)
{
check(values);
check(control);
check(leaf_cw);
return compute(opcodes::fss_apply_leaf, {values, control, leaf_cw},
info_[values.id].dom, info_[values.id].value_bytes);
}
template <typename Ring>
beavers::session<Ring> & aby()
{
return aby_lane<Ring>(0);
}
/// @brief Independent ABY session on sink lane `lane` (Grotto multi-α).
/// @details Barriers from different lanes do not chain; same ready_round
/// flushes pack into one compose wave when data-independent.
template <typename Ring>
beavers::session<Ring> & aby_lane(std::size_t lane)
{
detail::aby_slot_key key{std::type_index(typeid(Ring)), lane};
auto it = aby_.find(key);
if (it == aby_.end())
{
auto holder = std::make_unique<detail::aby_holder<Ring>>();
holder->session.set_schedule_objective(
beavers::schedule_objective::rounds);
auto * ptr = &holder->session;
aby_.emplace(key, std::move(holder));
aby_order_.push_back(key);
return *ptr;
}
return static_cast<detail::aby_holder<Ring> &>(*it->second).session;
}
template <typename Ring>
node opened(typename beavers::session<Ring>::wire w, std::size_t lane = 0)
{
auto & s = aby_lane<Ring>(lane);
const auto barriers = s.exchange_barriers();
std::size_t barrier_index = barriers.size();
for (std::size_t i = 0; i < barriers.size(); ++i)
{
for (auto id : barriers[i].wire_ids)
{
if (id == w.id())
{
barrier_index = i;
break;
}
}
if (barrier_index < barriers.size())
break;
}
detail::aby_slot_key key{std::type_index(typeid(Ring)), lane};
std::size_t sess_index = 0;
for (; sess_index < aby_order_.size(); ++sess_index)
{
if (aby_order_[sess_index].type == key.type
&& aby_order_[sess_index].lane == key.lane)
break;
}
if (barriers.empty())
{
node out = emplace(effect::compute, domain::a, opcodes::beaver_delta,
{}, sizeof(Ring), 0, 0, op_flags::none, false, domain::a);
info_[out.id].aby_wire = static_cast<int>(w.id());
info_[out.id].aby_session = static_cast<int>(sess_index);
return out;
}
if (barrier_index >= barriers.size())
{
const int rw = s.round_of(w);
barrier_index = 0;
for (std::size_t i = 0; i < barriers.size(); ++i)
{
if (barriers[i].ready_round <= rw)
barrier_index = i;
}
}
node ex{};
const std::size_t elem = s.has_mac()
? sizeof(beavers::auth_opening<Ring>) : sizeof(Ring);
for (std::size_t i = 0; i <= barrier_index; ++i)
{
const node * pred = (i == 0) ? nullptr : &ex;
ex = beaver_barrier_node(sess_index, i, elem,
barriers[i].wire_ids, pred);
}
std::vector<std::uint32_t> out_ins = {ex.id};
for (std::size_t i = 0; i <= barrier_index; ++i)
{
for (auto wid : barriers[i].wire_ids)
{
auto it = aby_import_.find({static_cast<int>(sess_index), wid});
if (it == aby_import_.end())
continue;
if (std::find(out_ins.begin(), out_ins.end(), it->second)
== out_ins.end())
out_ins.push_back(it->second);
}
}
{
auto it = aby_import_.find(
{static_cast<int>(sess_index), w.id()});
if (it != aby_import_.end()
&& std::find(out_ins.begin(), out_ins.end(), it->second)
== out_ins.end())
out_ins.push_back(it->second);
}
node out = emplace(effect::compute, domain::a, opcodes::beaver_delta,
std::move(out_ins), sizeof(Ring), 0, 0, op_flags::none, false,
domain::a);
info_[out.id].aby_wire = static_cast<int>(w.id());
info_[out.id].aby_session = static_cast<int>(sess_index);
return out;
}
template <typename Ring>
node beaver_triple(node x, node y, std::size_t lane = 0)
{
x = as(x, domain::a);
y = as(y, domain::a);
auto & s = aby_lane<Ring>(lane);
auto wx = ensure_aby_wire<Ring>(x, lane);
auto wy = ensure_aby_wire<Ring>(y, lane);
auto wz = s.product(wx, wy);
return opened<Ring>(wz, lane);
}
template <typename Ring>
node aby_product(node x, node y, std::size_t lane = 0)
{
return beaver_triple<Ring>(x, y, lane);
}
node rss_product(node x, node y)
{
x = as(x, domain::rss);
y = as(y, domain::rss);
std::size_t out_bytes = info_[x.id].value_bytes;
if (out_bytes >= 2)
out_bytes /= 2;
return compute(opcodes::rss_mul, {x, y}, domain::y, out_bytes,
op_flags::commutative);
}
/// @brief Arithmetic→boolean conversion node (kernel + open of delta).
HEDLEY_WARN_UNUSED_RESULT
node bin_a2b(node x, std::size_t out_bytes)
{
check(x);
return compute(opcodes::bin_a2b, {x}, domain::bin, out_bytes);
}
/// @brief Local `x - y` (masked Beaver limb).
HEDLEY_WARN_UNUSED_RESULT
node arith_sub(node x, node y)
{
check(x);
check(y);
return emplace(effect::compute, domain::a, opcodes::share_input,
{x.id, y.id}, 8, 0, 0, op_flags::none, false, domain::a, 0);
}
/// @brief Local XOR of two equal-width bit buffers.
HEDLEY_WARN_UNUSED_RESULT
node bit_xor(node a, node b, std::size_t bytes)
{
check(a);
check(b);
return emplace(effect::compute, domain::bin, opcodes::bin_and,
{a.id, b.id}, bytes, 0, 0, op_flags::commutative, false, domain::bin,
/*aux=*/1);
}
/// @brief B2A: XOR-open `bits ⊕ r_bits`, then daBit correction.
/// @details `r_arith` is one `uint64` additive share per bit, packed back to back.
HEDLEY_WARN_UNUSED_RESULT
node bin_b2a(node bits, node r_bits, node r_arith, unsigned ell)
{
check(bits);
check(r_bits);
check(r_arith);
const std::size_t nbytes = (static_cast<std::size_t>(ell) + 7u) / 8u;
node mask = bit_xor(bits, r_bits, nbytes);
node opened = exchange(mask);
return emplace(effect::compute, domain::a, opcodes::bin_b2a,
{r_arith.id, opened.id}, 8, 0, 0, op_flags::none, false, domain::a,
ell);
}
/// @brief Bit × arithmetic via one bit×ring triple and two opens.
/// @details `bit` and `a_bit` are additive 0/1 shares (`uint64`). `b_arith`
/// and `c_arith` are the additive shares of `b` and `c`.
HEDLEY_WARN_UNUSED_RESULT
node bin_inject(node bit, node x, node a_bit, node b_arith, node c_arith)
{
check(bit);
check(x);
check(a_bit);
check(b_arith);
check(c_arith);
node d = arith_sub(bit, a_bit);
node e = arith_sub(x, b_arith);
node od = exchange(d);
node oe = exchange(e);
return emplace(effect::compute, domain::a, opcodes::bin_inject,
{a_bit.id, b_arith.id, c_arith.id, od.id, oe.id}, 8, 0, 0,
op_flags::none, false, domain::a, 0);
}
/// @brief RSS AND: local factor from replicated components, ring-send, refresh.
HEDLEY_WARN_UNUSED_RESULT
node bin_rss_and(node a_own, node a_next, node b_own, node b_next)
{
check(a_own);
check(a_next);
check(b_own);
check(b_next);
node local = emplace(effect::compute, domain::y, opcodes::bin_rss_and,
{a_own.id, a_next.id, b_own.id, b_next.id}, 1, 0, 0,
op_flags::none, false, domain::y, 0);
// Ring send. The own component (this node) is what the three parties XOR.
(void)exchange(local);
return local;
}
/// @brief RSS bit injection: local product factor plus a ring send of it.
HEDLEY_WARN_UNUSED_RESULT
node bin_rss_inject(node b_own, node b_next, node x_own, node x_next)
{
check(b_own);
check(b_next);
check(x_own);
check(x_next);
node local = emplace(effect::compute, domain::y, opcodes::bin_rss_inject,
{b_own.id, b_next.id, x_own.id, x_next.id}, 8, 0, 0,
op_flags::none, false, domain::y, 0);
(void)exchange(local);
return local;
}
/// @brief GMW AND: open `p⊕a` and `q⊕b`, then the local triple correction.
HEDLEY_WARN_UNUSED_RESULT
node gmw_and(node p, node q, node a, node b, node c)
{
check(p);
check(q);
check(a);
check(b);
check(c);
node d = bit_xor(p, a, 1);
node e = bit_xor(q, b, 1);
node od = exchange(d);
node oe = exchange(e);
return emplace(effect::compute, domain::bin, opcodes::bin_and,
{a.id, b.id, c.id, od.id, oe.id}, 1, 0, 0, op_flags::none, false,
domain::bin, /*aux=*/2);
}
/// @brief Beaver product then local shift. Opens `x-a` and `y-b`.
HEDLEY_WARN_UNUSED_RESULT
node mul_trunc_open(node x, node y, node a, node b, node c, unsigned s)
{
check(x);
check(y);
check(a);
check(b);
check(c);
node od = exchange(arith_sub(x, a));
node oe = exchange(arith_sub(y, b));
return emplace(effect::compute, domain::a, opcodes::mul_trunc,
{a.id, b.id, c.id, od.id, oe.id}, 8, 0, 0, op_flags::none, false,
domain::a, s);
}
/// @brief `b + sel·(a-b)` with one bit×ring inject of `sel` and `a-b`.
HEDLEY_WARN_UNUSED_RESULT
node share_mux_open(node sel, node a, node b, node t_a, node t_b, node t_c)
{
check(b);
node diff = arith_sub(a, b);
node prod = bin_inject(sel, diff, t_a, t_b, t_c);
return emplace(effect::compute, domain::a, opcodes::reshare_mask,
{b.id, prod.id}, 8, 0, 0, op_flags::none, false, domain::a, 0);
}
HEDLEY_WARN_UNUSED_RESULT
node trunc_prob_op(node x, unsigned shift)
{
check(x);
return emplace(effect::compute, info_[x.id].dom, opcodes::trunc_prob,
{x.id}, info_[x.id].value_bytes, 0, 0, op_flags::none, false,
info_[x.id].dom, shift);
}
HEDLEY_WARN_UNUSED_RESULT
node trunc_exact_op(node x, unsigned /*n*/, unsigned s)
{
check(x);
return emplace(effect::compute, info_[x.id].dom, opcodes::trunc_exact,
{x.id}, info_[x.id].value_bytes, 0, 0, op_flags::none, false,
info_[x.id].dom, s);
}
/// @brief Exact trunc after `delta` has been opened. Inputs are this party's
/// `x`, `r`, public `delta`, and wrap share. `aux` is the shift.
HEDLEY_WARN_UNUSED_RESULT
node trunc_exact_open(node x, node r, node delta, node wrap, unsigned s)
{
check(x);
check(r);
check(delta);
check(wrap);
return emplace(effect::compute, domain::a, opcodes::trunc_exact,
{x.id, r.id, delta.id, wrap.id}, 8, 0, 0, op_flags::none, false,
domain::a, s);
}
HEDLEY_WARN_UNUSED_RESULT
node mul_trunc_op(node x, node y, unsigned s)
{
check(x);
check(y);
return emplace(effect::compute, domain::a, opcodes::mul_trunc,
{x.id, y.id}, info_[x.id].value_bytes, 0, 0, op_flags::commutative,
false, domain::a, s);
}
HEDLEY_WARN_UNUSED_RESULT
node share_gt(node x, node y)
{
check(x);
check(y);
return compute(opcodes::share_cmp, {x, y}, domain::bin, 1);
}
HEDLEY_WARN_UNUSED_RESULT
node share_mux_op(node sel, node a, node b)
{
check(sel);
check(a);
check(b);
return compute(opcodes::share_mux, {sel, a, b}, domain::a,
info_[a.id].value_bytes);
}
HEDLEY_WARN_UNUSED_RESULT
domain domain_of(node n) const
{
check(n);
return info_[n.id].dom;
}
HEDLEY_WARN_UNUSED_RESULT
effect effect_of(node n) const
{
check(n);
return info_[n.id].kind;
}
HEDLEY_WARN_UNUSED_RESULT
std::size_t value_bytes(node n) const
{
check(n);
return info_[n.id].value_bytes;
}
HEDLEY_NO_THROW
HEDLEY_WARN_UNUSED_RESULT
std::size_t node_count() const noexcept { return info_.size(); }
HEDLEY_WARN_UNUSED_RESULT
plan schedule() const
{
plan p;
const std::size_t n = info_.size();
p.node_wave_.assign(n, 0);
p.value_bytes_.resize(n);
p.domains_.resize(n);
p.effects_.resize(n);
p.opcodes_.resize(n);
p.levels_.resize(n);
p.aux_.resize(n);
p.inputs_.resize(n);
p.alias_bytes_.resize(n);
p.all_nodes_.resize(n);
for (std::uint32_t id = 0; id < n; ++id)
{
p.all_nodes_[id] = node{id};
p.value_bytes_[id] = info_[id].value_bytes;
p.domains_[id] = info_[id].dom;
p.effects_[id] = info_[id].kind;
p.opcodes_[id] = info_[id].opcode;
p.levels_[id] = info_[id].level;
p.aux_[id] = info_[id].aux;
p.inputs_[id] = info_[id].inputs;
p.alias_bytes_[id] = info_[id].alias_bytes;
++p.effect_counts_[info_[id].kind];
if (info_[id].kind == effect::convert)
++p.conversion_counts_[{info_[id].from_domain, info_[id].dom}];
}
std::vector<char> seen(n, 0);
std::function<std::size_t(std::uint32_t)> wave_of =
[&](std::uint32_t id) -> std::size_t {
if (seen[id])
return p.node_wave_[id];
seen[id] = 1;
std::size_t w = 0;
for (auto in : info_[id].inputs)
{
// Only real compose nodes are dependencies (skip stale ids).
if (in >= n)
continue;
w = std::max(w, wave_of(in));
}
p.node_wave_[id] = w;
// If any input was an exchange, this node is after that barrier.
for (auto in : info_[id].inputs)
{
if (in >= n)
continue;
if (info_[in].kind == effect::exchange)
p.node_wave_[id] = std::max(p.node_wave_[id],
p.node_wave_[in] + 1);
}
return p.node_wave_[id];
};
for (std::uint32_t id = 0; id < n; ++id)
wave_of(id);
std::size_t max_wave = 0;
bool any_exchange = false;
for (std::uint32_t id = 0; id < n; ++id)
{
max_wave = std::max(max_wave, p.node_wave_[id]);
if (info_[id].kind == effect::exchange)
any_exchange = true;
}
const std::size_t wave_count =
n == 0 ? 0 : (any_exchange ? max_wave + 1 : max_wave + 1);
p.waves_.assign(wave_count, wave_info{});
for (std::size_t w = 0; w < wave_count; ++w)
{
p.waves_[w].index = w;
std::map<std::pair<effect, std::uint32_t>, std::vector<node>> groups;
for (std::uint32_t id = 0; id < n; ++id)
{
if (p.node_wave_[id] != w)
continue;
if (info_[id].kind == effect::exchange)
{
p.waves_[w].exchanges.push_back(node{id});
continue;
}
if (info_[id].kind == effect::input)
continue;
groups[{info_[id].kind, info_[id].opcode}].push_back(node{id});
}
for (auto & kv : groups)
{
simd_group g;
g.kind = kv.first.first;
g.opcode = kv.first.second;
g.nodes = std::move(kv.second);
p.waves_[w].groups.push_back(std::move(g));
}
std::sort(p.waves_[w].exchanges.begin(), p.waves_[w].exchanges.end(),
[](node a, node b) { return a.id < b.id; });
std::size_t slot = 0;
for (auto ex : p.waves_[w].exchanges)
slot += info_[ex.id].value_bytes;
p.waves_[w].slot_bytes = slot;
for (auto ex : p.waves_[w].exchanges)
{
const auto op = info_[ex.id].opcode;
const bool is_setup = (op == opcodes::dealer_pad
|| op == opcodes::dealer_zero
|| op == opcodes::client_upload);
if (is_setup)
p.setup_bytes_ += info_[ex.id].value_bytes;
else
p.online_bytes_ += info_[ex.id].value_bytes;
}
}
return p;
}
private:
struct node_info
{
effect kind = effect::input;
domain dom = domain::a;
std::size_t value_bytes = 0;
std::uint32_t opcode = 0;
std::vector<std::uint32_t> inputs;
std::uint16_t level = 0;
std::uint8_t child = 0;
op_flags flags = op_flags::none;
bool alias_bytes = false;
domain from_domain = domain::a;
int aby_wire = -1;
int aby_session = -1;
std::uint32_t aux = 0;
};
void check(node n) const
{
if (n.id >= info_.size())
throw std::out_of_range("compose node");
}
walk_result level_walk_impl(node seed, std::size_t depth,
std::size_t slot_bytes, node trailer, std::size_t trailer_bytes,
std::uint32_t expand_op, std::uint32_t step_op)
{
check(seed);
if (depth == 0)
throw std::invalid_argument("level_walk depth must be positive");
walk_result wr;
node state = seed;
for (std::size_t level = 0; level < depth; ++level)
{
node kids = expand_pair(state, static_cast<std::uint16_t>(level),
expand_op, slot_bytes);
node step = compute(step_op, {kids}, info_[state.id].dom,
slot_bytes);
const bool last = (level + 1 == depth);
if (last && trailer_bytes != 0)
{
// Express/Sabre: sketch or proof share rides with the final CW.
node packed = compute(opcodes::fss_fuse_trailer, {step, trailer},
info_[state.id].dom, slot_bytes + trailer_bytes);
wr.fused_exchange = exchange(packed);
wr.trailer_open = compute(opcodes::fss_fuse_trailer + 1,
{wr.fused_exchange}, info_[trailer.id].dom, trailer_bytes);
state = compute(step_op + 1, {kids, wr.fused_exchange},
info_[state.id].dom, slot_bytes);
}
else
{
node msg = exchange(step);
state = compute(step_op + 1, {kids, msg}, info_[state.id].dom,
slot_bytes);
if (last)
wr.fused_exchange = msg;
}
}
wr.leaf = state;
return wr;
}
walk_result level_walk_sized_impl(node seed,
const std::vector<std::size_t> & slot_bytes_per_level, node trailer,
std::size_t trailer_bytes, std::uint32_t expand_op,
std::uint32_t step_op)
{
check(seed);
if (slot_bytes_per_level.empty())
throw std::invalid_argument("level_walk_sized needs at least one level");
walk_result wr;
node state = seed;
const std::size_t depth = slot_bytes_per_level.size();
for (std::size_t level = 0; level < depth; ++level)
{
const std::size_t slot_bytes = slot_bytes_per_level[level];
if (slot_bytes == 0)
throw std::invalid_argument("level_walk_sized slot_bytes must be > 0");
node kids = expand_pair(state, static_cast<std::uint16_t>(level),
expand_op, slot_bytes);
node step = compute(step_op, {kids}, info_[state.id].dom,
slot_bytes);
const bool last = (level + 1 == depth);
if (last && trailer_bytes != 0)
{
node packed = compute(opcodes::fss_fuse_trailer, {step, trailer},
info_[state.id].dom, slot_bytes + trailer_bytes);
wr.fused_exchange = exchange(packed);
wr.trailer_open = compute(opcodes::fss_fuse_trailer + 1,
{wr.fused_exchange}, info_[trailer.id].dom, trailer_bytes);
state = compute(step_op + 1, {kids, wr.fused_exchange},
info_[state.id].dom, slot_bytes);
}
else
{
node msg = exchange(step);
state = compute(step_op + 1, {kids, msg}, info_[state.id].dom,
slot_bytes);
if (last)
wr.fused_exchange = msg;
}
}
wr.leaf = state;
return wr;
}
node emplace(effect kind, domain dom, std::uint32_t opcode,
std::vector<std::uint32_t> inputs, std::size_t value_bytes,
std::uint16_t level, std::uint8_t child, op_flags flags, bool alias,
domain from_domain, std::uint32_t aux = 0)
{
detail::intern_key key;
key.kind = kind;
key.dom = dom;
key.opcode = opcode;
key.level = level;
key.child = child;
key.aux = aux;
key.inputs = inputs;
auto it = intern_.find(key);
if (it != intern_.end())
return node{it->second};
node_info ni;
ni.kind = kind;
ni.dom = dom;
ni.value_bytes = value_bytes;
ni.opcode = opcode;
ni.inputs = std::move(inputs);
ni.level = level;
ni.child = child;
ni.aux = aux;
ni.flags = flags;
ni.alias_bytes = alias;
ni.from_domain = from_domain;
const auto id = static_cast<std::uint32_t>(info_.size());
info_.push_back(std::move(ni));
intern_.emplace(std::move(key), id);
return node{id};
}
/// @brief One ABY δ flush. `pred` is the previous barrier in this session
/// (real graph edge). Session/barrier identity lives in the intern
/// key's opcode/level — never as fake input node ids (those used to
/// pull waves onto unrelated FSS exchanges).
/// @details Compose nodes imported into ABY wires listed in `wire_ids` are
/// real inputs so Duoram/SUBLEQ scale waits for the FSS leaf.
node beaver_barrier_node(std::size_t sess_index, std::size_t barrier_index,
std::size_t ring_bytes, const std::vector<std::uint32_t> & wire_ids,
const node * pred)
{
detail::intern_key key;
key.kind = effect::exchange;
key.dom = domain::a;
// Unique per (session, barrier); not a graph dependency.
key.opcode = opcodes::beaver_delta
+ static_cast<std::uint32_t>(barrier_index);
key.level = static_cast<std::uint16_t>(sess_index);
key.child = 0;
if (pred != nullptr)
key.inputs.push_back(pred->id);
for (auto wid : wire_ids)
{
auto it = aby_import_.find({static_cast<int>(sess_index), wid});
if (it == aby_import_.end())
continue;
if (std::find(key.inputs.begin(), key.inputs.end(), it->second)
== key.inputs.end())
key.inputs.push_back(it->second);
}
auto it = intern_.find(key);
if (it != intern_.end())
return node{it->second};
node_info ni;
ni.kind = effect::exchange;
ni.dom = domain::a;
ni.value_bytes = sizeof(std::uint32_t) + wire_ids.size() * ring_bytes;
ni.opcode = key.opcode;
ni.inputs = key.inputs;
ni.level = key.level;
ni.child = key.child;
ni.aby_session = static_cast<int>(sess_index);
const auto id = static_cast<std::uint32_t>(info_.size());
info_.push_back(std::move(ni));
intern_.emplace(std::move(key), id);
return node{id};
}
template <typename Ring>
typename beavers::session<Ring>::wire ensure_aby_wire(node n,
std::size_t lane = 0)
{
auto & s = aby_lane<Ring>(lane);
if (info_[n.id].aby_wire >= 0
&& info_[n.id].aby_session >= 0
&& static_cast<std::size_t>(info_[n.id].aby_session) < aby_order_.size()
&& aby_order_[static_cast<std::size_t>(info_[n.id].aby_session)].lane
== lane
&& aby_order_[static_cast<std::size_t>(info_[n.id].aby_session)].type
== std::type_index(typeid(Ring)))
return s.wire_at(static_cast<std::uint32_t>(info_[n.id].aby_wire));
auto w = s.input();
info_[n.id].aby_wire = static_cast<int>(w.id());
detail::aby_slot_key key{std::type_index(typeid(Ring)), lane};
for (std::size_t i = 0; i < aby_order_.size(); ++i)
{
if (aby_order_[i] == key)
{
info_[n.id].aby_session = static_cast<int>(i);
aby_import_[{static_cast<int>(i), w.id()}] = n.id;
break;
}
}
return w;
}
std::size_t party_ = 0;
std::vector<node_info> info_;
std::unordered_map<detail::intern_key, std::uint32_t, detail::intern_key_hash>
intern_;
std::map<detail::aby_slot_key, std::unique_ptr<detail::aby_holder_base>> aby_;
std::vector<detail::aby_slot_key> aby_order_;
/// ABY input wire id → compose node that produced the share (per session).
std::map<std::pair<int, std::uint32_t>, std::uint32_t> aby_import_;
};
/// @brief Optional beaver δ host for composed drives (plain openings).
/// @details When set, exchange nodes with `opcode ∈ [beaver_delta, +1024)` call
/// `pack` / `apply` instead of treating the slot as a reconstruct sum.
/// Wire layout matches `evaluate_online_on_sink`: `uint32 n ‖ Ring[n]`.
struct beaver_host
{
virtual ~beaver_host() = default;
virtual void pack(std::size_t session_index, std::size_t barrier_index,
std::uint8_t * dst, std::size_t n) = 0;
virtual void apply(std::size_t session_index, std::size_t barrier_index,
const std::uint8_t * src, std::size_t n) = 0;
};
/// @brief How non-`y` openings combine peer bytes.
enum class field_open { words, fp61 };
/// @brief Options for `drive` / `drive_composed*`.
/// @brief Party-aware local work for a bench cell. `aux` is the cell id
/// (shuffle passes pack the pass index in the high half).
using cell_fn = void (*)(std::uint32_t opcode, std::size_t party,
std::uint32_t aux, std::uint8_t * out, std::size_t out_n);
struct drive_options
{
/// Skip plan exchange-waves `[0, from)` (already flushed). Local compute
/// for those waves is also skipped — `values` must retain prior results.
std::size_t from_exchange_wave = 0;
/// When true, RoundSink round 0 maps to plan exchange-wave
/// `from_exchange_wave` (sink sized with `slot_bytes_from`).
bool compact_sink = false;
beaver_host * beavers = nullptr;
/// `words`: u64 add/sub (or XOR if not 8 bytes). `fp61`: Mersenne field lanes.
field_open field = field_open::words;
/// Return instead of spinning when the peer has not flushed this round.
bool park_if_waiting = false;
/// Stop after one completed exchange (used by the fleet scheduler).
bool one_exchange = false;
struct drive_cursor * cursor = nullptr;
/// Optional per-round cost probe (`experiment::probe()`).
const round_probe * probe = nullptr;
/// Allow `schedule_session` / drive to flush up to `k` independent rounds ahead.
std::size_t pipeline_credit = 0;
/// Trivial shares (party 0 holds the secret); check reconstructing opens.
bool cleartext = false;
/// Optional clear outputs keyed by exchange node id (cleartext mode).
const std::map<std::uint32_t, std::vector<std::uint8_t>> * clear_oracle = nullptr;
/// Extra consistency check on reconstructed opens (default off).
bool check_open = false;
/// Pairwise RSS seeds for `rss_zero` kernels (required when that opcode runs).
const rss::party_seeds * rss_seeds = nullptr;
/// Base index added to `aux` seed indices for RSS kernels.
std::uint64_t rss_seed_base = 0;
/// When set, custom and builtin walk kernels run on `workers->post_compute`.
net::io_pool * workers = nullptr;
/// Context serviced while a kernel runs on `workers`. Async drives set it
/// to the sink's context; otherwise `workers->context()`.
asio::io_context * pump = nullptr;
/// Stream lanes. `0` → `net::default_lanes` (8);
/// `net::lanes_one_per_round` → one stream per round.
std::size_t n_lanes = 0;
/// Round headers on lanes (`automatic` frames only when rounds share lanes).
net::framing_mode framing = net::framing_mode::automatic;
/// Longest wait for one round's peer bytes, measured from when the
/// instance entered that round. `0` keeps the historical spin cap.
std::chrono::milliseconds wait_timeout{30000};
/// Per-edge override of `wait_timeout` (edge id → budget).
std::map<edge_id, std::chrono::milliseconds> edge_timeout;
/// Per-round override (schedule round index → budget); wins over edges.
std::map<std::uint16_t, std::chrono::milliseconds> round_timeout;
/// Bench cells (`bench_emit` / `bench_apply`). Null for ordinary plans.
cell_fn cell = nullptr;
};
/// @brief How many duplex lanes a plan should open.
inline std::size_t lanes_for_plan(std::size_t n_rounds, const drive_options & opt)
{
return net::lane_count_for_rounds(n_rounds, opt.n_lanes);
}
/// @brief Budget for one wait: round override, else edge override, else default.
inline std::chrono::milliseconds wait_budget(const drive_options & opt,
edge_id edge, std::chrono::milliseconds round_budget = {})
{
if (round_budget.count() > 0)
return round_budget;
auto it = opt.edge_timeout.find(edge);
if (it != opt.edge_timeout.end())
return it->second;
return opt.wait_timeout;
}
inline std::chrono::milliseconds round_budget(const drive_options & opt,
std::size_t round)
{
auto it = opt.round_timeout.find(static_cast<std::uint16_t>(round));
return it == opt.round_timeout.end() ? std::chrono::milliseconds(0)
: it->second;
}
namespace detail
{
/// @brief One wait step on `sink`: sleep in its reactor, or yield if it cannot.
inline bool pump_sink(net::RoundSink & sink, std::chrono::milliseconds slice)
{
if (sink.can_block())
return sink.wait_io_for(slice);
const bool ran = sink.wait_io();
if (!ran)
std::this_thread::yield();
return ran;
}
/// @brief Pump `sink` until `ready()`; `budget` 0 keeps the spin cap.
template <typename Ready>
inline void wait_ready(net::RoundSink & sink, std::chrono::milliseconds budget,
Ready ready, const std::string & what)
{
const auto started = std::chrono::steady_clock::now();
unsigned spins = 0;
while (!ready())
{
if (budget.count() == 0)
{
if (pump_sink(sink, std::chrono::milliseconds(1)))
spins = 0;
else if (++spins > 100000u)
throw std::runtime_error(what + ": peer not ready (drive both ends)");
continue;
}
const auto waited = std::chrono::duration_cast<std::chrono::milliseconds>(
std::chrono::steady_clock::now() - started);
if (waited >= budget)
throw std::runtime_error(what + ": no peer bytes for "
+ std::to_string(waited.count()) + " ms (budget "
+ std::to_string(budget.count())
+ " ms; is the peer driving the same plan?)");
pump_sink(sink, std::min(budget - waited, std::chrono::milliseconds(1000)));
}
}
} // namespace detail
/// @brief Run a kernel inline, or on the compute pool while servicing I/O.
/// @details A kernel writes into `values`, so it cannot be abandoned: this
/// waits for it, running socket completions on `opt.pump` meanwhile.
/// On the pool the kernel draws from this thread's randomness source,
/// and its symmetric-key blocks, random bytes, and CPU time are
/// charged to this thread (`dpf/thread_work.hpp`).
inline void invoke_kernel(const drive_options & opt, const kernel_fn & fn,
std::uint32_t opcode, const std::vector<node> & nodes,
const std::vector<block_span> & inputs, block_span output, std::size_t lanes)
{
if (opt.workers == nullptr)
{
fn(opcode, nodes, inputs, output, lanes);
return;
}
asio::io_context & io = opt.pump != nullptr ? *opt.pump : opt.workers->context();
asio::io_context * wake = &io;
std::atomic<bool> done{false};
std::exception_ptr err;
const auto draws = draw_context::capture();
work_meter before;
work_meter after;
opt.workers->post_compute([&, wake] {
{
const adopt_draws adopted(draws);
before = work_meter::now();
try
{
fn(opcode, nodes, inputs, output, lanes);
}
catch (...)
{
err = std::current_exception();
}
after = work_meter::now();
}
done.store(true, std::memory_order_release);
// Only by-value captures may be touched once `done` is set.
asio::post(*wake, [] {});
});
while (!done.load(std::memory_order_acquire))
{
if (io.stopped())
io.restart();
io.run_one_for(std::chrono::milliseconds(1));
}
absorb_work(before, after);
if (err)
std::rethrow_exception(err);
}
/// @brief Resume point for `drive` when `park_if_waiting` yields.
struct drive_cursor
{
std::size_t wave = 0;
std::size_t exchange_i = 0;
std::uint16_t sink_round = 0;
bool awaiting_peer = false;
bool done = false;
};
inline std::uint64_t open_hash(const std::uint8_t * word, std::size_t n) noexcept
{
std::uint64_t h = 14695981039346656037ull;
for (std::size_t i = 0; i < n; ++i)
{
h ^= word[i];
h *= 1099511628211ull;
}
return h;
}
/// @brief Second-round hash of a reconstructed word on one duplex edge.
inline void exchange_check_hash(net::RoundSink & sink, const std::uint8_t * word,
std::size_t n, std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
const auto h = open_hash(word, n);
std::uint8_t buf[8];
std::memcpy(buf, &h, 8);
sink.submit(0, 0, buf, 8);
sink.flush();
net::wait_peer_ready(sink, 0, 0, budget, "compose check_open: peer hash");
std::uint8_t peer[8];
sink.read_peer(0, 0, peer, 8);
std::uint64_t ph = 0;
std::memcpy(&ph, peer, 8);
if (ph != h)
throw std::runtime_error("compose check_open: hash exchange mismatch");
}
/// @brief 3PC ring: send the hash to the next party, compare with the previous.
inline void exchange_check_hash_ring(net::RoundSink & to_next,
net::RoundSink & from_prev, const std::uint8_t * word, std::size_t n,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
const auto h = open_hash(word, n);
std::uint8_t buf[8];
std::memcpy(buf, &h, 8);
to_next.submit(0, 0, buf, 8);
to_next.flush();
net::wait_peer_ready(from_prev, 0, 0, budget, "compose check_open: ring hash");
std::uint8_t peer[8];
from_prev.read_peer(0, 0, peer, 8);
std::uint64_t ph = 0;
std::memcpy(&ph, peer, 8);
if (ph != h)
throw std::runtime_error("compose check_open: hash exchange mismatch");
}
namespace detail
{
HEDLEY_CONST
inline bool is_beaver_opcode(std::uint32_t op) noexcept
{
return op >= opcodes::beaver_delta && op < opcodes::beaver_delta + 1024u;
}
HEDLEY_CONST
inline bool is_rss_neighbor_open(const plan & p, std::uint32_t ex_id) noexcept
{
if (p.domain_of(ex_id) == domain::y)
return true;
const auto & ins = p.inputs_of(ex_id);
return !ins.empty() && p.domain_of(ins[0]) == domain::y;
}
/// domain::a / fss: sum; domain::b: p0−p1 style difference (party 0 keeps
/// mine−peer, party 1 keeps peer−mine via the same mine−peer on each side
/// when shares are subtractive); domain::y: keep peer bytes.
inline void reconstruct_open(domain d, std::size_t party,
const std::uint8_t * mine, const std::uint8_t * peer, std::size_t n,
std::uint8_t * out, field_open field = field_open::words)
{
if (d == domain::y)
{
std::memcpy(out, peer, n);
return;
}
if (d == domain::bin || d == domain::bin_rss)
{
for (std::size_t i = 0; i < n; ++i)
out[i] = static_cast<std::uint8_t>(mine[i] ^ peer[i]);
return;
}
auto combine_word = [&](std::uint64_t m, std::uint64_t pr) -> std::uint64_t {
if (field == field_open::fp61)
{
constexpr std::uint64_t mod = (std::uint64_t{1} << 61) - 1;
auto red = [mod](std::uint64_t x) {
x = (x & mod) + (x >> 61);
if (x >= mod)
x -= mod;
return x;
};
m = red(m);
pr = red(pr);
if (d == domain::b)
return red((party == 0) ? (m + mod - pr) : (pr + mod - m));
return red(m + pr);
}
if (d == domain::b)
return (party == 0) ? (m - pr) : (pr - m);
return m + pr;
};
if (n >= 8 && n % 8 == 0)
{
for (std::size_t i = 0; i < n; i += 8)
{
std::uint64_t m = 0, pr = 0;
std::memcpy(&m, mine + i, 8);
std::memcpy(&pr, peer + i, 8);
const std::uint64_t r = combine_word(m, pr);
std::memcpy(out + i, &r, 8);
}
return;
}
for (std::size_t i = 0; i < n; ++i)
out[i] = static_cast<std::uint8_t>(mine[i] ^ peer[i]);
}
inline void apply_fuse_segment(const plan & p,
std::vector<std::vector<std::uint8_t>> & values, std::uint32_t id,
std::size_t lanes)
{
const auto & ins = p.inputs_of(id);
if (ins.empty())
return;
const auto src = ins[0];
const auto off = static_cast<std::size_t>(p.aux_of(id));
const auto nb = p.value_bytes_of(id);
const auto sb = p.value_bytes_of(src);
for (std::size_t lane = 0; lane < lanes; ++lane)
{
if (off + nb > sb)
throw std::runtime_error("compose fuse segment out of range");
std::memcpy(values[id].data() + lane * nb,
values[src].data() + lane * sb + off, nb);
}
}
inline bool is_arith_walk_opcode(std::uint32_t opcode) noexcept
{
return opcode == opcodes::bin_a2b || opcode == opcodes::bin_b2a
|| opcode == opcodes::bin_inject || opcode == opcodes::bin_and
|| opcode == opcodes::bin_rss_and || opcode == opcodes::bin_rss_inject
|| opcode == opcodes::rss_zero || opcode == opcodes::trunc_prob
|| opcode == opcodes::trunc_exact || opcode == opcodes::mul_trunc
|| opcode == opcodes::share_cmp || opcode == opcodes::share_mux
|| opcode == opcodes::share_input;
}
inline kernel_fn builtin_walk_kernel(std::uint32_t opcode)
{
return [opcode](std::uint32_t, const std::vector<node> &,
const std::vector<block_span> & inputs, block_span output,
std::size_t) {
if (output.data == nullptr || output.lanes == 0)
return;
const std::size_t nbytes = output.lanes * output.value_bytes;
std::fill(output.data, output.data + nbytes, 0);
if (inputs.empty() || inputs[0].data == nullptr)
return;
if (opcode == opcodes::fss_expand_pair && inputs[0].value_bytes >= 16
&& output.value_bytes >= 16 && output.value_bytes % 16 == 0)
{
for (std::size_t lane = 0; lane < output.lanes; ++lane)
{
simde__m128i seed{};
std::memcpy(&seed, inputs[0].at(lane), 16);
if (output.value_bytes == 32)
{
auto both = dpf::prg::aes128::eval01(seed);
std::memcpy(output.at(lane), both.data(), 32);
}
else
{
for (std::uint32_t b = 0; b < output.value_bytes / 16; ++b)
{
auto blk = dpf::prg::aes128::eval(seed, b);
std::memcpy(output.at(lane) + static_cast<std::size_t>(b) * 16,
&blk, 16);
}
}
}
return;
}
if (opcode == opcodes::fss_step + 1 && inputs.size() >= 2
&& inputs[1].data != nullptr)
{
for (std::size_t lane = 0; lane < output.lanes; ++lane)
{
const std::size_t n0 = std::min(output.value_bytes, inputs[0].value_bytes);
std::memcpy(output.at(lane), inputs[0].at(lane), n0);
const std::size_t n1 = std::min(output.value_bytes, inputs[1].value_bytes);
auto * dst = output.at(lane);
const auto * src = inputs[1].at(lane);
for (std::size_t i = 0; i < n1; ++i)
dst[i] = static_cast<std::uint8_t>(dst[i] ^ src[i]);
}
return;
}
if (opcode == opcodes::fss_answer_pack || opcode == opcodes::fss_prefix_pair
|| opcode == opcodes::fss_fuse_trailer || opcode == opcodes::reshare_mask)
{
if (opcode == opcodes::reshare_mask && inputs.size() >= 2
&& inputs[0].value_bytes == 8 && inputs[1].value_bytes == 8
&& output.value_bytes == 8)
{
for (std::size_t lane = 0; lane < output.lanes; ++lane)
{
std::uint64_t a = 0, b = 0;
std::memcpy(&a, inputs[0].at(lane), 8);
std::memcpy(&b, inputs[1].at(lane), 8);
const std::uint64_t s = a + b;
std::memcpy(output.at(lane), &s, 8);
}
return;
}
std::size_t off = 0;
for (const auto & in : inputs)
{
const std::size_t chunk = in.lanes * in.value_bytes;
if (off + chunk > nbytes)
break;
std::memcpy(output.data + off, in.data, chunk);
off += chunk;
}
return;
}
if (is_arith_walk_opcode(opcode))
{
// Local arith defaults: trunc_prob shifts; others copy / XOR-and.
if (opcode == opcodes::trunc_prob && !inputs.empty())
{
for (std::size_t lane = 0; lane < output.lanes; ++lane)
{
if (output.value_bytes == 8 && inputs[0].value_bytes >= 8)
{
std::uint64_t v = 0;
std::memcpy(&v, inputs[0].at(lane), 8);
// Shift amount is applied in run_arith_local via aux.
std::memcpy(output.at(lane), &v, 8);
}
else
{
const std::size_t n = std::min(output.value_bytes,
inputs[0].value_bytes);
std::memcpy(output.at(lane), inputs[0].at(lane), n);
}
}
return;
}
if ((opcode == opcodes::bin_and || opcode == opcodes::bin_rss_and)
&& inputs.size() >= 2)
{
for (std::size_t lane = 0; lane < output.lanes; ++lane)
{
const std::size_t n = std::min({output.value_bytes,
inputs[0].value_bytes, inputs[1].value_bytes});
for (std::size_t i = 0; i < n; ++i)
output.at(lane)[i] = static_cast<std::uint8_t>(
inputs[0].at(lane)[i] & inputs[1].at(lane)[i]);
}
return;
}
if (!inputs.empty())
{
const std::size_t chunk = std::min(nbytes,
inputs[0].lanes * inputs[0].value_bytes);
std::memcpy(output.data, inputs[0].data, chunk);
}
return;
}
const std::size_t chunk = std::min(nbytes, inputs[0].lanes * inputs[0].value_bytes);
std::memcpy(output.data, inputs[0].data, chunk);
};
}
inline void ensure_value_storage(const plan & p,
std::vector<std::vector<std::uint8_t>> & values, std::uint32_t id,
std::size_t lanes);
/// @brief Run one arith/share opcode with party + drive_options (rss seeds, shift).
inline void run_arith_local(const plan & p, std::uint32_t id, std::size_t party,
std::size_t lane, std::size_t lanes,
std::vector<std::vector<std::uint8_t>> & values, const drive_options & opt)
{
const auto op = p.opcode_of(id);
ensure_value_storage(p, values, id, lanes);
const auto vb = p.value_bytes_of(id);
auto * out = values[id].data() + lane * vb;
if (op == opcodes::rss_zero)
{
if (opt.rss_seeds == nullptr)
throw std::runtime_error("compose rss_zero: drive_options::rss_seeds required");
const auto idx = opt.rss_seed_base + p.aux_of(id);
if (vb == 8)
{
const auto z = rss::zero_share<std::uint64_t>(*opt.rss_seeds, idx);
std::memcpy(out, &z, 8);
}
else if (vb == 1)
{
const auto z = rss::zero_share<std::uint8_t>(*opt.rss_seeds, idx);
out[0] = z;
}
else
{
for (std::size_t i = 0; i < vb; ++i)
{
const auto z = rss::zero_share<std::uint8_t>(
*opt.rss_seeds, idx + i);
out[i] = z;
}
}
(void)party;
return;
}
const auto & ins = p.inputs_of(id);
if (op == opcodes::trunc_prob && !ins.empty() && vb == 8
&& p.value_bytes_of(ins[0]) >= 8)
{
ensure_value_storage(p, values, ins[0], lanes);
std::uint64_t v = 0;
std::memcpy(&v, values[ins[0]].data() + lane * p.value_bytes_of(ins[0]), 8);
const unsigned s = static_cast<unsigned>(p.aux_of(id));
v >>= s;
std::memcpy(out, &v, 8);
return;
}
if (op == opcodes::trunc_exact && ins.size() >= 4 && vb >= 8)
{
auto load64 = [&](std::uint32_t src) {
ensure_value_storage(p, values, src, lanes);
std::uint64_t v = 0;
std::memcpy(&v, values[src].data() + lane * p.value_bytes_of(src),
std::min<std::size_t>(8, p.value_bytes_of(src)));
return v;
};
const auto x = load64(ins[0]);
const auto r = load64(ins[1]);
const auto delta = load64(ins[2]);
const auto wrap = load64(ins[3]);
const std::uint64_t r_public = ins.size() >= 5 ? load64(ins[4]) : 0;
const unsigned s = static_cast<unsigned>(p.aux_of(id));
const auto o = trunc::trunc_exact_party<std::uint64_t>(x, r, delta, wrap,
s, static_cast<unsigned>(party), r_public);
std::memcpy(out, &o, 8);
return;
}
if (op == opcodes::bin_a2b && ins.size() >= 3)
{
(void)party;
(void)out;
(void)vb;
throw std::runtime_error(
"bin_a2b opened-r conversion was removed; use a2b_gmw (carry AND)");
}
auto load64 = [&](std::uint32_t src) {
ensure_value_storage(p, values, src, lanes);
std::uint64_t v = 0;
const auto sb = p.value_bytes_of(src);
std::memcpy(&v, values[src].data() + lane * sb, std::min<std::size_t>(8, sb));
return v;
};
auto load8 = [&](std::uint32_t src) {
ensure_value_storage(p, values, src, lanes);
const auto sb = p.value_bytes_of(src);
if (sb == 0)
return static_cast<std::uint8_t>(0);
return values[src][lane * sb];
};
if (op == opcodes::share_input && ins.size() >= 2 && vb >= 8)
{
const auto d = static_cast<std::uint64_t>(load64(ins[0]) - load64(ins[1]));
std::memcpy(out, &d, 8);
return;
}
if (op == opcodes::bin_and && p.aux_of(id) == 1 && ins.size() >= 2)
{
ensure_value_storage(p, values, ins[0], lanes);
ensure_value_storage(p, values, ins[1], lanes);
const auto n = std::min(vb, std::min(p.value_bytes_of(ins[0]),
p.value_bytes_of(ins[1])));
for (std::size_t i = 0; i < n; ++i)
out[i] = static_cast<std::uint8_t>(
values[ins[0]][lane * p.value_bytes_of(ins[0]) + i]
^ values[ins[1]][lane * p.value_bytes_of(ins[1]) + i]);
return;
}
if (op == opcodes::bin_and && p.aux_of(id) == 2 && ins.size() >= 5)
{
ot::bit_triple mine{load8(ins[0]), load8(ins[1]), load8(ins[2])};
out[0] = edabit::and_finish(mine, load8(ins[3]), load8(ins[4]),
static_cast<unsigned>(party));
return;
}
if (op == opcodes::bin_b2a && ins.size() >= 2 && vb >= 8)
{
const unsigned ell = static_cast<unsigned>(p.aux_of(id));
ensure_value_storage(p, values, ins[0], lanes);
ensure_value_storage(p, values, ins[1], lanes);
const auto ab = p.value_bytes_of(ins[0]);
const auto mb = p.value_bytes_of(ins[1]);
const auto * arith = values[ins[0]].data() + lane * ab;
const auto * maskb = values[ins[1]].data() + lane * mb;
std::uint64_t acc = 0;
for (unsigned i = 0; i < ell; ++i)
{
std::uint64_t r = 0;
if (static_cast<std::size_t>(i + 1) * 8 <= ab)
std::memcpy(&r, arith + static_cast<std::size_t>(i) * 8, 8);
const std::uint8_t mask = (static_cast<std::size_t>(i / 8) < mb)
? static_cast<std::uint8_t>((maskb[i / 8] >> (i % 8)) & 1u)
: 0;
ot::dabit<std::uint64_t> dab{0, r};
acc += edabit::b2a_party_bit<std::uint64_t>(0, dab, mask,
static_cast<unsigned>(party), i);
}
std::memcpy(out, &acc, 8);
return;
}
if (op == opcodes::bin_inject && ins.size() >= 5 && vb >= 8)
{
const auto a = load64(ins[0]);
const auto b = load64(ins[1]);
const auto c = load64(ins[2]);
const auto d = load64(ins[3]);
const auto e = load64(ins[4]);
std::uint64_t z = c;
z += d * b;
z += e * a;
if (party == 0)
z += d * e;
std::memcpy(out, &z, 8);
return;
}
if (op == opcodes::bin_rss_and && ins.size() >= 4)
{
if (opt.rss_seeds == nullptr)
throw std::runtime_error("compose bin_rss_and: rss_seeds required");
out[0] = bit_inject::rss_and_local(*opt.rss_seeds, load8(ins[0]),
load8(ins[1]), load8(ins[2]), load8(ins[3]), p.aux_of(id));
return;
}
if (op == opcodes::bin_rss_inject && ins.size() >= 4 && vb >= 8)
{
if (opt.rss_seeds == nullptr)
throw std::runtime_error("compose bin_rss_inject: rss_seeds required");
const auto y = bit_inject::rss_inject_local<std::uint64_t>(*opt.rss_seeds,
load64(ins[0]), load64(ins[1]), load64(ins[2]), load64(ins[3]),
p.aux_of(id));
std::memcpy(out, &y, 8);
return;
}
if (op == opcodes::mul_trunc && ins.size() >= 5 && vb >= 8)
{
const unsigned s = static_cast<unsigned>(p.aux_of(id));
const auto z = trunc::mul_trunc_party<std::uint64_t>(0, 0, load64(ins[0]),
load64(ins[1]), load64(ins[2]), load64(ins[3]), load64(ins[4]), s,
static_cast<unsigned>(party));
std::memcpy(out, &z, 8);
return;
}
// Default: copy first input (or zero).
std::fill(out, out + vb, 0);
if (!ins.empty())
{
ensure_value_storage(p, values, ins[0], lanes);
const auto sb = p.value_bytes_of(ins[0]);
std::memcpy(out, values[ins[0]].data() + lane * sb, std::min(sb, vb));
}
if ((op == opcodes::bin_and || op == opcodes::bin_rss_and) && ins.size() >= 2)
{
ensure_value_storage(p, values, ins[1], lanes);
const auto sb = p.value_bytes_of(ins[1]);
for (std::size_t i = 0; i < std::min(sb, vb); ++i)
out[i] = static_cast<std::uint8_t>(out[i]
& values[ins[1]].data()[lane * sb + i]);
}
}
inline bool has_builtin_walk_kernel(std::uint32_t opcode) noexcept
{
return opcode == opcodes::fss_expand || opcode == opcodes::fss_expand_pair
|| opcode == opcodes::fss_step || opcode == opcodes::fss_step + 1
|| opcode == opcodes::fss_fuse_trailer || opcode == opcodes::fss_fuse_trailer + 1
|| opcode == opcodes::fss_early_stop_pack
|| opcode == opcodes::fss_prefix_emit || opcode == opcodes::fss_prefix_emit + 1
|| opcode == opcodes::fss_prefix_pair
|| opcode == opcodes::fss_prefix_retain || opcode == opcodes::fss_prefix_retain + 1
|| opcode == opcodes::fss_defer_expand || opcode == opcodes::fss_rotate
|| opcode == opcodes::fss_leaf_later || opcode == opcodes::fss_apply_leaf
|| opcode == opcodes::fss_answer_pack
|| opcode == opcodes::ds_blind || opcode == opcodes::ds_cw
|| opcode == opcodes::ds_advice || opcode == opcodes::ds_and
|| opcode == opcodes::ds_and2
|| opcode == opcodes::ds_oh || opcode == opcodes::ds_leaf_mux
|| opcode == opcodes::ds_leaf_mux + 1
|| opcode == opcodes::reshare_mask || opcode == opcodes::rss_mul
|| opcode == opcodes::client_answer
|| is_arith_walk_opcode(opcode);
}
/// @brief Channel for one exchange node (pair / RSS neighbor / dealer).
inline edge_channel exchange_channel(const plan & p, std::uint32_t ex_id) noexcept
{
const auto op = p.opcode_of(ex_id);
if (op == opcodes::dealer_pad || op == opcodes::dealer_zero)
return edge_channel::dealer;
if (op == opcodes::shuffle_send || op == opcodes::bench_ring)
return edge_channel::rss_next;
if (is_rss_neighbor_open(p, ex_id))
return edge_channel::rss_next;
return edge_channel::peer;
}
/// @brief Receive rule for one exchange node.
inline receive_rule exchange_receive_rule(const plan & p,
std::uint32_t ex_id) noexcept
{
const auto op = p.opcode_of(ex_id);
if (is_beaver_opcode(op))
return receive_rule::beaver;
if (op == opcodes::client_upload || op == opcodes::client_answer
|| op == opcodes::dealer_pad || op == opcodes::dealer_zero
|| op == opcodes::bench_wire || op == opcodes::bench_ring)
return receive_rule::copy_peer;
if (is_rss_neighbor_open(p, ex_id))
return receive_rule::copy_peer;
// Aux high bit marks special receive rules recorded by the composer.
const auto aux = p.aux_of(ex_id);
if ((aux & 0xff000000u) == 0x01000000u)
return receive_rule::field_sum;
if ((aux & 0xff000000u) == 0x02000000u)
return receive_rule::any_two;
if ((aux & 0xff000000u) == 0x03000000u)
return receive_rule::verify_sketch;
if ((aux & 0xff000000u) == 0x04000000u)
return receive_rule::verify_proof;
if ((aux & 0xff000000u) == 0x05000000u)
return receive_rule::eq_check;
return receive_rule::domain_open;
}
/// @brief Dominant channel for a packed exchange wave (dealer > rss > peer).
inline edge_channel wave_channel(const plan & p, const wave_info & wave) noexcept
{
edge_channel ch = edge_channel::peer;
for (auto ex : wave.exchanges)
{
const auto c = exchange_channel(p, ex.id);
if (c == edge_channel::dealer)
return edge_channel::dealer;
if (c == edge_channel::rss_next)
ch = edge_channel::rss_next;
}
return ch;
}
inline receive_rule wave_receive_rule(const plan & p,
const wave_info & wave) noexcept
{
for (auto ex : wave.exchanges)
{
const auto r = exchange_receive_rule(p, ex.id);
if (r != receive_rule::domain_open)
return r;
}
return receive_rule::domain_open;
}
inline void ensure_value_storage(const plan & p,
std::vector<std::vector<std::uint8_t>> & values, std::uint32_t id,
std::size_t lanes)
{
const std::size_t need = lanes * p.value_bytes_of(id);
auto & v = values[id];
if (v.size() >= need)
return;
// An empty input runs as zeros; a short one is a caller bug.
if (!v.empty() && p.effect_of(id) == effect::input)
throw std::invalid_argument("compose: input node " + std::to_string(id)
+ " has " + std::to_string(v.size()) + " bytes; "
+ std::to_string(lanes) + " instance(s) of "
+ std::to_string(p.value_bytes_of(id)) + " bytes need "
+ std::to_string(need));
v.assign(need, 0);
}
/// @brief Apply one packed peer slot to `exchanges` (in slot order).
inline void apply_peer_exchanges(const plan & p, const std::vector<node> & exchanges,
std::size_t slot_bytes, std::size_t party, std::size_t lane, std::size_t lanes,
const std::uint8_t * peer, std::size_t peer_n,
std::vector<std::vector<std::uint8_t>> & values, const drive_options & opt)
{
if (exchanges.empty() || peer == nullptr)
return;
if (peer_n < slot_bytes)
throw std::runtime_error("compose apply_peer_slot: short peer");
std::size_t off = 0;
for (auto ex : exchanges)
{
const auto nb = p.value_bytes_of(ex.id);
ensure_value_storage(p, values, ex.id, lanes);
const auto op = p.opcode_of(ex.id);
const auto rule = exchange_receive_rule(p, ex.id);
if (rule == receive_rule::beaver && opt.beavers != nullptr)
{
const auto sess = static_cast<std::size_t>(p.level_of(ex.id));
const auto bi = static_cast<std::size_t>(op - opcodes::beaver_delta);
opt.beavers->apply(sess, bi, peer + off, nb);
std::vector<std::uint8_t> mine(nb, 0);
std::memcpy(mine.data(), values[ex.id].data() + lane * nb, nb);
reconstruct_open(domain::a, party, mine.data(), peer + off, nb,
values[ex.id].data() + lane * nb, opt.field);
}
else if (rule == receive_rule::copy_peer
|| rule == receive_rule::xor_bytes)
{
if (rule == receive_rule::xor_bytes)
{
auto * dst = values[ex.id].data() + lane * nb;
for (std::size_t i = 0; i < nb; ++i)
dst[i] = static_cast<std::uint8_t>(dst[i] ^ peer[off + i]);
}
else
{
std::memcpy(values[ex.id].data() + lane * nb, peer + off, nb);
}
}
else if (rule == receive_rule::field_sum
|| rule == receive_rule::any_two)
{
// field_sum / any_two: same word-lane add as domain::a (Shamir
// any-two of two shares collapses to a sum of the held pair).
std::vector<std::uint8_t> mine(nb, 0);
const auto & ins = p.inputs_of(ex.id);
if (!ins.empty())
{
const auto sb = p.value_bytes_of(ins[0]);
ensure_value_storage(p, values, ins[0], lanes);
std::memcpy(mine.data(), values[ins[0]].data() + lane * sb,
std::min(sb, nb));
}
else
{
std::memcpy(mine.data(), values[ex.id].data() + lane * nb, nb);
}
reconstruct_open(domain::a, party, mine.data(), peer + off, nb,
values[ex.id].data() + lane * nb,
rule == receive_rule::field_sum ? field_open::fp61 : opt.field);
}
else if (rule == receive_rule::verify_sketch
|| rule == receive_rule::verify_proof
|| rule == receive_rule::eq_check)
{
// Store peer bytes; host/kernel verifies (sketch/proof/tag).
std::memcpy(values[ex.id].data() + lane * nb, peer + off, nb);
if (rule == receive_rule::eq_check)
{
const auto & ins = p.inputs_of(ex.id);
if (!ins.empty())
{
ensure_value_storage(p, values, ins[0], lanes);
const auto sb = p.value_bytes_of(ins[0]);
const auto n = std::min(sb, nb);
if (std::memcmp(values[ins[0]].data() + lane * sb,
peer + off, n)
!= 0)
throw std::runtime_error(
"compose apply_peer_slot: eq_check failed");
}
}
}
else
{
const auto & ins = p.inputs_of(ex.id);
domain d = p.domain_of(ex.id);
if (!ins.empty())
d = p.domain_of(ins[0]);
std::vector<std::uint8_t> mine(nb, 0);
if (!ins.empty())
{
const auto sb = p.value_bytes_of(ins[0]);
ensure_value_storage(p, values, ins[0], lanes);
std::memcpy(mine.data(), values[ins[0]].data() + lane * sb,
std::min(sb, nb));
}
reconstruct_open(d, party, mine.data(), peer + off, nb,
values[ex.id].data() + lane * nb, opt.field);
}
if (opt.cleartext && opt.clear_oracle != nullptr)
{
auto it = opt.clear_oracle->find(ex.id);
if (it != opt.clear_oracle->end() && it->second.size() == nb
&& std::memcmp(values[ex.id].data() + lane * nb,
it->second.data(), nb)
!= 0)
throw std::runtime_error(
"compose cleartext: open mismatch vs clear_oracle");
}
if (opt.check_open)
{
// Consistency: hash of reconstructed word must match oracle when
// provided; otherwise store a fingerprint peers can compare.
if (opt.clear_oracle != nullptr)
{
auto it = opt.clear_oracle->find(ex.id);
if (it != opt.clear_oracle->end() && it->second.size() == nb
&& std::memcmp(values[ex.id].data() + lane * nb,
it->second.data(), nb)
!= 0)
throw std::runtime_error(
"compose check_open: reconstructed open rejected");
}
std::uint64_t h = 14695981039346656037ull;
for (std::size_t i = 0; i < nb; ++i)
{
h ^= values[ex.id][lane * nb + i];
h *= 1099511628211ull;
}
// Peer hash lives in trailing bytes of the peer slot when present.
if (peer_n >= off + nb + 8)
{
std::uint64_t peer_h = 0;
std::memcpy(&peer_h, peer + off + nb, 8);
if (peer_h != 0 && peer_h != h)
throw std::runtime_error(
"compose check_open: hash exchange mismatch");
}
}
off += nb;
}
}
/// @brief Apply one packed peer slot to the exchange nodes of `wave`.
inline void apply_peer_slot(const plan & p, const wave_info & wave,
std::size_t party, std::size_t lane, std::size_t lanes,
const std::uint8_t * peer, std::size_t peer_n,
std::vector<std::vector<std::uint8_t>> & values, const drive_options & opt)
{
apply_peer_exchanges(p, wave.exchanges, wave.slot_bytes, party, lane, lanes,
peer, peer_n, values, opt);
}
/// @brief Pack one lane's payloads for `exchanges` into `out`.
inline void pack_exchanges(const plan & p, const std::vector<node> & exchanges,
std::size_t lane, std::size_t lanes,
std::vector<std::vector<std::uint8_t>> & values, const drive_options & opt,
std::uint8_t * out)
{
std::size_t off = 0;
for (auto ex : exchanges)
{
ensure_value_storage(p, values, ex.id, lanes);
const auto nb = p.value_bytes_of(ex.id);
const auto op = p.opcode_of(ex.id);
if (op == opcodes::dealer_pad || op == opcodes::dealer_zero)
throw std::runtime_error(
"compose: dealer waves need a dealer edge; drive with "
"drive_via_schedule(plan, edge_mesh) or app::run_parties");
if (is_beaver_opcode(op) && opt.beavers != nullptr)
{
const auto sess = static_cast<std::size_t>(p.level_of(ex.id));
const auto bi = static_cast<std::size_t>(op - opcodes::beaver_delta);
opt.beavers->pack(sess, bi, out + off, nb);
std::memcpy(values[ex.id].data() + lane * nb, out + off, nb);
}
else
{
const auto & ins = p.inputs_of(ex.id);
if (!ins.empty() && !is_beaver_opcode(op))
{
const auto src = ins[0];
ensure_value_storage(p, values, src, lanes);
const auto sb = p.value_bytes_of(src);
std::memcpy(out + off, values[src].data() + lane * sb,
std::min(sb, nb));
std::memcpy(values[ex.id].data() + lane * nb,
values[src].data() + lane * sb, std::min(sb, nb));
}
else
{
std::memcpy(out + off, values[ex.id].data() + lane * nb, nb);
}
}
off += nb;
}
}
/// @brief Pack one lane's exchange payloads into `out` (size `wave.slot_bytes`).
inline void pack_exchange_lane(const plan & p, const wave_info & wave,
std::size_t lane, std::size_t lanes,
std::vector<std::vector<std::uint8_t>> & values, const drive_options & opt,
std::uint8_t * out)
{
pack_exchanges(p, wave.exchanges, lane, lanes, values, opt, out);
}
/// @brief Bench-cell local work. Returns true when this group was a cell.
inline bool dispatch_cell(const plan & p, const simd_group & g, std::size_t party,
std::size_t lane, std::size_t lanes,
std::vector<std::vector<std::uint8_t>> & values, const drive_options & opt)
{
if (opt.cell == nullptr)
return false;
if (g.opcode != opcodes::bench_emit && g.opcode != opcodes::bench_apply)
return false;
for (auto n : g.nodes)
{
ensure_value_storage(p, values, n.id, lanes);
const auto vb = p.value_bytes_of(n.id);
opt.cell(g.opcode, party, p.aux_of(n.id),
values[n.id].data() + lane * vb, vb);
}
return true;
}
/// @brief Run local groups of `wave` for a single lane (schedule_session path).
inline void run_wave_locals_lane(const plan & p, const wave_info & wave,
std::size_t party, std::size_t lane, std::size_t lanes,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels,
const drive_options & opt = {})
{
for (const auto & g : wave.groups)
{
if (is_arith_walk_opcode(g.opcode))
{
for (auto n : g.nodes)
run_arith_local(p, n.id, party, lane, lanes, values, opt);
continue;
}
if (g.kind == effect::convert)
{
for (auto n : g.nodes)
{
ensure_value_storage(p, values, n.id, lanes);
const auto & ins = p.inputs_of(n.id);
if (ins.empty())
continue;
if (p.alias_of(n.id) && ins.size() == 1
&& p.value_bytes_of(n.id) == p.value_bytes_of(ins[0]))
{
const auto vb = p.value_bytes_of(n.id);
std::memcpy(values[n.id].data() + lane * vb,
values[ins[0]].data() + lane * vb, vb);
continue;
}
if (ins.size() == 1)
{
run_builtin_convert(g.opcode, party,
values[ins[0]].data() + lane * p.value_bytes_of(ins[0]),
p.value_bytes_of(ins[0]),
values[n.id].data() + lane * p.value_bytes_of(n.id),
p.value_bytes_of(n.id));
}
else if (ins.size() == 2 && g.opcode == opcodes::conv_y2rss)
{
auto * outp = values[n.id].data() + lane * p.value_bytes_of(n.id);
const auto half = p.value_bytes_of(ins[0]);
std::memcpy(outp, values[ins[0]].data() + lane * half, half);
std::memcpy(outp + half, values[ins[1]].data() + lane * half,
half);
}
}
continue;
}
if (is_beaver_opcode(g.opcode))
continue;
if (g.opcode == opcodes::dealer_pad || g.opcode == opcodes::dealer_zero)
throw std::runtime_error(
"compose: dealer delivery needs a dealer edge sink");
if (g.opcode == opcodes::fss_fuse_segment)
{
for (auto n : g.nodes)
{
ensure_value_storage(p, values, n.id, lanes);
apply_fuse_segment(p, values, n.id, lanes);
}
continue;
}
if (dispatch_cell(p, g, party, lane, lanes, values, opt))
continue;
auto kit = kernels.find(g.opcode);
kernel_fn builtin;
const kernel_fn * fn = nullptr;
if (kit != kernels.end())
fn = &kit->second;
else if (has_builtin_walk_kernel(g.opcode))
{
builtin = builtin_walk_kernel(g.opcode);
fn = &builtin;
}
else if (g.kind == effect::blind)
{
for (auto n : g.nodes)
{
ensure_value_storage(p, values, n.id, lanes);
const auto vb = p.value_bytes_of(n.id);
std::fill(values[n.id].data() + lane * vb,
values[n.id].data() + lane * vb + vb, 0);
}
continue;
}
else
{
throw std::runtime_error(
"compose missing kernel for opcode " + std::to_string(g.opcode));
}
for (auto n : g.nodes)
{
ensure_value_storage(p, values, n.id, lanes);
std::vector<block_span> ins;
for (auto in_id : p.inputs_of(n.id))
{
ensure_value_storage(p, values, in_id, lanes);
ins.push_back(block_span{values[in_id].data() + lane * p.value_bytes_of(in_id),
1, p.value_bytes_of(in_id)});
}
block_span local_out{values[n.id].data() + lane * p.value_bytes_of(n.id),
1, p.value_bytes_of(n.id)};
invoke_kernel(opt, *fn, g.opcode, {n}, ins, local_out, 1);
}
}
}
inline std::vector<std::size_t> exchange_wave_indices(const plan & p)
{
std::vector<std::size_t> idx;
for (std::size_t w = 0; w < p.waves(); ++w)
if (!p.wave(w).exchanges.empty())
idx.push_back(w);
return idx;
}
} // namespace detail
inline void drive(const plan & p, net::RoundSink & sink,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party = 0,
const drive_options & opt = {})
{
const std::size_t lanes = sink.count();
if (values.size() < p.nodes().size())
values.resize(p.nodes().size());
auto ensure_storage = [&](std::uint32_t id) {
detail::ensure_value_storage(p, values, id, lanes);
};
for (std::uint32_t id = 0; id < p.nodes().size(); ++id)
ensure_storage(id);
std::size_t exchange_i = 0;
std::size_t wave_begin = 0;
bool resume_recv = false;
std::uint16_t resume_round = 0;
if (opt.cursor != nullptr && !opt.cursor->done)
{
wave_begin = opt.cursor->wave;
exchange_i = opt.cursor->exchange_i;
resume_recv = opt.cursor->awaiting_peer;
resume_round = opt.cursor->sink_round;
}
auto park = [&](std::size_t wave, std::uint16_t sink_round) {
if (opt.cursor == nullptr)
return;
opt.cursor->wave = wave;
opt.cursor->exchange_i = exchange_i;
opt.cursor->sink_round = sink_round;
opt.cursor->awaiting_peer = true;
opt.cursor->done = false;
};
for (std::size_t w = wave_begin; w < p.waves(); ++w)
{
const auto & wave = p.wave(w);
const bool skip_wave = !resume_recv && !wave.exchanges.empty()
&& exchange_i < opt.from_exchange_wave;
if (!skip_wave && !resume_recv)
{
for (const auto & g : wave.groups)
{
if (detail::is_arith_walk_opcode(g.opcode))
{
for (auto n : g.nodes)
for (std::size_t lane = 0; lane < lanes; ++lane)
detail::run_arith_local(p, n.id, party, lane, lanes,
values, opt);
continue;
}
if (g.kind == effect::convert)
{
for (auto n : g.nodes)
{
ensure_storage(n.id);
const auto & ins = p.inputs_of(n.id);
if (ins.empty())
continue;
if (p.alias_of(n.id) && ins.size() == 1
&& p.value_bytes_of(n.id) == p.value_bytes_of(ins[0]))
{
values[n.id] = values[ins[0]];
continue;
}
for (std::size_t lane = 0; lane < lanes; ++lane)
{
if (ins.size() == 1)
{
detail::run_builtin_convert(g.opcode, party,
values[ins[0]].data()
+ lane * p.value_bytes_of(ins[0]),
p.value_bytes_of(ins[0]),
values[n.id].data()
+ lane * p.value_bytes_of(n.id),
p.value_bytes_of(n.id));
}
else if (ins.size() == 2
&& g.opcode == opcodes::conv_y2rss)
{
auto * out = values[n.id].data()
+ lane * p.value_bytes_of(n.id);
const auto half = p.value_bytes_of(ins[0]);
std::memcpy(out,
values[ins[0]].data() + lane * half, half);
std::memcpy(out + half,
values[ins[1]].data() + lane * half, half);
}
}
}
continue;
}
if (detail::is_beaver_opcode(g.opcode))
continue; // δ results land via beaver_host / exchange path
if (g.opcode == opcodes::dealer_pad || g.opcode == opcodes::dealer_zero)
throw std::runtime_error(
"compose::drive: dealer delivery requires drive_composed_trio");
if (g.opcode == opcodes::fss_fuse_segment)
{
for (auto n : g.nodes)
{
ensure_storage(n.id);
const auto & ins = p.inputs_of(n.id);
if (!ins.empty())
ensure_storage(ins[0]);
detail::apply_fuse_segment(p, values, n.id, lanes);
}
continue;
}
if (opt.cell != nullptr
&& (g.opcode == opcodes::bench_emit
|| g.opcode == opcodes::bench_apply))
{
for (auto n : g.nodes)
{
ensure_storage(n.id);
const auto vb = p.value_bytes_of(n.id);
for (std::size_t lane = 0; lane < lanes; ++lane)
opt.cell(g.opcode, party, p.aux_of(n.id),
values[n.id].data() + lane * vb, vb);
}
continue;
}
auto kit = kernels.find(g.opcode);
kernel_fn builtin;
const kernel_fn * fn = nullptr;
if (kit != kernels.end())
fn = &kit->second;
else if (detail::has_builtin_walk_kernel(g.opcode))
{
builtin = detail::builtin_walk_kernel(g.opcode);
fn = &builtin;
}
else if (g.kind == effect::blind)
{
for (auto n : g.nodes)
{
ensure_storage(n.id);
std::fill(values[n.id].begin(), values[n.id].end(), 0);
}
continue;
}
else
{
throw std::runtime_error(
"compose::drive missing kernel for opcode "
+ std::to_string(g.opcode));
}
for (auto n : g.nodes)
ensure_storage(n.id);
bool uniform = !g.nodes.empty();
for (auto n : g.nodes)
{
if (p.inputs_of(n.id).size() != 1
|| p.value_bytes_of(n.id)
!= p.value_bytes_of(g.nodes[0].id))
{
uniform = false;
break;
}
}
if (uniform && g.nodes.size() > 1)
{
const auto vb = p.value_bytes_of(g.nodes[0].id);
const std::size_t packed_lanes = g.nodes.size() * lanes;
std::vector<std::uint8_t> packed_in(packed_lanes * vb);
std::vector<std::uint8_t> packed_out(packed_lanes * vb);
for (std::size_t ni = 0; ni < g.nodes.size(); ++ni)
{
const auto in_id = p.inputs_of(g.nodes[ni].id)[0];
ensure_storage(in_id);
std::memcpy(packed_in.data() + ni * lanes * vb,
values[in_id].data(), lanes * vb);
}
block_span in_span{packed_in.data(), packed_lanes, vb};
block_span out_span{packed_out.data(), packed_lanes, vb};
invoke_kernel(opt, *fn, g.opcode, g.nodes, {in_span}, out_span,
lanes);
for (std::size_t ni = 0; ni < g.nodes.size(); ++ni)
{
std::memcpy(values[g.nodes[ni].id].data(),
packed_out.data() + ni * lanes * vb, lanes * vb);
}
}
else
{
for (auto n : g.nodes)
{
std::vector<block_span> ins;
for (auto in_id : p.inputs_of(n.id))
{
ensure_storage(in_id);
ins.push_back(block_span{values[in_id].data(), lanes,
p.value_bytes_of(in_id)});
}
block_span local_out{values[n.id].data(), lanes,
p.value_bytes_of(n.id)};
invoke_kernel(opt, *fn, g.opcode, {n}, ins, local_out, lanes);
}
}
}
}
if (wave.exchanges.empty() || wave.slot_bytes == 0)
{
resume_recv = false;
continue;
}
if (!resume_recv && exchange_i < opt.from_exchange_wave)
{
++exchange_i;
continue;
}
std::uint16_t sink_round = 0;
if (resume_recv)
sink_round = resume_round;
else
{
sink_round = static_cast<std::uint16_t>(
opt.compact_sink ? (exchange_i - opt.from_exchange_wave)
: exchange_i);
++exchange_i;
for (std::size_t lane = 0; lane < lanes; ++lane)
{
std::vector<std::uint8_t> buf(wave.slot_bytes, 0);
std::size_t off = 0;
for (auto ex : wave.exchanges)
{
ensure_storage(ex.id);
const auto nb = p.value_bytes_of(ex.id);
const auto op = p.opcode_of(ex.id);
if (op == opcodes::dealer_pad || op == opcodes::dealer_zero)
throw std::runtime_error(
"compose::drive: dealer waves need a dealer edge; drive "
"with drive_via_schedule(plan, edge_mesh) or app::run_parties");
if (detail::is_beaver_opcode(op) && opt.beavers != nullptr)
{
const auto sess = static_cast<std::size_t>(p.level_of(ex.id));
const auto bi = static_cast<std::size_t>(op - opcodes::beaver_delta);
opt.beavers->pack(sess, bi, buf.data() + off, nb);
std::memcpy(values[ex.id].data() + lane * nb,
buf.data() + off, nb);
}
else
{
const auto & ins = p.inputs_of(ex.id);
if (!ins.empty() && !detail::is_beaver_opcode(op))
{
const auto src = ins[0];
ensure_storage(src);
const auto sb = p.value_bytes_of(src);
std::memcpy(buf.data() + off,
values[src].data() + lane * sb,
std::min(sb, nb));
std::memcpy(values[ex.id].data() + lane * nb,
values[src].data() + lane * sb,
std::min(sb, nb));
}
else
{
std::memcpy(buf.data() + off,
values[ex.id].data() + lane * nb, nb);
}
}
off += nb;
}
sink.submit(sink_round, lane, buf.data(), buf.size());
}
}
if (!resume_recv)
{
sink.flush();
sink.poll();
}
if (opt.park_if_waiting)
{
bool ready = true;
for (std::size_t lane = 0; lane < lanes; ++lane)
{
if (!sink.peer_ready(sink_round, lane))
ready = false;
}
if (!ready)
{
park(w, sink_round);
return;
}
}
for (std::size_t lane = 0; lane < lanes; ++lane)
{
detail::wait_ready(sink,
wait_budget(opt, edge_peer, round_budget(opt, sink_round)),
[&] { return sink.peer_ready(sink_round, lane); },
"compose::drive round " + std::to_string(sink_round));
std::vector<std::uint8_t> peer(wave.slot_bytes);
sink.read_peer(sink_round, lane, peer.data(), peer.size());
std::size_t off = 0;
for (auto ex : wave.exchanges)
{
const auto nb = p.value_bytes_of(ex.id);
ensure_storage(ex.id);
const auto op = p.opcode_of(ex.id);
if (detail::is_beaver_opcode(op) && opt.beavers != nullptr)
{
const auto sess = static_cast<std::size_t>(p.level_of(ex.id));
const auto bi = static_cast<std::size_t>(op - opcodes::beaver_delta);
opt.beavers->apply(sess, bi, peer.data() + off, nb);
// δ sum into the exchange node for downstream compute.
const auto & ins = p.inputs_of(ex.id);
std::vector<std::uint8_t> mine(nb, 0);
std::memcpy(mine.data(), values[ex.id].data() + lane * nb, nb);
detail::reconstruct_open(domain::a, party, mine.data(),
peer.data() + off, nb, values[ex.id].data() + lane * nb,
opt.field);
(void)ins;
}
else
{
const auto & ins = p.inputs_of(ex.id);
domain d = p.domain_of(ex.id);
if (!ins.empty())
d = p.domain_of(ins[0]);
const bool client_msg = op == opcodes::client_upload
|| op == opcodes::client_answer;
if (detail::is_rss_neighbor_open(p, ex.id)
|| detail::is_beaver_opcode(op) || client_msg)
{
if (detail::is_beaver_opcode(op) && !client_msg)
{
std::vector<std::uint8_t> mine(nb, 0);
if (!ins.empty())
{
const auto sb = p.value_bytes_of(ins[0]);
std::memcpy(mine.data(),
values[ins[0]].data() + lane * sb,
std::min(sb, nb));
}
else
{
std::memcpy(mine.data(),
values[ex.id].data() + lane * nb, nb);
}
detail::reconstruct_open(domain::a, party,
mine.data(), peer.data() + off, nb,
values[ex.id].data() + lane * nb, opt.field);
}
else
{
std::memcpy(values[ex.id].data() + lane * nb,
peer.data() + off, nb);
}
}
else
{
std::vector<std::uint8_t> mine(nb, 0);
if (!ins.empty())
{
const auto sb = p.value_bytes_of(ins[0]);
std::memcpy(mine.data(),
values[ins[0]].data() + lane * sb,
std::min(sb, nb));
}
detail::reconstruct_open(d, party, mine.data(),
peer.data() + off, nb,
values[ex.id].data() + lane * nb, opt.field);
}
}
if ((opt.cleartext || opt.check_open) && opt.clear_oracle != nullptr)
{
auto it = opt.clear_oracle->find(ex.id);
if (it != opt.clear_oracle->end() && it->second.size() == nb
&& std::memcmp(values[ex.id].data() + lane * nb,
it->second.data(), nb)
!= 0)
throw std::runtime_error(
opt.check_open
? "compose check_open: reconstructed open rejected"
: "compose cleartext: open mismatch vs clear_oracle");
}
off += nb;
}
}
resume_recv = false;
if (opt.one_exchange)
{
if (opt.cursor != nullptr)
{
opt.cursor->wave = w + 1;
opt.cursor->exchange_i = exchange_i;
opt.cursor->awaiting_peer = false;
opt.cursor->done = false;
}
return;
}
}
if (opt.cursor != nullptr)
{
opt.cursor->done = true;
opt.cursor->awaiting_peer = false;
}
}
namespace detail
{
/// @brief The exchanges of one wave that travel on one channel.
struct wave_part
{
std::size_t wave = 0;
edge_channel channel = edge_channel::peer;
std::vector<node> exchanges;
std::size_t slot_bytes = 0;
};
/// @brief Exchange waves split per channel (dealer, then rss_next, then peer).
/// @details A wave that mixes dealer, RSS, and 2PC exchanges becomes one round
/// per channel; exchanges in one wave never depend on each other, so
/// the order only fixes the wire layout.
inline std::vector<wave_part> schedule_parts(const plan & p)
{
std::vector<wave_part> parts;
for (const auto w : exchange_wave_indices(p))
{
const auto & wave = p.wave(w);
wave_part by[3];
for (auto ex : wave.exchanges)
{
const auto c = exchange_channel(p, ex.id);
auto & part = by[static_cast<int>(c)];
part.wave = w;
part.channel = c;
part.exchanges.push_back(ex);
part.slot_bytes += p.value_bytes_of(ex.id);
}
for (int c : {2, 1, 0})
if (!by[c].exchanges.empty())
parts.push_back(std::move(by[c]));
}
return parts;
}
/// @brief One schedule round from `party`'s point of view.
/// @details Parties 0 and 1 run every part: peer parts with each other, RSS
/// parts on their ring edge, dealer parts as receivers. Party 2 is the
/// dealer / third party: idle on 2PC opens, on the ring for RSS parts,
/// and the sender of dealer parts (one round to party 0 on
/// `edge_dealer`, one to party 1 on `edge_dealer_p1`).
struct party_step
{
wave_part part;
bool dealer_side = false;
unsigned dest = 0;
edge_id edge = edge_peer;
std::uint16_t sink_round = 0;
};
inline std::vector<party_step> party_steps(const plan & p, std::size_t party)
{
const auto parts = schedule_parts(p);
std::vector<party_step> out;
if (party > 2)
{
if (!parts.empty())
throw std::invalid_argument("plan_to_schedule: party "
+ std::to_string(party) + " has no role (0 and 1 are online, "
"2 is the dealer / third party)");
return out;
}
auto add = [&](const wave_part & pt, bool dealer_side, unsigned dest,
edge_id e) {
party_step st;
st.part = pt;
st.dealer_side = dealer_side;
st.dest = dest;
st.edge = e;
out.push_back(std::move(st));
};
bool peer_only = true;
for (const auto & pt : parts)
{
switch (pt.channel)
{
case edge_channel::peer:
if (party != 2)
add(pt, false, 0, edge_peer);
break;
case edge_channel::rss_next:
peer_only = false;
add(pt, false, 0, edge_rss_next);
break;
case edge_channel::dealer:
peer_only = false;
if (party == 2)
{
add(pt, true, 0, edge_dealer);
add(pt, true, 1, edge_dealer_p1);
}
else
add(pt, false, 0, edge_dealer);
break;
}
}
if (party == 2 && peer_only && !parts.empty())
throw std::invalid_argument("plan_to_schedule: party 2 has no dealer or "
"rss_next rounds; 2PC opens run between parties 0 and 1");
std::map<edge_id, std::uint16_t> next;
for (auto & st : out)
st.sink_round = next[st.edge]++;
return out;
}
inline receive_rule part_receive_rule(const plan & p, const wave_part & part)
{
for (auto ex : part.exchanges)
{
const auto r = exchange_receive_rule(p, ex.id);
if (r != receive_rule::domain_open)
return r;
}
return receive_rule::domain_open;
}
using zero_share_cache =
std::map<std::pair<std::uint32_t, std::size_t>, std::vector<std::uint8_t>>;
/// @brief Dealer side of a dealer part: the payload, or a zero share for `dest`.
inline void pack_dealer_part(const plan & p, const wave_part & part, unsigned dest,
std::size_t lane, std::size_t lanes,
std::vector<std::vector<std::uint8_t>> & values, zero_share_cache & zeros,
std::uint8_t * out)
{
std::size_t off = 0;
for (auto ex : part.exchanges)
{
ensure_value_storage(p, values, ex.id, lanes);
const auto nb = p.value_bytes_of(ex.id);
auto * dst = out + off;
if (p.opcode_of(ex.id) == opcodes::dealer_pad)
{
const auto & ins = p.inputs_of(ex.id);
const std::uint8_t * src = values[ex.id].data() + lane * nb;
std::size_t n = nb;
if (!ins.empty())
{
ensure_value_storage(p, values, ins[0], lanes);
const auto sb = p.value_bytes_of(ins[0]);
src = values[ins[0]].data() + lane * sb;
n = std::min(sb, nb);
}
std::memset(dst, 0, nb);
std::memcpy(dst, src, n);
std::memcpy(values[ex.id].data() + lane * nb, dst, nb);
}
else
{
const auto key = std::make_pair(ex.id, lane);
if (dest == 0)
{
std::vector<std::uint8_t> a(nb), b(nb);
for (auto & byte : a)
byte = dpf::uniform_sample<std::uint8_t>();
if (p.domain_of(ex.id) != domain::b && nb % 8 == 0)
{
for (std::size_t i = 0; i < nb; i += 8)
{
std::uint64_t w = 0;
std::memcpy(&w, a.data() + i, 8);
w = static_cast<std::uint64_t>(0) - w;
std::memcpy(b.data() + i, &w, 8);
}
}
else
b = a;
std::memcpy(dst, a.data(), nb);
zeros[key] = std::move(b);
}
else
{
auto it = zeros.find(key);
if (it == zeros.end())
throw std::logic_error(
"dealer: zero share for party 1 before party 0");
std::memcpy(dst, it->second.data(), nb);
zeros.erase(it);
}
}
off += nb;
}
}
/// @brief Last round's peer slot and trailing compute-only waves.
inline void finish_steps(const plan & p, const std::vector<party_step> & steps,
edge_mesh & mesh, std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party,
std::size_t lanes, const drive_options & opt)
{
if (values.size() < p.nodes().size())
values.resize(p.nodes().size());
for (std::uint32_t id = 0; id < p.nodes().size(); ++id)
ensure_value_storage(p, values, id, lanes);
if (steps.empty())
{
for (std::size_t wi = 0; wi < p.waves(); ++wi)
for (std::size_t lane = 0; lane < lanes; ++lane)
run_wave_locals_lane(p, p.wave(wi), party, lane, lanes, values,
kernels, opt);
return;
}
const auto & last = steps.back();
auto & sink = mesh.at(last.edge);
sink.flush();
sink.poll();
const auto budget = wait_budget(opt, last.edge,
round_budget(opt, steps.size() - 1));
for (std::size_t lane = 0; lane < lanes; ++lane)
{
wait_ready(sink, budget,
[&] {
if (sink.peer_ready(last.sink_round, lane))
return true;
sink.flush();
return false;
},
"finish on " + net::edge_name(last.edge));
std::vector<std::uint8_t> peer(last.part.slot_bytes);
if (!peer.empty())
sink.read_peer(last.sink_round, lane, peer.data(), peer.size());
if (!last.dealer_side)
apply_peer_exchanges(p, last.part.exchanges, last.part.slot_bytes,
party, lane, lanes, peer.data(), peer.size(), values, opt);
}
for (std::size_t wi = last.part.wave + 1; wi < p.waves(); ++wi)
for (std::size_t lane = 0; lane < lanes; ++lane)
run_wave_locals_lane(p, p.wave(wi), party, lane, lanes, values,
kernels, opt);
}
/// @brief Drive `sess` to completion under per-instance round budgets.
/// @details An instance's clock starts when it enters a round and resets when
/// it advances, so a long healthy protocol never times out; only a
/// round whose peer bytes do not arrive does. Completions for other
/// edges do not reset it.
inline void run_session(schedule_session & sess, const drive_options & opt,
const std::string & what)
{
using clock = std::chrono::steady_clock;
const std::size_t count = sess.count();
std::vector<std::uint64_t> mark(count, ~std::uint64_t{0});
std::vector<clock::time_point> since(count, clock::now());
std::vector<net::RoundSink *> sinks;
for (auto * s : sess.mesh().sinks)
if (s != nullptr && std::find(sinks.begin(), sinks.end(), s) == sinks.end())
sinks.push_back(s);
unsigned spins = 0;
for (;;)
{
sess.drive();
const auto now = clock::now();
bool all_done = true;
bool advanced = false;
bool bounded = false;
std::chrono::milliseconds slice(1000);
for (std::size_t i = 0; i < count; ++i)
{
if (sess.done(i))
continue;
all_done = false;
const auto m = sess.mark(i);
if (m != mark[i])
{
mark[i] = m;
since[i] = now;
advanced = true;
}
const auto budget = wait_budget(opt, sess.current_edge(i),
sess.current_timeout(i));
if (budget.count() == 0)
continue;
bounded = true;
const auto waited =
std::chrono::duration_cast<std::chrono::milliseconds>(now - since[i]);
if (waited >= budget)
throw std::runtime_error(what + ": instance " + std::to_string(i)
+ " waited " + std::to_string(waited.count()) + " ms in round "
+ std::to_string(sess.current_round(i)) + " on "
+ net::edge_name(sess.current_edge(i)) + " (budget "
+ std::to_string(budget.count())
+ " ms); that peer stopped sending or is not driving");
slice = std::min(slice, budget - waited);
}
if (all_done)
return;
if (slice.count() < 1)
slice = std::chrono::milliseconds(1);
bool ran = false;
bool one_domain = true;
for (auto * s : sinks)
one_domain = one_domain && s->wait_domain() == sinks.front()->wait_domain();
if (sinks.size() == 1 || one_domain)
ran = pump_sink(*sinks.front(), slice);
else
{
for (auto * s : sinks)
s->poll();
for (auto * s : sinks)
if (pump_sink(*s, std::chrono::milliseconds(1)))
ran = true;
}
if (!bounded)
{
if (ran || advanced)
spins = 0;
else if (++spins > 100000u)
throw std::runtime_error(what + ": peer not ready (drive both ends)");
}
}
}
} // namespace detail
/// @brief Lower exchange waves of `p` to `schedule_round`s for `schedule_session`.
/// @details Round `r`'s `produce` applies the previous round's peer bytes,
/// runs local groups since the prior round's wave, and packs this
/// round's slot. See `detail::party_steps` for the per-party roles.
/// Call a finish (`finish_schedule` or the mesh drive) afterwards so
/// the last peer slot and trailing compute-only waves land in `values`.
inline std::vector<schedule_round> plan_to_schedule(const plan & p,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party,
std::size_t lanes, const drive_options & opt = {})
{
if (values.size() < p.nodes().size())
values.resize(p.nodes().size());
for (std::uint32_t id = 0; id < p.nodes().size(); ++id)
detail::ensure_value_storage(p, values, id, lanes);
auto steps = std::make_shared<std::vector<detail::party_step>>(
detail::party_steps(p, party));
auto zeros = std::make_shared<detail::zero_share_cache>();
std::vector<schedule_round> rounds;
rounds.reserve(steps->size());
for (std::size_t ri = 0; ri < steps->size(); ++ri)
{
const auto & st = (*steps)[ri];
schedule_round step;
step.slot_bytes = st.part.slot_bytes;
step.channel = st.part.channel;
step.edge = st.edge;
step.recv = st.dealer_side ? receive_rule::copy_peer
: detail::part_receive_rule(p, st.part);
step.sink_round = st.sink_round;
step.timeout = round_budget(opt, ri);
const std::size_t local_from = ri == 0 ? 0 : (*steps)[ri - 1].part.wave + 1;
step.produce = [&p, &values, &kernels, party, lanes, opt, ri, local_from,
steps, zeros](std::size_t index, const std::uint8_t * peer,
std::size_t peer_n, std::uint8_t * out) {
const auto & me = (*steps)[ri];
if (ri > 0)
{
const auto & prev = (*steps)[ri - 1];
if (!prev.dealer_side)
detail::apply_peer_exchanges(p, prev.part.exchanges,
prev.part.slot_bytes, party, index, lanes, peer, peer_n,
values, opt);
}
for (std::size_t wi = local_from; wi <= me.part.wave; ++wi)
detail::run_wave_locals_lane(p, p.wave(wi), party, index, lanes,
values, kernels, opt);
if (out == nullptr)
return;
if (me.dealer_side)
detail::pack_dealer_part(p, me.part, me.dest, index, lanes, values,
*zeros, out);
else if (me.part.channel == edge_channel::dealer)
std::memset(out, 0, me.part.slot_bytes);
else
detail::pack_exchanges(p, me.part.exchanges, index, lanes, values,
opt, out);
};
rounds.push_back(std::move(step));
}
return rounds;
}
/// @brief Apply the last peer slot and trailing compute-only waves.
inline void finish_schedule(const plan & p, RoundSink & sink,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party = 0,
const drive_options & opt = {})
{
const auto steps = detail::party_steps(p, party);
for (const auto & st : steps)
if (st.edge != edge_peer)
throw std::invalid_argument("finish_schedule: plan uses the "
+ net::edge_name(st.edge) + " edge; drive it on an edge_mesh");
edge_mesh m{{&sink}};
detail::finish_steps(p, steps, m, values, kernels, party, sink.count(), opt);
}
/// @brief Drive a peer-only plan through `schedule_session` on one sink.
inline void drive_via_schedule(const plan & p, RoundSink & sink,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party = 0,
const drive_options & opt = {})
{
for (const auto & st : detail::party_steps(p, party))
if (st.edge != edge_peer)
throw std::invalid_argument("drive_via_schedule: plan has "
+ net::edge_name(st.edge) + " exchanges; bind that edge with "
"drive_via_schedule(plan, edge_mesh) or app::run_parties");
auto rounds = plan_to_schedule(p, values, kernels, party, sink.count(), opt);
if (rounds.empty())
{
finish_schedule(p, sink, values, kernels, party, opt);
return;
}
schedule_session sess(sink.count(), edge_mesh{{&sink}}, std::move(rounds),
/*rebase_single=*/true, opt.probe, opt.pipeline_credit);
for (std::size_t i = 0; i < sink.count(); ++i)
sess.submit(i);
detail::run_session(sess, opt, "drive_via_schedule");
finish_schedule(p, sink, values, kernels, party, opt);
}
/// @brief Drive a plan on an edge mesh (peer / rss_next / dealer / star).
inline void drive_via_schedule(const plan & p, edge_mesh mesh,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party = 0,
const drive_options & opt = {})
{
std::size_t lanes = 0;
for (std::size_t e = 0; e < mesh.size() && lanes == 0; ++e)
if (mesh.has(static_cast<edge_id>(e)))
lanes = mesh.at(static_cast<edge_id>(e)).count();
if (lanes == 0)
throw std::logic_error("drive_via_schedule(mesh): no edges bound");
const auto steps = detail::party_steps(p, party);
for (const auto & st : steps)
if (!mesh.has(st.edge))
throw std::invalid_argument("drive_via_schedule(mesh): party "
+ std::to_string(party) + " needs the " + net::edge_name(st.edge)
+ " edge, which the mesh does not bind");
auto rounds = plan_to_schedule(p, values, kernels, party, lanes, opt);
if (rounds.empty())
{
detail::finish_steps(p, steps, mesh, values, kernels, party, lanes, opt);
return;
}
schedule_session sess(lanes, std::move(mesh), std::move(rounds),
/*rebase_single=*/false, opt.probe, opt.pipeline_credit);
for (std::size_t i = 0; i < lanes; ++i)
sess.submit(i);
detail::run_session(sess, opt, "drive_via_schedule(mesh)");
detail::finish_steps(p, steps, sess.mesh(), values, kernels, party, lanes, opt);
}
/// @brief Pad / setup frames as schedule rounds, then the online plan rounds.
/// @details Pass `iknp::setup_rounds(...)` (or any setup count). Pads occupy
/// peer sink rounds `[0, setup)`; online peer rounds are rebased.
inline std::vector<schedule_round> plan_with_pad_setup(const plan & p,
std::size_t setup_rounds, std::size_t pad_slot_bytes,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party,
std::size_t lanes, const std::shared_ptr<std::vector<std::uint8_t>> & tape,
const drive_options & opt = {})
{
auto pads = make_pad_rounds(setup_rounds, pad_slot_bytes, tape,
edge_channel::peer);
auto online = plan_to_schedule(p, values, kernels, party, lanes, opt);
for (auto & r : online)
if (r.channel == edge_channel::peer)
r.sink_round = static_cast<std::uint16_t>(
static_cast<std::size_t>(r.sink_round) + setup_rounds);
return splice_rounds(std::move(pads), std::move(online));
}
/// @brief Drive a peer-only plan on a `stream_array` (index = exchange round).
/// @details Builds a `stream_array_sink` from `p.slot_bytes_all()` and runs
/// `drive_via_schedule`. Memory, file, and mux arrays all work; the
/// schedule loop flushes one round at a time. Plans that use
/// `rss_next` / `dealer` edges need `drive_via_schedule(plan,
/// edge_mesh{...})` instead.
/// @brief Per-channel exchange slot widths (peer / rss_next / dealer order).
struct channel_slot_bytes
{
std::vector<std::size_t> peer;
std::vector<std::size_t> rss_next;
std::vector<std::size_t> dealer;
};
/// @brief Per-channel round widths (a mixed wave contributes to each channel).
inline channel_slot_bytes slot_bytes_by_channel(const plan & p)
{
channel_slot_bytes out;
for (const auto & part : detail::schedule_parts(p))
{
switch (part.channel)
{
case edge_channel::peer:
out.peer.push_back(part.slot_bytes);
break;
case edge_channel::rss_next:
out.rss_next.push_back(part.slot_bytes);
break;
case edge_channel::dealer:
out.dealer.push_back(part.slot_bytes);
break;
}
}
return out;
}
/// @brief Build stream sinks from a plan's channel slot layout.
inline net::stream_edge_sinks make_stream_edge_sinks_for_plan(const plan & p,
net::stream_array * peer, net::stream_array * rss_next,
net::stream_array * dealer, std::size_t lanes = 1)
{
const auto slots = slot_bytes_by_channel(p);
return net::make_stream_edge_sinks(peer, slots.peer, rss_next, slots.rss_next,
dealer, slots.dealer, lanes);
}
/// @brief Drive on a pre-built multi-edge stream mesh (`edge_sinks` adapter).
inline void drive_plan_on_stream_mesh(const plan & p,
net::stream_edge_sinks & sinks,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party = 0,
const drive_options & opt = {})
{
drive_via_schedule(p, sinks.mesh(), values, kernels, party, opt);
}
inline void drive_plan_on_streams(const plan & p, net::stream_array & streams,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party = 0,
std::size_t lanes = 1, const drive_options & opt = {})
{
auto slots = p.slot_bytes_all();
if (!slots.empty() && streams.size() == 0)
throw std::invalid_argument("drive_plan_on_streams: empty stream_array");
net::stream_array_sink sink(streams, std::move(slots), lanes, opt.framing);
drive_via_schedule(p, sink, values, kernels, party, opt);
}
/// @brief Two parties, two threads, one memory stream pair; drive both plans.
/// @details `p0` / `p1` must share the same `slot_bytes_all()` shape (typical
/// when both composers record the same protocol).
inline void drive_both_on_streams(const plan & p0, const plan & p1,
std::vector<std::vector<std::uint8_t>> & v0,
std::vector<std::vector<std::uint8_t>> & v1,
const std::map<std::uint32_t, kernel_fn> & kernels = {},
std::size_t lanes = 1, const drive_options & opt = {})
{
const auto slots = p0.slot_bytes_all();
if (slots != p1.slot_bytes_all())
throw std::invalid_argument(
"drive_both_on_streams: party slot shapes differ");
const std::size_t nstreams =
slots.empty() ? 1 : lanes_for_plan(slots.size(), opt);
auto ends = net::make_memory_stream_pair(nstreams);
std::mutex err_mu;
std::exception_ptr err;
auto note = [&](std::exception_ptr e) {
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::move(e);
};
std::thread t0([&] {
try
{
drive_plan_on_streams(p0, ends.first, v0, kernels, 0, lanes, opt);
}
catch (...)
{
note(std::current_exception());
}
});
std::thread t1([&] {
try
{
drive_plan_on_streams(p1, ends.second, v1, kernels, 1, lanes, opt);
}
catch (...)
{
note(std::current_exception());
}
});
t0.join();
t1.join();
if (err)
std::rethrow_exception(err);
}
/// @brief Same plan on both parties (common for symmetric compose graphs).
inline void drive_both_on_streams(const plan & p,
std::vector<std::vector<std::uint8_t>> & v0,
std::vector<std::vector<std::uint8_t>> & v1,
const std::map<std::uint32_t, kernel_fn> & kernels = {},
std::size_t lanes = 1, const drive_options & opt = {})
{
drive_both_on_streams(p, p, v0, v1, kernels, lanes, opt);
}
} // namespace protocol
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_COMPOSE_HPP__