libdpf/include/dpf/compose.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

4293 lines
162 KiB
C++
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @file dpf/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__