libdpf/party/flow_util.hpp

811 lines
31 KiB
C++
Raw Normal View History

/// @file party/flow_util.hpp
/// @brief Shared helpers for (2+1) party flow bodies.
#ifndef LIBDPF_PARTY_FLOW_UTIL_HPP__
#define LIBDPF_PARTY_FLOW_UTIL_HPP__
#include <cstdint>
#include <stdexcept>
#include <utility>
#include <vector>
#include "dpf/beaver.hpp"
#include "dpf/compose.hpp"
#include "dpf/net/mux_sink.hpp" // trio::batch default
#include "dpf/net/party_tape_io.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/net/trio.hpp"
#include "dpf/secret_share.hpp"
#include "dpf/verifiable.hpp"
#include <algorithm>
#include <cstring>
#include <map>
#include <memory>
#include <type_traits>
#include <vector>
namespace dpf
{
namespace party
{
namespace util
{
using net::role;
using net::trio;
using u64 = std::uint64_t;
struct Counter
{
int draws = 0;
u64 operator()()
{
++draws;
return 0x9e3779b97f4a7c15ull * static_cast<u64>(draws);
}
};
inline void require(bool cond, const char * msg)
{
if (!cond)
throw std::runtime_error(msg);
}
inline std::pair<u64, u64> split_u64(u64 secret, unsigned tag = 1)
{
u64 p0 = 0x9e3779b97f4a7c15ull * (tag + 1u);
return {p0, secret - p0};
}
template <typename Ring>
std::pair<Ring, Ring> split_ring(Ring secret, unsigned tag = 1)
{
using traits = beavers::ring_traits<Ring>;
Ring p0 = Ring{static_cast<u64>(0x9e3779b97f4a7c15ull * (tag + 1u))};
return {p0, traits::sub(secret, p0)};
}
template <typename Ring>
Ring open_additive(trio & net, role self, Ring mine)
{
role peer = self == role::p0 ? role::p1 : role::p0;
return net.open_with(peer, mine);
}
/// @brief Open many independent additive shares in one vector exchange.
template <typename Ring>
std::vector<Ring> open_additive(trio & net, role self, const std::vector<Ring> & mine)
{
role peer = self == role::p0 ? role::p1 : role::p0;
return net.open_vec_with(peer, mine);
}
template <typename T>
T open_subtractive(trio & net, role self, T mine)
{
role peer = self == role::p0 ? role::p1 : role::p0;
T theirs = net.exchange_with(peer, mine);
return self == role::p0 ? static_cast<T>(mine - theirs)
: static_cast<T>(theirs - mine);
}
/// @brief Open many independent subtractive shares in one vector exchange.
template <typename T>
std::vector<T> open_subtractive(trio & net, role self, const std::vector<T> & mine)
{
role peer = self == role::p0 ? role::p1 : role::p0;
auto theirs = net.exchange_vec_with(peer, mine);
require(theirs.size() == mine.size(), "open_subtractive vec size");
std::vector<T> out(mine.size());
for (std::size_t i = 0; i < mine.size(); ++i)
{
out[i] = self == role::p0 ? static_cast<T>(mine[i] - theirs[i])
: static_cast<T>(theirs[i] - mine[i]);
}
return out;
}
template <typename T>
constexpr T share_bits(T s) noexcept
{
return s;
}
template <typename T, std::size_t Party, sharing Scheme>
constexpr T share_bits(const secret_share<T, Party, Scheme> & s) noexcept
{
return s.raw();
}
inline void install_and_bind_u64(beavers::session<u64> & s, trio & net, role self,
const std::vector<beavers::session<u64>::wire> & inputs,
const std::vector<u64> & secrets)
{
auto tape = net::accept_session<u64>(net);
s.install_party(self == role::p0 ? 0u : 1u, tape);
for (std::size_t i = 0; i < inputs.size(); ++i)
{
auto [p0, p1] = split_u64(secrets[i], static_cast<unsigned>(i + 1));
s.bind_party(inputs[i], self == role::p0 ? p0 : p1);
}
}
template <typename Ring>
void install_and_bind(beavers::session<Ring> & s, trio & net, role self,
const std::vector<typename beavers::session<Ring>::wire> & inputs,
const std::vector<Ring> & secrets)
{
auto tape = net::accept_session<Ring>(net);
s.install_party(self == role::p0 ? 0u : 1u, tape);
for (std::size_t i = 0; i < inputs.size(); ++i)
{
auto [p0, p1] = split_ring<Ring>(secrets[i], static_cast<unsigned>(i + 1));
s.bind_party(inputs[i], self == role::p0 ? p0 : p1);
}
}
/// @brief Host `evaluate_party_batch` on a RoundSink (one instance).
/// @details Each vector exchange is one schedule round. The sink must have been
/// sized with enough rounds and `slot_bytes >= 4 + n * sizeof(Ring)`
/// for the largest batch this session will open. `index` selects the
/// instance lane when the sink holds many circuits; all parties must
/// drive the same `index` sequence. Prefer `count == 1` sinks.
template <typename Ring>
void evaluate_online_on_sink(beavers::session<Ring> & s, net::RoundSink & sink,
std::size_t index = 0)
{
std::uint16_t flush_round = 0;
s.evaluate_party_batch([&](std::vector<Ring> mine) {
if (flush_round >= sink.rounds())
throw std::runtime_error("evaluate_online_on_sink: out of rounds");
const std::size_t slot = sink.slot_bytes(flush_round);
const std::size_t need = sizeof(std::uint32_t) + mine.size() * sizeof(Ring);
if (need > slot)
throw std::runtime_error("evaluate_online_on_sink: slot too small");
std::vector<std::uint8_t> buf(slot, 0);
const std::uint32_t n = static_cast<std::uint32_t>(mine.size());
std::memcpy(buf.data(), &n, 4);
if (!mine.empty())
std::memcpy(buf.data() + 4, mine.data(), mine.size() * sizeof(Ring));
sink.submit(flush_round, index, buf.data(), buf.size());
sink.flush_round(flush_round);
sink.poll();
if (!sink.peer_ready(flush_round, index))
throw std::runtime_error("evaluate_online_on_sink: peer missing");
std::vector<std::uint8_t> peer_buf(slot);
sink.read_peer(flush_round, index, peer_buf.data(), slot);
std::uint32_t pn = 0;
std::memcpy(&pn, peer_buf.data(), 4);
if (pn != n)
throw std::runtime_error("evaluate_online_on_sink: peer size");
std::vector<Ring> peer(n);
if (n != 0)
std::memcpy(peer.data(), peer_buf.data() + 4, n * sizeof(Ring));
++flush_round;
return peer;
});
}
/// @brief Sink round budget for `evaluate_party_batch` / auth-batch.
/// @details One input flush and one output flush per ready-round, plus a small
/// slack for same-round dependency flushes. Must not grow with
/// `wire_count`: a 2k-wire depth-1 circuit still needs only a handful
/// of vector exchanges, and allocating O(wires) round windows is a
/// multi-megabyte tax on every online evaluate.
inline int beaver_batch_round_budget(int max_ready_round,
std::size_t /*wire_count*/) noexcept
{
return std::max(1, max_ready_round * 2 + 16);
}
/// @brief Online beaver opens hosted on a count-1 batch sink from `trio::batch`.
template <typename Ring>
void evaluate_online(beavers::session<Ring> & s, trio & net, role self)
{
(void)self;
const int rounds =
beaver_batch_round_budget(s.max_ready_round(), s.wire_count());
// Upper bound: every wire δ in one vector exchange, plus length prefix.
const std::size_t slot =
sizeof(std::uint32_t) + s.wire_count() * sizeof(Ring);
std::vector<std::size_t> slots(static_cast<std::size_t>(rounds), slot);
auto sink = net.batch(1, std::move(slots));
evaluate_online_on_sink(s, *sink, 0);
}
/// @brief Drive many independent circuits on one RoundSink (one index each).
/// @details Each circuit runs to completion on its lane. Callers that need a
/// single prefix flush across instances must submit lockstep themselves
/// via `schedule_session`; this helper is the portable multi-lane host
/// of `evaluate_party_batch` on an existing sink.
template <typename Ring>
void evaluate_online_multi_on_sink(
std::vector<beavers::session<Ring> *> sessions, net::RoundSink & sink)
{
if (sessions.size() != sink.count())
throw std::invalid_argument("evaluate_online_multi_on_sink: count");
for (std::size_t i = 0; i < sessions.size(); ++i)
{
if (!sessions[i])
throw std::invalid_argument("evaluate_online_multi_on_sink: null");
evaluate_online_on_sink(*sessions[i], sink, i);
}
}
/// @brief Authenticated `evaluate_party_auth_batch` on a RoundSink.
template <typename Ring>
void evaluate_online_auth_on_sink(beavers::session<Ring> & s,
net::RoundSink & sink, std::size_t index = 0)
{
using opening = beavers::auth_opening<Ring>;
std::uint16_t flush_round = 0;
s.evaluate_party_auth_batch([&](std::vector<opening> mine) {
if (flush_round >= sink.rounds())
throw std::runtime_error("evaluate_online_auth_on_sink: out of rounds");
const std::size_t slot = sink.slot_bytes(flush_round);
const std::size_t need =
sizeof(std::uint32_t) + mine.size() * sizeof(opening);
if (need > slot)
throw std::runtime_error("evaluate_online_auth_on_sink: slot too small");
std::vector<std::uint8_t> buf(slot, 0);
const std::uint32_t n = static_cast<std::uint32_t>(mine.size());
std::memcpy(buf.data(), &n, 4);
if (!mine.empty())
std::memcpy(buf.data() + 4, mine.data(), mine.size() * sizeof(opening));
sink.submit(flush_round, index, buf.data(), buf.size());
sink.flush_round(flush_round);
sink.poll();
if (!sink.peer_ready(flush_round, index))
throw std::runtime_error("evaluate_online_auth_on_sink: peer missing");
std::vector<std::uint8_t> peer_buf(slot);
sink.read_peer(flush_round, index, peer_buf.data(), slot);
std::uint32_t pn = 0;
std::memcpy(&pn, peer_buf.data(), 4);
if (pn != n)
throw std::runtime_error("evaluate_online_auth_on_sink: peer size");
std::vector<opening> peer(n);
if (n != 0)
std::memcpy(peer.data(), peer_buf.data() + 4, n * sizeof(opening));
++flush_round;
return peer;
});
}
/// @brief Authenticated opening path when the session tape carries IT-MACs.
/// @details Call only after the dealer ran `set_mac_key` and `deal_session`.
/// Uses `trio::batch` so a installed `comm_hook` can replace the sink.
template <typename Ring>
void evaluate_online_auth(beavers::session<Ring> & s, trio & net, role self)
{
(void)self;
using opening = beavers::auth_opening<Ring>;
const int rounds =
beaver_batch_round_budget(s.max_ready_round(), s.wire_count());
const std::size_t slot =
sizeof(std::uint32_t) + s.wire_count() * sizeof(opening);
std::vector<std::size_t> slots(static_cast<std::size_t>(rounds), slot);
auto sink = net.batch(1, std::move(slots));
evaluate_online_auth_on_sink(s, *sink, 0);
}
/// @brief u64 beaver δ host using `party_batch_stepper` (one per ABY session).
struct u64_beaver_host final : protocol::beaver_host
{
explicit u64_beaver_host(
std::vector<beavers::session<std::uint64_t> *> sessions)
{
steppers_.reserve(sessions.size());
for (auto * s : sessions)
{
if (s == nullptr)
steppers_.push_back(nullptr);
else
steppers_.push_back(
std::make_unique<beavers::party_batch_stepper<std::uint64_t>>(
*s));
}
}
void pack(std::size_t session_index, std::size_t /*barrier_index*/,
std::uint8_t * dst, std::size_t n) override
{
if (session_index >= steppers_.size() || !steppers_[session_index])
throw std::runtime_error("u64_beaver_host: missing session");
auto mine = steppers_[session_index]->take_local();
const std::uint32_t count = static_cast<std::uint32_t>(mine.size());
const std::size_t need =
sizeof(std::uint32_t) + mine.size() * sizeof(std::uint64_t);
if (n < need)
throw std::runtime_error("u64_beaver_host: slot too small");
std::memset(dst, 0, n);
std::memcpy(dst, &count, sizeof(count));
if (!mine.empty())
std::memcpy(dst + sizeof(count), mine.data(),
mine.size() * sizeof(std::uint64_t));
}
void apply(std::size_t session_index, std::size_t /*barrier_index*/,
const std::uint8_t * src, std::size_t n) override
{
if (session_index >= steppers_.size() || !steppers_[session_index])
throw std::runtime_error("u64_beaver_host: missing session");
if (n < sizeof(std::uint32_t))
throw std::runtime_error("u64_beaver_host: peer short");
std::uint32_t count = 0;
std::memcpy(&count, src, sizeof(count));
const std::size_t need =
sizeof(std::uint32_t) + static_cast<std::size_t>(count)
* sizeof(std::uint64_t);
if (n < need)
throw std::runtime_error("u64_beaver_host: peer size");
std::vector<std::uint64_t> peer(count);
if (count != 0)
std::memcpy(peer.data(), src + sizeof(count),
peer.size() * sizeof(std::uint64_t));
steppers_[session_index]->apply_peer(peer);
}
private:
std::vector<std::unique_ptr<beavers::party_batch_stepper<std::uint64_t>>>
steppers_;
};
/// @brief IT-MAC δ host: `uint32 n ‖ auth_opening<Ring>[n]`.
/// @details `apply` checks each tag via `party_auth_batch_stepper`.
template <typename Ring>
struct auth_beaver_host final : protocol::beaver_host
{
using opening = beavers::auth_opening<Ring>;
explicit auth_beaver_host(std::vector<beavers::session<Ring> *> sessions)
{
static_assert(std::is_trivially_copyable_v<opening>,
"auth_opening must be trivially copyable");
steppers_.reserve(sessions.size());
for (auto * s : sessions)
{
if (s == nullptr)
steppers_.push_back(nullptr);
else
steppers_.push_back(
std::make_unique<beavers::party_auth_batch_stepper<Ring>>(*s));
}
}
void pack(std::size_t session_index, std::size_t, std::uint8_t * dst,
std::size_t n) override
{
if (session_index >= steppers_.size() || !steppers_[session_index])
throw std::runtime_error("auth_beaver_host: missing session");
auto mine = steppers_[session_index]->take_local();
const std::uint32_t count = static_cast<std::uint32_t>(mine.size());
const std::size_t need = sizeof(std::uint32_t) + mine.size() * sizeof(opening);
if (n < need)
throw std::runtime_error("auth_beaver_host: slot too small");
std::memset(dst, 0, n);
std::memcpy(dst, &count, sizeof(count));
if (!mine.empty())
std::memcpy(dst + sizeof(count), mine.data(),
mine.size() * sizeof(opening));
}
void apply(std::size_t session_index, std::size_t, const std::uint8_t * src,
std::size_t n) override
{
if (session_index >= steppers_.size() || !steppers_[session_index])
throw std::runtime_error("auth_beaver_host: missing session");
if (n < sizeof(std::uint32_t))
throw std::runtime_error("auth_beaver_host: peer short");
std::uint32_t count = 0;
std::memcpy(&count, src, sizeof(count));
const std::size_t need =
sizeof(std::uint32_t) + static_cast<std::size_t>(count) * sizeof(opening);
if (n < need)
throw std::runtime_error("auth_beaver_host: peer size");
std::vector<opening> peer(count);
if (count != 0)
std::memcpy(peer.data(), src + sizeof(count),
peer.size() * sizeof(opening));
steppers_[session_index]->apply_peer(peer);
}
private:
std::vector<std::unique_ptr<beavers::party_auth_batch_stepper<Ring>>> steppers_;
};
using u64_auth_beaver_host = auth_beaver_host<std::uint64_t>;
/// @brief Schedule a composer and drive it on a trio batch sink.
/// @details Uses `default_plan()`: RoundSink rounds == `plan::rounds()` ==
/// exchange-bearing waves. Pass `opt.from_exchange_wave` /
/// `compact_sink` for adaptive tails; `opt.beavers` for live δ.
inline void drive_composed(protocol::composer & c, trio & net,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, protocol::kernel_fn> & kernels,
std::size_t count = 1, protocol::drive_options opt = {})
{
auto p = c.default_plan();
auto slots = opt.compact_sink ? p.slot_bytes_from(opt.from_exchange_wave)
: p.slot_bytes_all();
if (slots.empty())
slots.push_back(0);
auto sink = net.batch(count, std::move(slots));
protocol::drive(p, *sink, values, kernels, c.party(), opt);
}
/// @brief Drive a compose plan on the trio mesh (2PC peer + 3PC RSS neighbor).
/// @details Shared semantics with `protocol::drive`: throws on missing kernels,
/// domain-correct opens, optional beaver host, incremental
/// `from_exchange_wave`. Per-exchange peer maps: `y` → neighbor ring,
/// else p0↔p1 (p2 idle).
inline void drive_composed_trio(protocol::composer & c, trio & net, role self,
std::vector<std::vector<std::uint8_t>> & values,
const std::map<std::uint32_t, protocol::kernel_fn> & kernels,
protocol::drive_options opt = {})
{
using protocol::domain;
using protocol::effect;
namespace opcodes = protocol::opcodes;
auto p = c.default_plan();
if (values.size() < p.nodes().size())
values.resize(p.nodes().size());
auto ensure = [&](std::uint32_t id) {
const std::size_t need = p.value_bytes_of(id);
if (values[id].size() < need)
values[id].assign(need, 0);
};
for (std::uint32_t id = 0; id < p.nodes().size(); ++id)
ensure(id);
const role peer = self == role::p0 ? role::p1 : role::p0;
const role rss_next =
static_cast<role>((static_cast<unsigned>(self) + 1u) % 3u);
auto run_groups = [&](const protocol::wave_info & wave) {
for (const auto & g : wave.groups)
{
if (g.kind == effect::convert)
{
for (auto n : g.nodes)
{
ensure(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;
}
if (ins.size() == 1)
{
protocol::detail::run_builtin_convert(g.opcode,
c.party(), values[ins[0]].data(),
p.value_bytes_of(ins[0]), values[n.id].data(),
p.value_bytes_of(n.id));
}
else if (ins.size() == 2
&& g.opcode == opcodes::conv_y2rss)
{
const auto half = p.value_bytes_of(ins[0]);
std::memcpy(values[n.id].data(),
values[ins[0]].data(), half);
std::memcpy(values[n.id].data() + half,
values[ins[1]].data(), half);
}
}
continue;
}
if (protocol::detail::is_beaver_opcode(g.opcode))
continue;
if (g.opcode == opcodes::fss_fuse_segment)
{
for (auto n : g.nodes)
{
ensure(n.id);
const auto & ins = p.inputs_of(n.id);
if (!ins.empty())
ensure(ins[0]);
protocol::detail::apply_fuse_segment(p, values, n.id, 1);
}
continue;
}
auto kit = kernels.find(g.opcode);
protocol::kernel_fn builtin;
const protocol::kernel_fn * fn = nullptr;
if (kit != kernels.end())
fn = &kit->second;
else if (protocol::detail::has_builtin_walk_kernel(g.opcode))
{
builtin = protocol::detail::builtin_walk_kernel(g.opcode);
fn = &builtin;
}
else if (g.kind == effect::blind)
{
for (auto n : g.nodes)
{
ensure(n.id);
std::fill(values[n.id].begin(), values[n.id].end(), 0);
}
continue;
}
else
{
throw std::runtime_error(
"drive_composed_trio missing kernel for opcode "
+ std::to_string(g.opcode));
}
for (auto n : g.nodes)
{
ensure(n.id);
std::vector<protocol::block_span> ins;
for (auto in_id : p.inputs_of(n.id))
{
ensure(in_id);
ins.push_back(protocol::block_span{values[in_id].data(), 1,
p.value_bytes_of(in_id)});
}
protocol::block_span out{values[n.id].data(), 1,
p.value_bytes_of(n.id)};
(*fn)(g.opcode, {n}, ins, out, 1);
}
}
};
std::size_t exchange_i = 0;
for (std::size_t w = 0; w < p.waves(); ++w)
{
const auto & wave = p.wave(w);
const bool skip = !wave.exchanges.empty()
&& exchange_i < opt.from_exchange_wave;
if (!skip)
run_groups(wave);
if (wave.exchanges.empty() || wave.slot_bytes == 0)
continue;
if (exchange_i < opt.from_exchange_wave)
{
++exchange_i;
continue;
}
++exchange_i;
// Split peer maps: y → RSS neighbor; dealer_pad → p2; else → 2PC peer.
std::vector<std::uint8_t> two_pc;
std::vector<std::uint8_t> rss;
std::vector<std::uint8_t> dealer;
struct zero_mask { std::size_t n = 0; domain dom = domain::a; };
std::vector<zero_mask> zeros;
std::vector<std::uint8_t> kinds; // 0 = 2pc, 1 = rss, 2 = dealer, 3 = zero
kinds.reserve(wave.exchanges.size());
for (auto ex : wave.exchanges)
{
ensure(ex.id);
const auto nb = p.value_bytes_of(ex.id);
const auto op = p.opcode_of(ex.id);
const bool neighbor = protocol::detail::is_rss_neighbor_open(p, ex.id);
const bool from_dealer = op == opcodes::dealer_pad;
const bool zero = op == opcodes::dealer_zero;
kinds.push_back(zero ? 3 : (from_dealer ? 2 : (neighbor ? 1 : 0)));
if (zero)
zeros.push_back(zero_mask{nb, p.domain_of(ex.id)});
std::vector<std::uint8_t> chunk(nb, 0);
if (protocol::detail::is_beaver_opcode(op) && opt.beavers != nullptr)
{
const auto sess = static_cast<std::size_t>(p.level_of(ex.id));
const auto bi =
static_cast<std::size_t>(op - opcodes::beaver_delta);
opt.beavers->pack(sess, bi, chunk.data(), nb);
std::memcpy(values[ex.id].data(), chunk.data(), nb);
}
else
{
const auto & ins = p.inputs_of(ex.id);
if (!ins.empty() && !protocol::detail::is_beaver_opcode(op))
{
const auto src = ins[0];
ensure(src);
const auto sb = p.value_bytes_of(src);
std::memcpy(chunk.data(), values[src].data(),
std::min(sb, nb));
std::memcpy(values[ex.id].data(), values[src].data(),
std::min(sb, nb));
}
else
std::memcpy(chunk.data(), values[ex.id].data(), nb);
}
if (from_dealer)
dealer.insert(dealer.end(), chunk.begin(), chunk.end());
else if (zero)
continue;
else if (neighbor)
rss.insert(rss.end(), chunk.begin(), chunk.end());
else
two_pc.insert(two_pc.end(), chunk.begin(), chunk.end());
}
std::vector<std::uint8_t> two_pc_peer;
std::vector<std::uint8_t> rss_peer;
std::vector<std::uint8_t> dealer_peer;
std::vector<std::vector<std::uint8_t>> zero_recv(zeros.size());
if (!rss.empty())
rss_peer = net.exchange_vec_with(rss_next, rss);
if (!two_pc.empty())
{
if (self == role::p2)
{
// Dealer idle on 2PC opens — leave peer zeros; reconstruct
// will not run for p2 on those slots.
two_pc_peer.assign(two_pc.size(), 0);
}
else
two_pc_peer = net.exchange_vec_with(peer, two_pc);
}
if (!dealer.empty())
{
if (self == role::p2)
{
net.send_bytes_to(role::p0, net::msg::bytes, dealer.data(),
dealer.size());
net.send_bytes_to(role::p1, net::msg::bytes, dealer.data(),
dealer.size());
dealer_peer = dealer;
}
else
dealer_peer = net.recv_bytes_from(role::p2, net::msg::bytes);
}
if (!zeros.empty())
{
if (self == role::p2)
{
for (std::size_t zi = 0; zi < zeros.size(); ++zi)
{
std::vector<std::uint8_t> a(zeros[zi].n), b(zeros[zi].n);
for (auto & byte : a)
byte = dpf::uniform_sample<std::uint8_t>();
if (zeros[zi].dom == domain::b)
b = a;
else if (zeros[zi].n % 8 == 0)
{
for (std::size_t i = 0; i < a.size(); i += 8)
{
std::uint64_t w = 0;
std::memcpy(&w, a.data() + i, 8);
w = static_cast<std::uint64_t>(0) - w;
std::memcpy(b.data() + i, &w, 8);
}
}
else
{
for (std::size_t i = 0; i < a.size(); ++i)
b[i] = a[i];
}
net.send_bytes_to(role::p0, net::msg::bytes, a.data(), a.size());
net.send_bytes_to(role::p1, net::msg::bytes, b.data(), b.size());
}
}
else
{
for (std::size_t zi = 0; zi < zeros.size(); ++zi)
zero_recv[zi] = net.recv_bytes_from(role::p2, net::msg::bytes);
}
}
std::size_t off2 = 0, offr = 0, offd = 0, zi = 0;
for (std::size_t i = 0; i < wave.exchanges.size(); ++i)
{
auto ex = wave.exchanges[i];
const auto nb = p.value_bytes_of(ex.id);
ensure(ex.id);
const auto op = p.opcode_of(ex.id);
const bool neighbor = kinds[i] == 1;
const bool from_dealer = kinds[i] == 2;
if (kinds[i] == 3)
{
if (self != role::p2)
{
if (zi >= zero_recv.size() || zero_recv[zi].size() < nb)
throw std::runtime_error("drive_composed_trio: zero mask");
std::memcpy(values[ex.id].data(), zero_recv[zi].data(), nb);
}
++zi;
continue;
}
if (from_dealer)
{
if (dealer_peer.size() < offd + nb)
throw std::runtime_error("drive_composed_trio: dealer size");
std::memcpy(values[ex.id].data(), dealer_peer.data() + offd, nb);
offd += nb;
continue;
}
const std::uint8_t * peer_bytes =
neighbor ? rss_peer.data() + offr : two_pc_peer.data() + off2;
if (neighbor)
offr += nb;
else
off2 += nb;
if (self == role::p2 && !neighbor)
continue;
if (protocol::detail::is_beaver_opcode(op) && opt.beavers != nullptr)
{
const auto sess = static_cast<std::size_t>(p.level_of(ex.id));
const auto bi =
static_cast<std::size_t>(op - opcodes::beaver_delta);
opt.beavers->apply(sess, bi, peer_bytes, nb);
}
domain d = p.domain_of(ex.id);
const auto & ins = p.inputs_of(ex.id);
if (!ins.empty())
d = p.domain_of(ins[0]);
if (neighbor)
{
std::memcpy(values[ex.id].data(), peer_bytes, nb);
}
else
{
std::vector<std::uint8_t> mine(nb, 0);
if (!ins.empty() && !protocol::detail::is_beaver_opcode(op))
{
const auto sb = p.value_bytes_of(ins[0]);
std::memcpy(mine.data(), values[ins[0]].data(),
std::min(sb, nb));
}
else
std::memcpy(mine.data(), values[ex.id].data(), nb);
protocol::detail::reconstruct_open(d, c.party(), mine.data(),
peer_bytes, nb, values[ex.id].data(), opt.field);
}
}
}
}
/// @brief Schedule a fused Express/Sabre point+audit walk on a composer.
/// @details Prefer this over `open_and_sketch` when the sketch can ride in the
/// last CW flush — one fewer exchange-bearing wave.
inline protocol::composer::walk_result schedule_fused_audit(
protocol::composer & c, protocol::node seed, std::size_t depth,
std::size_t slot_bytes, protocol::node sketch)
{
return c.fss_point_fused(seed, depth, slot_bytes, sketch);
}
/// @brief Fold opened outputs into a sketch and exchange for `sketch_verify`.
/// @details Used only by extractable party flows. Prefer Express/Sabre
/// `schedule_fused_audit` / `fss_point_fused` so the sketch rides in
/// the last CW flush instead of adding a round. Non-extractable keys
/// leave `note_sketch` as a no-op; callers still exchange empty shares.
template <typename KeyT, typename YRange, typename RRange>
bool open_and_sketch(trio & net, role self, sketch_share & local,
YRange && ys, RRange && rs)
{
note_sketch<KeyT>(local, std::forward<YRange>(ys), std::forward<RRange>(rs));
role peer = self == role::p0 ? role::p1 : role::p0;
sketch_share theirs = net.exchange_with(peer, local, net::msg::sketch_share);
if (self == role::p0)
return sketch_verify(local, theirs);
return sketch_verify(theirs, local);
}
/// @brief Keep parties aligned across `--repeat` iterations.
inline void sync_round(trio & net, role self, std::uint64_t round)
{
if (self == role::p2)
{
net.to(role::p0).send(net::msg::case_ok, round);
net.to(role::p1).send(net::msg::case_ok, round);
return;
}
auto got = net.to(role::p2).recv<std::uint64_t>(net::msg::case_ok);
require(got == round, "repeat sync");
}
} // namespace util
} // namespace party
} // namespace dpf
#endif // LIBDPF_PARTY_FLOW_UTIL_HPP__