/// @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 #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #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( static_cast(a) | static_cast(b)); } HEDLEY_CONST HEDLEY_NO_THROW inline constexpr bool has_flag(op_flags f, op_flags bit) noexcept { return (static_cast(f) & static_cast(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 & nodes, const std::vector & inputs, block_span output, std::size_t sink_lanes)>; struct simd_group { std::uint32_t opcode = 0; effect kind = effect::compute; std::vector nodes; }; struct wave_info { std::size_t index = 0; std::vector groups; std::vector 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 slot_bytes_all() const { std::vector 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 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(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 & 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 & 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> &, const std::map &, std::size_t, const struct drive_options &); std::vector waves_; std::map effect_counts_; std::map, std::size_t> conversion_counts_; std::vector all_nodes_; std::vector node_wave_; std::vector value_bytes_; std::vector domains_; std::vector effects_; std::vector opcodes_; std::vector levels_; std::vector aux_; std::vector> inputs_; std::vector 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 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(k.kind); h = h * 131u + static_cast(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 struct aby_holder final : aby_holder_base { beavers::session 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 tmp(in, in + in_bytes); unsigned borrow = 1; for (std::size_t i = 0; i < tmp.size(); ++i) { const unsigned x = static_cast(static_cast(~tmp[i])) + borrow; tmp[i] = static_cast(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(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 inputs, domain dom, std::size_t value_bytes, op_flags flags = op_flags::none) { std::vector 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 segments; }; HEDLEY_WARN_UNUSED_RESULT fuse_result exchange_fuse(const std::vector & 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(off)); fr.segments.push_back(seg); off += width; } return fr; } template std::vector fan(std::size_t n, Fn && fn) { std::vector 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(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(early_stop); return leaf; } /// @brief Poplar / idpf_agg: emit a prefix share after each CW — no extra rounds. struct prefix_walk_result { std::vector 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(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 & 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 & 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 & 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(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(level), static_cast(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(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 leaves; node answers_open{}; }; template 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 HEDLEY_WARN_UNUSED_RESULT multipoint_result schedule_cuckoo_probes( const std::vector & probes, SeedAt && seed_at, std::size_t depth, std::size_t slot_bytes, std::size_t answer_bytes) { std::vector 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 & payloads, std::size_t per_payload_bytes = 0) { if (payloads.empty()) throw std::invalid_argument("exchange_pack needs payloads"); std::vector 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(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(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 beavers::session & aby() { return aby_lane(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 beavers::session & 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>(); 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 &>(*it->second).session; } template node opened(typename beavers::session::wire w, std::size_t lane = 0) { auto & s = aby_lane(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(w.id()); info_[out.id].aby_session = static_cast(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) : 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 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(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(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(w.id()); info_[out.id].aby_session = static_cast(sess_index); return out; } template 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(lane); auto wx = ensure_aby_wire(x, lane); auto wy = ensure_aby_wire(y, lane); auto wz = s.product(wx, wy); return opened(wz, lane); } template node aby_product(node x, node y, std::size_t lane = 0) { return beaver_triple(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(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 seen(n, 0); std::function 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::vector> 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 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(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 & 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(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 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(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 & 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(barrier_index); key.level = static_cast(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(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(sess_index); const auto id = static_cast(info_.size()); info_.push_back(std::move(ni)); intern_.emplace(std::move(key), id); return node{id}; } template typename beavers::session::wire ensure_aby_wire(node n, std::size_t lane = 0) { auto & s = aby_lane(lane); if (info_[n.id].aby_wire >= 0 && info_[n.id].aby_session >= 0 && static_cast(info_[n.id].aby_session) < aby_order_.size() && aby_order_[static_cast(info_[n.id].aby_session)].lane == lane && aby_order_[static_cast(info_[n.id].aby_session)].type == std::type_index(typeid(Ring))) return s.wire_at(static_cast(info_[n.id].aby_wire)); auto w = s.input(); info_[n.id].aby_wire = static_cast(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(i); aby_import_[{static_cast(i), w.id()}] = n.id; break; } } return w; } std::size_t party_ = 0; std::vector info_; std::unordered_map intern_; std::map> aby_; std::vector aby_order_; /// ABY input wire id → compose node that produced the share (per session). std::map, 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> * 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_timeout; /// Per-round override (schedule round index → budget); wins over edges. std::map 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(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 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::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 & nodes, const std::vector & 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 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(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(mine[i] ^ peer[i]); } inline void apply_fuse_segment(const plan & p, std::vector> & 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(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 &, const std::vector & 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(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(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( 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> & 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> & 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(*opt.rss_seeds, idx); std::memcpy(out, &z, 8); } else if (vb == 1) { const auto z = rss::zero_share(*opt.rss_seeds, idx); out[0] = z; } else { for (std::size_t i = 0; i < vb; ++i) { const auto z = rss::zero_share( *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(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(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(p.aux_of(id)); const auto o = trunc::trunc_exact_party(x, r, delta, wrap, s, static_cast(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(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(0); return values[src][lane * sb]; }; if (op == opcodes::share_input && ins.size() >= 2 && vb >= 8) { const auto d = static_cast(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( 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(party)); return; } if (op == opcodes::bin_b2a && ins.size() >= 2 && vb >= 8) { const unsigned ell = static_cast(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(i + 1) * 8 <= ab) std::memcpy(&r, arith + static_cast(i) * 8, 8); const std::uint8_t mask = (static_cast(i / 8) < mb) ? static_cast((maskb[i / 8] >> (i % 8)) & 1u) : 0; ot::dabit dab{0, r}; acc += edabit::b2a_party_bit(0, dab, mask, static_cast(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(*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(p.aux_of(id)); const auto z = trunc::mul_trunc_party(0, 0, load64(ins[0]), load64(ins[1]), load64(ins[2]), load64(ins[3]), load64(ins[4]), s, static_cast(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(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> & 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 & 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> & 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(p.level_of(ex.id)); const auto bi = static_cast(op - opcodes::beaver_delta); opt.beavers->apply(sess, bi, peer + off, nb); std::vector 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(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 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 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> & 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 & exchanges, std::size_t lane, std::size_t lanes, std::vector> & 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(p.level_of(ex.id)); const auto bi = static_cast(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> & 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> & 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> & values, const std::map & 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 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 exchange_wave_indices(const plan & p) { std::vector 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> & values, const std::map & 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 packed_in(packed_lanes * vb); std::vector 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 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( opt.compact_sink ? (exchange_i - opt.from_exchange_wave) : exchange_i); ++exchange_i; for (std::size_t lane = 0; lane < lanes; ++lane) { std::vector 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(p.level_of(ex.id)); const auto bi = static_cast(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 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(p.level_of(ex.id)); const auto bi = static_cast(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 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 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 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 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 schedule_parts(const plan & p) { std::vector 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(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_steps(const plan & p, std::size_t party) { const auto parts = schedule_parts(p); std::vector 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 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::vector>; /// @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> & 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 a(nb), b(nb); for (auto & byte : a) byte = dpf::uniform_sample(); 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(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 & steps, edge_mesh & mesh, std::vector> & values, const std::map & 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 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 mark(count, ~std::uint64_t{0}); std::vector since(count, clock::now()); std::vector 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(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 plan_to_schedule(const plan & p, std::vector> & values, const std::map & 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>( detail::party_steps(p, party)); auto zeros = std::make_shared(); std::vector 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> & values, const std::map & 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> & values, const std::map & 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> & values, const std::map & 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(e))) lanes = mesh.at(static_cast(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 plan_with_pad_setup(const plan & p, std::size_t setup_rounds, std::size_t pad_slot_bytes, std::vector> & values, const std::map & kernels, std::size_t party, std::size_t lanes, const std::shared_ptr> & 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( static_cast(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 peer; std::vector rss_next; std::vector 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> & values, const std::map & 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> & values, const std::map & 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> & v0, std::vector> & v1, const std::map & 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 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> & v0, std::vector> & v1, const std::map & 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__