270 lines
8.5 KiB
C++
270 lines
8.5 KiB
C++
|
|
/// @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__
|