/// @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 #include #include #include #include #include #include #include #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 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 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 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 & 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(begin); hdr.count = static_cast(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 & payload) { // Always exchange so a party with nothing new still receives. std::vector 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 & 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(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 slot_bytes_; std::vector 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 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(peer_id)).send_bytes(tag, data, n); } std::vector recv_bytes(trio & net, unsigned peer_id, msg tag) override { return net.to(static_cast(peer_id)).recv_bytes(tag); } std::vector exchange_bytes(trio & net, unsigned peer_id, msg tag, const void * data, std::size_t n) override { auto & link = net.to(static_cast(peer_id)); const unsigned self_id = to_u(net.self()); std::vector 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 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 batch(trio & net, std::size_t count, std::vector slot_bytes) override { return std::make_unique(make_mux_sink(net, count, std::move(slot_bytes))); } }; inline std::unique_ptr trio::batch(std::size_t count, std::vector slot_bytes) { if (hook_) return hook_->batch(*this, count, std::move(slot_bytes)); return std::make_unique(make_mux_sink(*this, count, std::move(slot_bytes))); } } // namespace net } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_NET_MUX_SINK_HPP__