libdpf/include/dpf/protocol.hpp

707 lines
25 KiB
C++
Raw Permalink Normal View History

/// @file dpf/protocol.hpp
/// @brief Count-sized batch sessions over a RoundSink.
/// @details Interactive rounds share one instance index across a PRG lane,
/// a dealer cursor, and a round window. `two_move_session` is the
/// typed custom-protocol helper. `schedule_session` drives a flat
/// round list: ready instances run `produce` on this thread; others
/// wait while the scan continues. Flush sends the largest unflushed
/// prefix. Mixed FSS / ABY2.0 / Beaver / RSS strands lower onto this
/// loop with `plan_to_schedule` in `compose.hpp`.
#ifndef LIBDPF_INCLUDE_DPF_PROTOCOL_HPP__
#define LIBDPF_INCLUDE_DPF_PROTOCOL_HPP__
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <functional>
#include <memory>
#include <optional>
#include <stdexcept>
#include <type_traits>
#include <utility>
#include <vector>
#include <time.h>
#include "dpf/net/dealer_cursor.hpp"
#include "dpf/net/edge_mesh.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/prg_count.hpp"
#include "dpf/random.hpp"
#include "dpf/thread_work.hpp"
namespace dpf
{
namespace protocol
{
using net::empty_pad;
using net::dealer_cursor;
using net::RoundSink;
using net::edge_id;
using net::edge_mesh;
using net::edge_peer;
using net::edge_rss_next;
using net::edge_dealer;
using net::edge_dealer_p1;
namespace detail
{
template <typename T>
constexpr bool is_empty_pad_v = std::is_same_v<std::decay_t<T>, empty_pad>;
} // namespace detail
/// @brief Named trio edges (aliases of `net::edge_*` for existing call sites).
enum class edge_channel : unsigned char
{
peer = 0,
rss_next = 1,
dealer = 2
};
/// @brief How `produce` / the finish step interpret peer bytes for a round.
enum class receive_rule : unsigned char
{
domain_open = 0, ///< Domain algebra (`a`/`fss` sum, `b` subtract, `y` copy)
copy_peer = 1, ///< Client upload / answer, dealer delivery
beaver = 2, ///< Authenticated / classic Beaver δ apply
xor_bytes = 3,
field_sum = 4,
any_two = 5, ///< Any-two Shamir-style reconstruct
verify_sketch = 6,
verify_proof = 7,
eq_check = 8, ///< PSI / equality tag check
ring_next = 9 ///< Directional RSS ring: copy bytes, do not reconstruct
};
/// @brief Setup vs online for byte tallies and cost reports.
enum class phase : unsigned char
{
online = 0,
setup = 1
};
/// @brief Direction of a schedule round on a duplex edge.
enum class round_dir : unsigned char
{
duplex = 0, ///< Submit and expect a peer slot (default)
send_next = 1, ///< Submit on next edge; no peer payload required to advance
recv_prev = 2 ///< Wait for prev party's send; produce may be empty
};
/// @brief One RoundSink exchange observed by a live probe.
struct round_event
{
std::uint16_t round = 0;
edge_id edge = edge_peer;
edge_channel channel = edge_channel::peer;
phase round_phase = phase::online;
/// Slot bytes this round submits.
std::size_t bytes_out = 0;
/// Slot bytes this round receives on `channel` (0 for `send_next` rounds).
std::size_t bytes_in = 0;
/// Time in `produce` and `submit`; waiting for the peer is not included.
std::uint64_t wall_ns = 0;
std::uint64_t cpu_ns = 0;
std::uint64_t prg_evals = 0;
std::uint64_t random_bytes = 0;
};
/// @brief Accumulated setup / online bytes from `round_event`s.
struct phase_tally
{
std::uint64_t setup_bytes_out = 0;
std::uint64_t setup_bytes_in = 0;
std::uint64_t online_bytes_out = 0;
std::uint64_t online_bytes_in = 0;
phase_tally & operator+=(const round_event & ev) noexcept
{
if (ev.round_phase == phase::setup)
{
setup_bytes_out += ev.bytes_out;
setup_bytes_in += ev.bytes_in;
}
else
{
online_bytes_out += ev.bytes_out;
online_bytes_in += ev.bytes_in;
}
return *this;
}
};
/// @brief Optional live probe installed on `drive_options` / `schedule_session`.
struct round_probe
{
void * ctx = nullptr;
void (*fn)(void * ctx, const round_event & ev) = nullptr;
void operator()(const round_event & ev) const
{
if (fn != nullptr)
fn(ctx, ev);
}
};
/// @brief One interactive round in a type-erased schedule.
struct schedule_round
{
std::size_t slot_bytes = 0;
edge_id edge = edge_peer;
edge_channel channel = edge_channel::peer; ///< mirrors `edge` for trio APIs
receive_rule recv = receive_rule::domain_open;
phase round_phase = phase::online;
round_dir dir = round_dir::duplex;
/// @brief Index on the chosen channel's RoundSink (0..that sink's rounds).
std::uint16_t sink_round = 0;
/// @brief Longest wait for this round's peer bytes (0 = edge / drive default).
std::chrono::milliseconds timeout{0};
/// @brief After the previous round's peer arrives: return false to send
/// nothing and mark the instance done (branch after an open).
std::function<bool(std::size_t index, const std::uint8_t * peer,
std::size_t peer_n)> branch;
/// @brief Jump after open: return next schedule index, or nullopt = done.
/// Default (empty) advances to `round + 1`.
std::function<std::optional<std::uint16_t>(std::size_t index,
const std::uint8_t * peer, std::size_t peer_n)> next;
/// @brief Write this party's slot for `index`. `peer` is the previous
/// round's peer slot (nullptr / 0 on round 0).
std::function<void(std::size_t index, const std::uint8_t * peer,
std::size_t peer_n, std::uint8_t * out)> produce;
};
/// @brief Trio adapter → `edge_mesh` (peer / rss_next / dealer at ids 0..2).
struct edge_sinks
{
RoundSink * peer = nullptr;
RoundSink * rss_next = nullptr;
RoundSink * dealer = nullptr;
RoundSink & at(edge_channel c) const
{
return mesh().at(static_cast<edge_id>(c));
}
edge_mesh mesh() const
{
edge_mesh m;
m.sinks = {peer, rss_next, dealer};
return m;
}
};
/// @brief Concatenate two round lists, preserving each round's sink_round.
inline std::vector<schedule_round> splice_rounds(std::vector<schedule_round> head,
std::vector<schedule_round> tail)
{
head.reserve(head.size() + tail.size());
for (auto & r : tail)
head.push_back(std::move(r));
return head;
}
/// @brief Opaque pad / setup frames as schedule rounds (IKNP, dealer tape, …).
/// @details Each frame exchanges `slot_bytes` into `tape` (XOR of peer into the
/// running lane). Splice before the online plan so setup is rounds in
/// the list, not a count added after `plan::rounds()`.
inline std::vector<schedule_round> make_pad_rounds(std::size_t n_rounds,
std::size_t slot_bytes,
const std::shared_ptr<std::vector<std::uint8_t>> & tape,
edge_channel channel = edge_channel::peer)
{
if (slot_bytes == 0 && n_rounds != 0)
throw std::invalid_argument("make_pad_rounds slot_bytes");
if (tape)
{
const std::size_t need = n_rounds * slot_bytes;
if (tape->size() < need)
tape->assign(need, 0);
}
std::vector<schedule_round> out(n_rounds);
for (std::size_t r = 0; r < n_rounds; ++r)
{
out[r].slot_bytes = slot_bytes;
out[r].channel = channel;
out[r].edge = static_cast<edge_id>(channel);
out[r].recv = receive_rule::xor_bytes;
out[r].round_phase = phase::setup;
out[r].dir = round_dir::send_next;
out[r].sink_round = static_cast<std::uint16_t>(r);
out[r].produce = [r, slot_bytes, tape](std::size_t /*index*/,
const std::uint8_t * peer, std::size_t peer_n,
std::uint8_t * out_slot) {
if (slot_bytes == 0 || out_slot == nullptr)
return;
const std::size_t off = r * slot_bytes;
std::memset(out_slot, static_cast<int>(r + 1), slot_bytes);
if (tape && tape->size() >= off + slot_bytes)
std::memcpy(tape->data() + off, out_slot, slot_bytes);
// `peer` is the previous round's slot — fold it into that round's tape.
if (r > 0 && peer != nullptr && peer_n >= slot_bytes && tape)
{
const std::size_t prev = (r - 1) * slot_bytes;
if (tape->size() >= prev + slot_bytes)
{
for (std::size_t i = 0; i < slot_bytes; ++i)
(*tape)[prev + i] = static_cast<std::uint8_t>(
(*tape)[prev + i] ^ peer[i]);
}
}
};
}
return out;
}
/// @brief Drive a list of schedule rounds on an `edge_mesh`.
/// @details Ready instances run `produce` on this thread; others are skipped
/// until `peer_ready`. Flush sends the largest unflushed prefix per
/// live edge. Peer/out scratch is thread_local (no per-ready alloc).
class schedule_session
{
public:
schedule_session(std::size_t count, RoundSink & peer,
std::vector<schedule_round> rounds, std::size_t pipeline_credit = 0)
: schedule_session(count, edge_mesh{{&peer}}, std::move(rounds),
/*rebase_single=*/true, nullptr, pipeline_credit)
{
}
schedule_session(std::size_t count, edge_sinks sinks,
std::vector<schedule_round> rounds, std::size_t pipeline_credit = 0)
: schedule_session(count, sinks.mesh(), std::move(rounds),
/*rebase_single=*/sinks.rss_next == nullptr
&& sinks.dealer == nullptr,
nullptr, pipeline_credit)
{
}
schedule_session(std::size_t count, edge_mesh mesh,
std::vector<schedule_round> rounds, bool rebase_single = false,
const round_probe * probe = nullptr, std::size_t pipeline_credit = 0)
: count_(count),
mesh_(std::move(mesh)),
rounds_(std::move(rounds)),
round_of_(count, 0),
mark_(count, 0),
done_(count, false),
submitted_(rounds_.size() * count, false),
probe_(probe),
pipeline_credit_(pipeline_credit)
{
if (!mesh_.has(edge_peer) && mesh_.size() == 0)
throw std::invalid_argument("schedule_session empty mesh");
// Sync channel ↔ edge for trio-named rounds.
for (auto & r : rounds_)
{
if (r.edge == edge_peer
&& r.channel != edge_channel::peer)
r.edge = static_cast<edge_id>(r.channel);
else if (r.edge <= edge_dealer)
r.channel = static_cast<edge_channel>(r.edge);
}
if (rebase_single && mesh_.has(edge_peer)
&& !mesh_.has(edge_rss_next) && !mesh_.has(edge_dealer)
&& mesh_.size() <= 3)
{
auto & peer = mesh_.at(edge_peer);
if (peer.count() != count_)
throw std::invalid_argument("schedule_session count mismatch");
if (peer.rounds() != rounds_.size())
throw std::invalid_argument("schedule_session rounds mismatch");
for (std::uint16_t r = 0; r < rounds_.size(); ++r)
{
rounds_[r].edge = edge_peer;
rounds_[r].channel = edge_channel::peer;
rounds_[r].sink_round = r;
if (peer.slot_bytes(r) != rounds_[r].slot_bytes)
throw std::invalid_argument("schedule_session slot mismatch");
}
}
else
{
for (std::uint16_t r = 0; r < rounds_.size(); ++r)
{
auto & sink = mesh_.at(rounds_[r].edge);
if (sink.count() != count_)
throw std::invalid_argument("schedule_session edge count");
const auto sr = rounds_[r].sink_round;
if (static_cast<std::size_t>(sr) >= sink.rounds()
|| sink.slot_bytes(sr) != rounds_[r].slot_bytes)
throw std::invalid_argument("schedule_session edge slot");
}
}
}
std::size_t count() const noexcept { return count_; }
std::size_t rounds() const noexcept { return rounds_.size(); }
edge_mesh & mesh() noexcept { return mesh_; }
const edge_mesh & mesh() const noexcept { return mesh_; }
/// @brief Changes whenever instance `index` enters a new round or finishes.
std::uint64_t mark(std::size_t index) const { return mark_.at(index); }
std::uint16_t current_round(std::size_t index) const
{
return round_of_.at(index);
}
edge_id current_edge(std::size_t index) const
{
const auto r = round_of_.at(index);
return r < rounds_.size() ? rounds_[r].edge : edge_peer;
}
std::chrono::milliseconds current_timeout(std::size_t index) const
{
const auto r = round_of_.at(index);
return r < rounds_.size() ? rounds_[r].timeout
: std::chrono::milliseconds(0);
}
/// @brief Kick instance `index` into round 0 (no peer bytes yet).
void submit(std::size_t index)
{
if (index >= count_)
throw std::out_of_range("schedule_session submit");
if (submitted_at(0, index))
throw std::logic_error("schedule_session already submitted");
run_round(index, 0, nullptr, 0);
}
void drive()
{
bool progress = true;
while (progress)
{
progress = false;
mesh_.flush_all();
const std::uint64_t received = mesh_.progress_total();
for (std::size_t i = 0; i < count_; ++i)
{
if (done_[i])
continue;
// Pipeline: submit independent future rounds up to credit.
if (pipeline_credit_ > 0)
{
const auto cur = static_cast<std::uint16_t>(round_of_[i]);
for (std::size_t ahead = 1; ahead <= pipeline_credit_;
++ahead)
{
const std::size_t cand =
static_cast<std::size_t>(cur) + ahead;
if (cand >= rounds_.size() || submitted_at(
static_cast<std::uint16_t>(cand), i))
break;
if (rounds_[cand].dir != round_dir::send_next
&& rounds_[cand].dir != round_dir::duplex)
break;
// Independent: produce must not require peer bytes.
if (cand > 0 && rounds_[cand].dir != round_dir::send_next)
break;
if (!mesh_.at(rounds_[cand].edge).can_send_ahead())
break;
run_round(i, static_cast<std::uint16_t>(cand), nullptr,
0);
progress = true;
}
}
const auto r = static_cast<std::uint16_t>(round_of_[i]);
if (!submitted_at(r, i))
continue;
std::optional<std::uint16_t> nxt;
auto & prev_sink = mesh_.at(rounds_[r].edge);
const auto prev_sr = rounds_[r].sink_round;
if (static_cast<std::size_t>(r) + 1 >= rounds_.size()
&& !rounds_[r].next)
{
done_[i] = true;
++mark_[i];
progress = true;
continue;
}
const bool needs_peer =
rounds_[r].dir != round_dir::send_next;
if (needs_peer && !prev_sink.peer_ready(prev_sr, i))
continue;
auto & peer_buf = scratch_peer();
peer_buf.clear();
if (needs_peer)
{
peer_buf.resize(rounds_[r].slot_bytes);
if (!peer_buf.empty())
prev_sink.read_peer(prev_sr, i, peer_buf.data(),
peer_buf.size());
}
if (rounds_[r].next)
nxt = rounds_[r].next(i, peer_buf.empty() ? nullptr
: peer_buf.data(),
peer_buf.size());
else if (static_cast<std::size_t>(r) + 1 < rounds_.size())
nxt = static_cast<std::uint16_t>(r + 1);
else
{
done_[i] = true;
++mark_[i];
progress = true;
continue;
}
if (!nxt)
{
done_[i] = true;
++mark_[i];
progress = true;
continue;
}
if (*nxt >= rounds_.size())
throw std::out_of_range("schedule_session next");
if (submitted_at(*nxt, i))
{
// Already pipelined; just advance cursor.
round_of_[i] = *nxt;
++mark_[i];
progress = true;
continue;
}
run_round(i, *nxt, peer_buf.empty() ? nullptr : peer_buf.data(),
peer_buf.size());
progress = true;
}
mesh_.flush_all();
// A read that completed inside that poll must not wait for the
// caller's next blocking wait.
if (mesh_.progress_total() != received)
progress = true;
}
}
const phase_tally & tally() const noexcept { return tally_; }
bool done(std::size_t index) const
{
if (index >= count_)
throw std::out_of_range("schedule_session done");
return done_[index];
}
private:
static std::vector<std::uint8_t> & scratch_peer()
{
thread_local std::vector<std::uint8_t> buf;
return buf;
}
static std::vector<std::uint8_t> & scratch_out()
{
thread_local std::vector<std::uint8_t> buf;
return buf;
}
bool submitted_at(std::uint16_t round, std::size_t index) const
{
return submitted_[static_cast<std::size_t>(round) * count_ + index];
}
void mark_submitted(std::uint16_t round, std::size_t index)
{
submitted_[static_cast<std::size_t>(round) * count_ + index] = true;
}
void run_round(std::size_t index, std::uint16_t round,
const std::uint8_t * peer, std::size_t peer_n)
{
auto & step = rounds_[round];
if (step.branch && !step.branch(index, peer, peer_n))
{
done_[index] = true;
round_of_[index] = round;
++mark_[index];
return;
}
if (!step.produce)
throw std::logic_error("schedule_session missing produce");
const std::uint64_t prg0 = prg::eval_count();
const std::uint64_t rnd0 = random_bytes_count();
const std::uint64_t wall0 = probe_ != nullptr ? probe_wall_ns_() : 0;
const std::uint64_t cpu0 = probe_ != nullptr ? probe_cpu_ns_() : 0;
auto & out = scratch_out();
out.assign(step.slot_bytes, 0);
step.produce(index, peer, peer_n, out.empty() ? nullptr : out.data());
auto & sink = mesh_.at(step.edge);
sink.submit(step.sink_round, index, out.empty() ? nullptr : out.data(),
out.size());
mark_submitted(round, index);
round_of_[index] = round;
++mark_[index];
{
round_event ev;
ev.round = round;
ev.edge = step.edge;
ev.channel = step.channel;
ev.round_phase = step.round_phase;
ev.bytes_out = out.size();
ev.bytes_in = step.dir == round_dir::send_next ? 0 : step.slot_bytes;
ev.wall_ns = probe_ != nullptr ? probe_wall_ns_() - wall0 : 0;
ev.cpu_ns = probe_ != nullptr ? probe_cpu_ns_() - cpu0 : 0;
ev.prg_evals = prg::eval_count() - prg0;
ev.random_bytes = random_bytes_count() - rnd0;
tally_ += ev;
if (probe_ != nullptr && probe_->fn != nullptr)
(*probe_)(ev);
}
// Done only when there is no further default successor and no jump.
if (!step.next && static_cast<std::size_t>(round) + 1 >= rounds_.size())
done_[index] = true;
}
static std::uint64_t probe_wall_ns_()
{
using clock = std::chrono::steady_clock;
return static_cast<std::uint64_t>(
std::chrono::duration_cast<std::chrono::nanoseconds>(
clock::now().time_since_epoch())
.count());
}
static std::uint64_t probe_cpu_ns_()
{
if (have_thread_cpu_clock())
return thread_cpu_ns();
return probe_wall_ns_();
}
std::size_t count_;
edge_mesh mesh_;
std::vector<schedule_round> rounds_;
std::vector<std::uint16_t> round_of_;
std::vector<std::uint64_t> mark_;
std::vector<bool> done_;
std::vector<bool> submitted_;
const round_probe * probe_ = nullptr;
std::size_t pipeline_credit_ = 0;
phase_tally tally_{};
};
/// @brief Two-move batch: `(in, blind) -> (fwd, swap)` then
/// `(fwd, peer, blind, corr) -> out`.
template <typename Fwd, typename Swap, typename Blind, typename Corr,
typename Output, typename Prg, typename Move0, typename Move1>
class two_move_session
{
public:
two_move_session(std::size_t count, RoundSink & peer,
dealer_cursor & dealer, Prg & prg, Move0 move0, Move1 move1)
: count_(count),
peer_(peer),
dealer_(dealer),
prg_(prg),
move0_(std::move(move0)),
move1_(std::move(move1)),
fwds_(count),
outs_(count),
have_in_(count, false),
have_swap_(count, false),
have_out_(count, false)
{
if (peer_.count() != count || peer_.rounds() < 1)
throw std::invalid_argument("two_move_session sink size");
if (peer_.slot_bytes(0) != sizeof(Swap))
throw std::invalid_argument("two_move_session slot");
}
template <typename Input>
void submit(std::size_t index, Input in)
{
if (index >= count_ || have_in_[index])
throw std::logic_error("two_move_session submit");
have_in_[index] = true;
Blind blind{};
if constexpr (!detail::is_empty_pad_v<Blind>)
blind = prg_.template at<0>(static_cast<std::uint64_t>(index));
auto [fwd, swap] = move0_(in, blind);
fwds_[index] = std::move(fwd);
static_assert(std::is_trivially_copyable_v<Swap>,
"two_move_session swap must be trivially copyable");
peer_.submit(0, index, reinterpret_cast<const std::uint8_t *>(&swap),
sizeof(Swap));
have_swap_[index] = true;
}
void drive()
{
for (;;)
{
peer_.flush();
peer_.poll();
bool progress = false;
for (std::size_t i = 0; i < count_; ++i)
{
if (!have_swap_[i] || have_out_[i])
continue;
if (!peer_.peer_ready(0, i))
continue;
Swap peer_swap{};
peer_.read_peer(0, i,
reinterpret_cast<std::uint8_t *>(&peer_swap), sizeof(Swap));
Blind blind{};
if constexpr (!detail::is_empty_pad_v<Blind>)
{
if constexpr (Prg::stream_count > 1)
blind = prg_.template at<1>(
static_cast<std::uint64_t>(i));
else
blind = prg_.template at<0>(
static_cast<std::uint64_t>(i));
}
Corr corr{};
if constexpr (!detail::is_empty_pad_v<Corr>)
corr = dealer_.template at<Corr>(0, i);
outs_[i] = move1_(fwds_[i], peer_swap, blind, corr);
have_out_[i] = true;
progress = true;
}
peer_.flush();
if (!progress)
break;
}
for (std::size_t i = 0; i < count_; ++i)
{
if (have_in_[i] && !have_out_[i])
throw std::logic_error("two_move_session stuck");
}
}
Output take(std::size_t index) const
{
if (index >= count_ || !have_out_[index])
throw std::logic_error("two_move_session take");
return outs_[index];
}
private:
std::size_t count_;
RoundSink & peer_;
dealer_cursor & dealer_;
Prg & prg_;
Move0 move0_;
Move1 move1_;
std::vector<Fwd> fwds_;
std::vector<Output> outs_;
std::vector<bool> have_in_;
std::vector<bool> have_swap_;
std::vector<bool> have_out_;
};
template <typename Fwd, typename Swap, typename Blind, typename Corr,
typename Output, typename Prg, typename Move0, typename Move1>
auto make_two_move_session(std::size_t count, RoundSink & peer,
dealer_cursor & dealer, Prg & prg, Move0 move0, Move1 move1)
{
return two_move_session<Fwd, Swap, Blind, Corr, Output, Prg, Move0, Move1>(
count, peer, dealer, prg, std::move(move0), std::move(move1));
}
} // namespace protocol
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_PROTOCOL_HPP__