/// @file dpf/net/stream_sink.hpp /// @brief RoundSink with one duplex channel per interactive round. #ifndef LIBDPF_INCLUDE_DPF_NET_STREAM_SINK_HPP__ #define LIBDPF_INCLUDE_DPF_NET_STREAM_SINK_HPP__ #include #include #include #include #include #include #include #include #include #include #include #include "dpf/net/asio_ns.hpp" #include "dpf/net/channel.hpp" #include "dpf/net/round_sink.hpp" #include "dpf/net/trio.hpp" namespace dpf { namespace net { /// @brief Prefix-flush RoundSink with an independent stream per round. class stream_sink : public RoundSink { public: stream_sink(std::vector links, unsigned self_id, unsigned peer_id, std::size_t count, std::vector slot_bytes) : links_(std::move(links)), self_id_(self_id), peer_id_(peer_id), count_(count), slot_bytes_(std::move(slot_bytes)) { if (links_.size() != slot_bytes_.size()) throw std::invalid_argument("stream_sink link count"); 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("stream_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 { // Global barrier: every round channel must rendezvous. Parties can have // pending on different rounds (schedule_session holes / staggered // indices); skipping idle links deadlocks the peer that still waits // there. Lockstep callers should use flush_round to avoid O(rounds) // exchanges on large sinks. for (std::uint16_t r = 0; r < slot_bytes_.size(); ++r) flush_one(r); } void flush_round(std::uint16_t r) override { if (r >= slot_bytes_.size()) throw std::out_of_range("stream_sink flush_round"); flush_one(r); } 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("stream_sink round"); return windows_[round]; } const round_window & window(std::uint16_t round) const { if (round >= windows_.size()) throw std::out_of_range("stream_sink round"); return windows_[round]; } void flush_one(std::uint16_t r) { std::size_t begin = 0; std::size_t nslots = 0; const std::uint8_t * pending = windows_[r].pending_out(begin, nslots); std::vector payload; if (nslots != 0) { payload.resize(sizeof(std::uint32_t) * 2 + nslots * slot_bytes_[r]); const std::uint32_t b = static_cast(begin); const std::uint32_t c = static_cast(nslots); std::memcpy(payload.data(), &b, 4); std::memcpy(payload.data() + 4, &c, 4); if (nslots * slot_bytes_[r] != 0) std::memcpy(payload.data() + 8, pending, nslots * slot_bytes_[r]); windows_[r].mark_flushed(nslots); } else { // Peer may still be delivering; exchange an empty prefix header. payload.resize(8, 0); } std::vector theirs; if (self_id_ < peer_id_) { links_[r].send_bytes(msg::round_batch, payload); theirs = links_[r].recv_bytes(msg::round_batch); } else { theirs = links_[r].recv_bytes(msg::round_batch); links_[r].send_bytes(msg::round_batch, payload); } if (theirs.size() < 8) throw std::runtime_error("stream_sink short frame"); std::uint32_t pb = 0; std::uint32_t pc = 0; std::memcpy(&pb, theirs.data(), 4); std::memcpy(&pc, theirs.data() + 4, 4); const std::size_t body = static_cast(pc) * slot_bytes_[r]; if (theirs.size() != 8 + body) throw std::runtime_error("stream_sink size mismatch"); if (pc != 0) windows_[r].accept_peer_at(pb, theirs.data() + 8, pc); ++link_exchanges_; } std::vector links_; 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; }; inline std::string stream_link_path(const std::string & dir, role a, role b, std::uint16_t round) { if (to_u(a) > to_u(b)) std::swap(a, b); return dir + "/" + role_name(a) + "-" + role_name(b) + "-r" + std::to_string(round); } /// @brief Connect one unix-domain stream per round under `dir`. inline stream_sink connect_stream_sink(role self, role peer, const std::string & dir, std::size_t count, std::vector slot_bytes, unsigned retries = 200) { if (self == role::p2 || peer == role::p2) throw std::invalid_argument("stream_sink is for computing parties"); asio::io_context io; std::vector links; links.reserve(slot_bytes.size()); using proto = asio::local::stream_protocol; for (std::uint16_t r = 0; r < slot_bytes.size(); ++r) { const auto path = stream_link_path(dir, self, peer, r); const bool accept = to_u(self) < to_u(peer); bool opened = false; for (unsigned i = 0; i < retries && !opened; ++i) { asio::error_code ec; if (accept) { ::unlink(path.c_str()); proto::endpoint ep(path); proto::acceptor acc(io, ep); proto::socket sock(io); acc.accept(sock, ec); if (!ec) { links.emplace_back(std::move(sock)); opened = true; } } else { proto::endpoint ep(path); proto::socket sock(io); sock.connect(ep, ec); if (!ec) { links.emplace_back(std::move(sock)); opened = true; } } if (!opened) std::this_thread::sleep_for(std::chrono::milliseconds(10)); } if (!opened) throw std::runtime_error("stream_sink connect failed: " + path); } return stream_sink(std::move(links), to_u(self), to_u(peer), count, std::move(slot_bytes)); } } // namespace net } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_NET_STREAM_SINK_HPP__