libdpf/include/dpf/net/stream_sink.hpp

236 lines
7.4 KiB
C++
Raw Permalink Normal View History

/// @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 <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <memory>
#include <stdexcept>
#include <string>
#include <thread>
#include <utility>
#include <vector>
#include <unistd.h>
#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<channel> links, unsigned self_id, unsigned peer_id,
std::size_t count, std::vector<std::size_t> 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<std::uint8_t> payload;
if (nslots != 0)
{
payload.resize(sizeof(std::uint32_t) * 2 + nslots * slot_bytes_[r]);
const std::uint32_t b = static_cast<std::uint32_t>(begin);
const std::uint32_t c = static_cast<std::uint32_t>(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<std::uint8_t> 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<std::size_t>(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<channel> links_;
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;
};
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<std::size_t> 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<channel> 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__