libdpf/include/dpf/net/mux_sink.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
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>
2026-09-28 05:59:19 -06:00

269 lines
8.5 KiB
C++
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @file dpf/net/mux_sink.hpp
/// @brief RoundSink multiplexed on one framed trio channel.
#ifndef LIBDPF_INCLUDE_DPF_NET_MUX_SINK_HPP__
#define LIBDPF_INCLUDE_DPF_NET_MUX_SINK_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <memory>
#include <stdexcept>
#include <utility>
#include <vector>
#include "dpf/net/channel.hpp"
#include "dpf/net/comm_hook.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/net/trio.hpp"
namespace dpf
{
namespace net
{
#pragma pack(push, 1)
struct round_batch_hdr
{
std::uint16_t round = 0;
std::uint32_t begin = 0;
std::uint32_t count = 0;
};
#pragma pack(pop)
/// @brief Prefix-flush RoundSink over a single duplex `channel`.
/// @details Each flush sends one `msg::round_batch` frame per round that has
/// a new contiguous prefix: header `(round, begin, count)` then
/// `count * slot_bytes` payload. The peer's matching frames are
/// read in the same flush (lower role sends first).
class mux_sink : public RoundSink
{
public:
mux_sink(channel & link, unsigned self_id, unsigned peer_id,
std::size_t count, std::vector<std::size_t> slot_bytes)
: link_(link),
self_id_(self_id),
peer_id_(peer_id),
count_(count),
slot_bytes_(std::move(slot_bytes))
{
windows_.reserve(slot_bytes_.size());
for (std::size_t sb : slot_bytes_)
windows_.emplace_back(count_, sb);
}
std::size_t count() const noexcept override { return count_; }
std::size_t rounds() const noexcept override { return slot_bytes_.size(); }
std::size_t slot_bytes(std::uint16_t round) const override
{
if (round >= slot_bytes_.size())
throw std::out_of_range("mux_sink round");
return slot_bytes_[round];
}
void submit(std::uint16_t round, std::size_t index,
const std::uint8_t * bytes, std::size_t n) override
{
window(round).submit(index, bytes, n);
}
bool peer_ready(std::uint16_t round, std::size_t index) const override
{
return window(round).peer_ready(index);
}
void read_peer(std::uint16_t round, std::size_t index, std::uint8_t * out,
std::size_t n) const override
{
window(round).read_peer(index, out, n);
}
void flush() override
{
std::vector<std::uint8_t> payload;
for (std::uint16_t r = 0; r < slot_bytes_.size(); ++r)
append_pending(r, payload);
exchange_payload(payload);
}
void flush_round(std::uint16_t round) override
{
if (round >= slot_bytes_.size())
throw std::out_of_range("mux_sink flush_round");
std::vector<std::uint8_t> payload;
append_pending(round, payload);
exchange_payload(payload);
}
void poll() override {}
std::uint64_t exchanges() const noexcept { return link_exchanges_; }
private:
round_window & window(std::uint16_t round)
{
if (round >= windows_.size())
throw std::out_of_range("mux_sink round");
return windows_[round];
}
const round_window & window(std::uint16_t round) const
{
if (round >= windows_.size())
throw std::out_of_range("mux_sink round");
return windows_[round];
}
void append_pending(std::uint16_t r, std::vector<std::uint8_t> & payload)
{
std::size_t begin = 0;
std::size_t nslots = 0;
const std::uint8_t * pending = windows_[r].pending_out(begin, nslots);
if (nslots == 0)
return;
round_batch_hdr hdr{};
hdr.round = r;
hdr.begin = static_cast<std::uint32_t>(begin);
hdr.count = static_cast<std::uint32_t>(nslots);
const std::size_t body = nslots * slot_bytes_[r];
const std::size_t old = payload.size();
payload.resize(old + sizeof(hdr) + body);
std::memcpy(payload.data() + old, &hdr, sizeof(hdr));
if (body != 0)
std::memcpy(payload.data() + old + sizeof(hdr), pending, body);
windows_[r].mark_flushed(nslots);
}
void exchange_payload(const std::vector<std::uint8_t> & payload)
{
// Always exchange so a party with nothing new still receives.
std::vector<std::uint8_t> theirs;
if (self_id_ < peer_id_)
{
link_.send_bytes(msg::round_batch, payload);
theirs = link_.recv_bytes(msg::round_batch);
}
else if (self_id_ > peer_id_)
{
theirs = link_.recv_bytes(msg::round_batch);
link_.send_bytes(msg::round_batch, payload);
}
else
throw std::invalid_argument("mux_sink flush with self");
++link_exchanges_;
ingest(theirs);
}
void ingest(const std::vector<std::uint8_t> & bytes)
{
std::size_t off = 0;
while (off < bytes.size())
{
if (off + sizeof(round_batch_hdr) > bytes.size())
throw std::runtime_error("mux_sink truncated header");
round_batch_hdr hdr{};
std::memcpy(&hdr, bytes.data() + off, sizeof(hdr));
off += sizeof(hdr);
if (hdr.round >= slot_bytes_.size())
throw std::runtime_error("mux_sink bad round");
const std::size_t body =
static_cast<std::size_t>(hdr.count) * slot_bytes_[hdr.round];
if (off + body > bytes.size())
throw std::runtime_error("mux_sink truncated body");
windows_[hdr.round].accept_peer_at(hdr.begin, bytes.data() + off,
hdr.count);
off += body;
}
}
channel & link_;
unsigned self_id_ = 0;
unsigned peer_id_ = 0;
std::size_t count_ = 0;
std::vector<std::size_t> slot_bytes_;
std::vector<round_window> windows_;
std::uint64_t link_exchanges_ = 0;
};
/// @brief Bind a mux sink to the p0–p1 link of an existing trio.
inline mux_sink make_mux_sink(trio & net, std::size_t count,
std::vector<std::size_t> slot_bytes)
{
const role self = net.self();
if (self == role::p2)
throw std::logic_error("mux_sink is for computing parties");
const role peer = self == role::p0 ? role::p1 : role::p0;
return mux_sink(net.to(peer), to_u(self), to_u(peer), count,
std::move(slot_bytes));
}
/// @brief Default `comm_hook`: framed mesh channels and a mux batch sink.
/// @details Subclass to reroute selected calls; leave the rest to these
/// defaults. Installing this hook is equivalent to a null hook for
/// framing, but lets overrides replace individual methods.
class mesh_comm_hook : public comm_hook
{
public:
void send_bytes(trio & net, unsigned peer_id, msg tag, const void * data,
std::size_t n) override
{
net.to(static_cast<role>(peer_id)).send_bytes(tag, data, n);
}
std::vector<std::uint8_t> recv_bytes(trio & net, unsigned peer_id,
msg tag) override
{
return net.to(static_cast<role>(peer_id)).recv_bytes(tag);
}
std::vector<std::uint8_t> exchange_bytes(trio & net, unsigned peer_id,
msg tag, const void * data, std::size_t n) override
{
auto & link = net.to(static_cast<role>(peer_id));
const unsigned self_id = to_u(net.self());
std::vector<std::uint8_t> theirs(n);
if (self_id < peer_id)
{
link.send_bytes(tag, data, n);
theirs = link.recv_bytes(tag);
}
else if (self_id > peer_id)
{
theirs = link.recv_bytes(tag);
link.send_bytes(tag, data, n);
}
else
throw std::invalid_argument("mesh_comm_hook exchange with self");
if (theirs.size() != n)
throw std::runtime_error("mesh_comm_hook exchange size");
return theirs;
}
std::vector<std::uint8_t> exchange_vec_bytes(trio & net, unsigned peer_id,
msg tag, const void * data, std::size_t nbytes) override
{
return exchange_bytes(net, peer_id, tag, data, nbytes);
}
std::unique_ptr<RoundSink> batch(trio & net, std::size_t count,
std::vector<std::size_t> slot_bytes) override
{
return std::make_unique<mux_sink>(make_mux_sink(net, count,
std::move(slot_bytes)));
}
};
inline std::unique_ptr<RoundSink> trio::batch(std::size_t count,
std::vector<std::size_t> slot_bytes)
{
if (hook_)
return hook_->batch(*this, count, std::move(slot_bytes));
return std::make_unique<mux_sink>(make_mux_sink(*this, count,
std::move(slot_bytes)));
}
} // namespace net
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_NET_MUX_SINK_HPP__