libdpf/include/dpf/net/mux_sink.hpp

270 lines
8.5 KiB
C++
Raw Normal View History

/// @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__