811 lines
31 KiB
C++
811 lines
31 KiB
C++
|
|
/// @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__
|