Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
706 lines
25 KiB
C++
706 lines
25 KiB
C++
/// @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__
|