libdpf/include/dpf/net/async_sctp_stream_array.hpp

774 lines
25 KiB
C++
Raw Permalink Normal View History

/// @file dpf/net/async_sctp_stream_array.hpp
/// @brief Truly asynchronous SCTP-backed indexed byte streams (Linux/libsctp).
/// @details One SCTP association carries every lane: index `i` maps to SCTP
/// stream `i`. Receive is head-of-line free across streams. Writes are
/// split into `wire_policy::chunk_bytes` messages and sent round-robin
/// across streams, so a large write on one stream does not hold small
/// writes on another behind it. The window covers the association.
///
/// Platform contract:
/// * Linux with `<netinet/sctp.h>` (libsctp) → real backend, needs `-lsctp`.
/// * Everything else → the class exists, every constructor throws.
#ifndef LIBDPF_INCLUDE_DPF_NET_ASYNC_SCTP_STREAM_ARRAY_HPP__
#define LIBDPF_INCLUDE_DPF_NET_ASYNC_SCTP_STREAM_ARRAY_HPP__
#if !defined(DPF_HAS_LIBSCTP)
# if defined(__linux__) && defined(__has_include)
# if __has_include(<netinet/sctp.h>)
# define DPF_HAS_LIBSCTP 1
# else
# define DPF_HAS_LIBSCTP 0
# endif
# else
# define DPF_HAS_LIBSCTP 0
# endif
#endif
#include <atomic>
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <deque>
#include <exception>
#include <functional>
#include <memory>
#include <mutex>
#include <stdexcept>
#include <string>
#include <system_error>
#include <utility>
#include <vector>
#include "dpf/net/asio_ns.hpp"
#include "hedley/hedley.h"
#include "dpf/net/async_stream_array.hpp"
#include "dpf/net/connect.hpp"
#include "dpf/net/policy.hpp"
#if DPF_HAS_LIBSCTP
# include <arpa/inet.h>
# include <fcntl.h>
# include <netinet/in.h>
# include <netinet/sctp.h>
# include <sys/socket.h>
# include <unistd.h>
# include <cerrno>
#endif
namespace dpf
{
namespace net
{
/// @brief True when this build has a real SCTP backend.
inline constexpr bool sctp_available() noexcept
{
return DPF_HAS_LIBSCTP != 0;
}
#if DPF_HAS_LIBSCTP
namespace detail
{
/// @brief Stream counts, per-message `sinfo`, and socket options.
inline void sctp_configure(int fd, std::size_t nstreams,
const socket_options & o = {})
{
struct sctp_initmsg im;
std::memset(&im, 0, sizeof(im));
im.sinit_num_ostreams = static_cast<std::uint16_t>(nstreams);
im.sinit_max_instreams = static_cast<std::uint16_t>(nstreams);
(void)::setsockopt(fd, IPPROTO_SCTP, SCTP_INITMSG, &im, sizeof(im));
struct sctp_event_subscribe ev;
std::memset(&ev, 0, sizeof(ev));
ev.sctp_data_io_event = 1;
(void)::setsockopt(fd, IPPROTO_SCTP, SCTP_EVENTS, &ev, sizeof(ev));
const int nodelay = o.no_delay ? 1 : 0;
(void)::setsockopt(fd, IPPROTO_SCTP, SCTP_NODELAY, &nodelay, sizeof(nodelay));
if (o.send_buffer > 0)
(void)::setsockopt(fd, SOL_SOCKET, SO_SNDBUF, &o.send_buffer,
sizeof(o.send_buffer));
if (o.recv_buffer > 0)
(void)::setsockopt(fd, SOL_SOCKET, SO_RCVBUF, &o.recv_buffer,
sizeof(o.recv_buffer));
}
} // namespace detail
/// @brief N lanes over one SCTP association, fully overlapped.
class async_sctp_stream_array final : public async_stream_array
{
public:
/// @brief Adopt a connected one-to-one SCTP fd.
async_sctp_stream_array(asio::io_context & io, int fd, std::size_t nstreams,
const wire_policy & pol = {})
{
if (nstreams == 0)
throw std::invalid_argument("async_sctp_stream_array empty");
if (nstreams > 0xffffu)
throw std::invalid_argument("async_sctp_stream_array: stream ids are u16");
pol.validate();
detail::sctp_configure(fd, nstreams, pol.socket);
impl_ = std::make_shared<impl>(io, fd, nstreams, pol);
impl_->start();
}
async_sctp_stream_array(const async_sctp_stream_array &) = delete;
async_sctp_stream_array & operator=(const async_sctp_stream_array &) = delete;
~async_sctp_stream_array() override
{
if (impl_)
impl_->close_graceful();
}
std::size_t size() const noexcept override { return impl_->n; }
asio::io_context & context() noexcept override { return impl_->io; }
void async_write(std::size_t i, const void * src, std::size_t n,
async_handler h) override
{
impl_->write(i, n == 0 ? nullptr : copy_buffer(src, n), std::move(h));
}
void async_write_owned(std::size_t i,
std::shared_ptr<std::vector<std::uint8_t>> buf, async_handler h) override
{
impl_->write(i, std::move(buf), std::move(h));
}
void async_read(std::size_t i, void * dst, std::size_t n,
async_handler h) override
{
impl_->read(i, dst, n, std::move(h));
}
std::size_t buffered_bytes() const noexcept override
{
return impl_->buffered.load(std::memory_order_relaxed);
}
std::size_t window_bytes() const noexcept override
{
return impl_->window.load(std::memory_order_relaxed);
}
void set_window_bytes(std::size_t bytes) override
{
impl_->window.store(bytes, std::memory_order_relaxed);
}
stream_stats stats() const override { return impl_->stats(); }
void close() noexcept override
{
impl_->close(asio::error::operation_aborted);
}
private:
struct chunk
{
std::shared_ptr<std::vector<std::uint8_t>> payload;
std::size_t off = 0;
std::size_t len = 0;
std::shared_ptr<detail::write_op> op;
};
struct impl : std::enable_shared_from_this<impl>
{
impl(asio::io_context & io_, int fd, std::size_t nstreams,
const wire_policy & pol_)
: io(io_),
sd(io_),
strand(asio::make_strand(io_)),
n(nstreams),
pol(pol_),
inbox(nstreams),
partial(nstreams),
wait(nstreams),
outq(nstreams),
rbuf(std::max<std::size_t>(pol_.chunk_bytes, std::size_t{1} << 16)),
window(pol_.window_bytes)
{
const int fl = ::fcntl(fd, F_GETFL, 0);
if (fl >= 0)
(void)::fcntl(fd, F_SETFL, fl | O_NONBLOCK);
sd.assign(fd);
}
void check(std::size_t i) const
{
if (i >= n)
throw std::out_of_range("async_sctp_stream_array index "
+ std::to_string(i));
}
void start()
{
auto self = shared_from_this();
asio::post(strand, [self] { self->arm_read(); });
}
// --- reads ---------------------------------------------------------
void read(std::size_t i, void * dst, std::size_t nbytes, async_handler h)
{
check(i);
std::lock_guard<std::mutex> lock(mu);
if (fail)
{
detail::post_handler(io, std::move(h), fail);
return;
}
auto & w = wait[i];
if (w.active)
throw std::logic_error(
"async_sctp_stream_array: overlapping read on stream "
+ std::to_string(i));
const std::size_t k = inbox[i].take(dst, nbytes, pol.compact_bytes);
if (k == nbytes)
{
detail::post_handler(io, std::move(h), {});
return;
}
w.active = true;
w.dst = dst;
w.n = nbytes;
w.filled = k;
w.h = std::move(h);
}
void satisfy_locked(std::size_t i)
{
auto & w = wait[i];
if (!w.active)
return;
w.filled += inbox[i].take(static_cast<std::uint8_t *>(w.dst) + w.filled,
w.n - w.filled, pol.compact_bytes);
if (w.filled == w.n)
{
w.active = false;
detail::post_handler(io, std::move(w.h), {});
}
}
void arm_read()
{
auto self = shared_from_this();
sd.async_wait(asio::posix::stream_descriptor::wait_read,
asio::bind_executor(strand, [self](const std::error_code & ec) {
self->on_readable(ec);
}));
}
void on_readable(const std::error_code & ec)
{
if (ec)
{
deliver_error(ec);
return;
}
for (;;)
{
struct sctp_sndrcvinfo sinfo;
std::memset(&sinfo, 0, sizeof(sinfo));
int flags = 0;
const ssize_t r = ::sctp_recvmsg(sd.native_handle(), rbuf.data(),
rbuf.size(), nullptr, nullptr, &sinfo, &flags);
if (r < 0)
{
if (errno == EAGAIN || errno == EWOULDBLOCK)
break;
deliver_error(std::error_code(errno, std::generic_category()));
return;
}
if (r == 0)
{
deliver_error(asio::error::eof);
return;
}
if (flags & MSG_NOTIFICATION)
continue;
const std::size_t stream = sinfo.sinfo_stream;
if (stream >= n)
{
deliver_error(std::make_error_code(std::errc::protocol_error));
return;
}
bool too_big = false;
{
std::lock_guard<std::mutex> lock(mu);
if (closed)
return;
auto & part = partial[stream];
part.insert(part.end(), rbuf.begin(), rbuf.begin() + r);
if (part.size() > pol.max_frame)
too_big = true;
else if ((flags & MSG_EOR) != 0)
{
counters.read(part.size(), part.size(), 1);
auto & box = inbox[stream].bytes;
box.insert(box.end(), part.begin(), part.end());
part.clear();
satisfy_locked(stream);
}
}
if (too_big)
{
deliver_error(std::make_error_code(std::errc::message_size));
return;
}
}
arm_read();
}
void fail_waiters_locked(const std::error_code & ec)
{
for (auto & w : wait)
{
if (!w.active)
continue;
w.active = false;
detail::post_handler(io, std::move(w.h), ec);
}
}
void deliver_error(const std::error_code & ec)
{
{
std::lock_guard<std::mutex> lock(mu);
if (!fail)
fail = ec;
fail_waiters_locked(fail);
}
auto self = shared_from_this();
asio::post(strand, [self, ec] { self->fail_queued(ec); });
}
void close(const std::error_code & ec)
{
{
std::lock_guard<std::mutex> lock(mu);
if (!fail)
fail = ec;
if (closed && aborted)
return;
closed = true;
aborted = true;
fail_waiters_locked(ec);
}
auto self = shared_from_this();
asio::post(strand, [self, ec] {
std::error_code e;
self->sd.close(e);
self->fail_queued(ec);
});
}
void close_graceful()
{
{
std::lock_guard<std::mutex> lock(mu);
if (closed)
return;
closed = true;
fail_waiters_locked(asio::error::operation_aborted);
}
auto self = shared_from_this();
asio::post(strand, [self] {
self->draining = true;
if (!self->writing)
{
std::error_code e;
self->sd.close(e);
}
});
}
stream_stats stats() const
{
stream_stats s;
counters.fill(s);
s.buffered = buffered.load(std::memory_order_relaxed);
std::lock_guard<std::mutex> lock(mu);
for (const auto & b : inbox)
s.unread += b.avail();
s.error = fail;
s.closed = closed;
return s;
}
// --- writes (strand) ----------------------------------------------
void write(std::size_t i, std::shared_ptr<std::vector<std::uint8_t>> payload,
async_handler h)
{
check(i);
const std::size_t len = payload ? payload->size() : 0;
{
std::lock_guard<std::mutex> lock(mu);
if (fail || closed)
{
detail::post_handler(io, std::move(h),
fail ? fail : asio::error::operation_aborted);
return;
}
}
if (len == 0)
{
detail::post_handler(io, std::move(h), {});
return;
}
const auto pieces = detail::split_chunks(len, pol.chunk_bytes);
auto op = std::make_shared<detail::write_op>();
op->h = std::move(h);
op->left = pieces.size();
std::vector<chunk> cs;
cs.reserve(pieces.size());
for (const auto & pc : pieces)
cs.push_back(chunk{payload, pc.first, pc.second, op});
buffered.fetch_add(len, std::memory_order_relaxed);
auto self = shared_from_this();
asio::post(strand, [self, i, cs = std::move(cs)]() mutable {
std::error_code ec;
{
std::lock_guard<std::mutex> lock(self->mu);
ec = self->fail;
}
if (ec)
{
for (auto & c : cs)
self->finish_chunk(c, ec);
return;
}
for (auto & c : cs)
self->outq[i].push_back(std::move(c));
if (!self->writing)
self->write_next();
});
}
void finish_chunk(chunk & c, const std::error_code & ec)
{
buffered.fetch_sub(std::min(buffered.load(std::memory_order_relaxed),
c.len),
std::memory_order_relaxed);
auto & op = *c.op;
if (ec && !op.ec)
op.ec = ec;
if (--op.left == 0)
detail::post_handler(io, std::move(op.h), op.ec);
}
bool pick(std::size_t & lane)
{
for (std::size_t k = 0; k < n; ++k)
{
const std::size_t l = (rr + k) % n;
if (!outq[l].empty())
{
lane = l;
return true;
}
}
return false;
}
void write_next()
{
std::size_t lane = 0;
if (!pick(lane))
{
writing = false;
if (draining)
{
std::error_code e;
sd.close(e);
}
return;
}
writing = true;
auto & front = outq[lane].front();
const ssize_t r = ::sctp_sendmsg(sd.native_handle(),
front.payload->data() + front.off, front.len, nullptr, 0, 0, 0,
static_cast<std::uint16_t>(lane), 0, 0);
if (r < 0)
{
if (errno == EAGAIN || errno == EWOULDBLOCK)
{
auto self = shared_from_this();
sd.async_wait(asio::posix::stream_descriptor::wait_write,
asio::bind_executor(strand,
[self](const std::error_code & ec) {
if (ec)
{
self->writing = false;
self->deliver_error(ec);
return;
}
self->write_next();
}));
return;
}
writing = false;
deliver_error(std::error_code(errno, std::generic_category()));
return;
}
const std::size_t sent = static_cast<std::size_t>(r);
counters.wrote(sent, sent, 1);
if (sent < front.len)
{
buffered.fetch_sub(std::min(buffered.load(), sent),
std::memory_order_relaxed);
front.off += sent;
front.len -= sent;
write_next();
return;
}
chunk done = std::move(front);
outq[lane].pop_front();
rr = (lane + 1) % n;
finish_chunk(done, {});
auto self = shared_from_this();
asio::post(strand, [self] { self->write_next(); });
}
void fail_queued(const std::error_code & ec)
{
writing = false;
for (auto & q : outq)
{
for (auto & c : q)
finish_chunk(c, ec);
q.clear();
}
}
asio::io_context & io;
asio::posix::stream_descriptor sd;
asio::strand<asio::io_context::executor_type> strand;
std::size_t n = 0;
wire_policy pol;
mutable std::mutex mu;
std::vector<detail::stream_inbox> inbox;
std::vector<std::vector<std::uint8_t>> partial;
std::vector<detail::stream_waiter> wait;
std::error_code fail;
bool closed = false;
bool aborted = false;
std::vector<std::deque<chunk>> outq;
std::size_t rr = 0;
bool writing = false;
bool draining = false;
std::vector<std::uint8_t> rbuf;
std::atomic<std::size_t> buffered{0};
std::atomic<std::size_t> window;
detail::io_counters counters;
};
std::shared_ptr<impl> impl_;
};
// ---------------------------------------------------------------------------
// Association setup (deadline-bounded)
// ---------------------------------------------------------------------------
/// @brief Listening one-to-one SCTP socket.
class sctp_listener
{
public:
sctp_listener(unsigned short port, std::size_t nstreams,
const socket_options & o = {})
: nstreams_(nstreams), opts_(o)
{
if (nstreams == 0)
throw std::invalid_argument("sctp_listener: no streams");
fd_ = ::socket(AF_INET, SOCK_STREAM, IPPROTO_SCTP);
if (fd_ < 0)
throw std::system_error(errno, std::generic_category(),
"sctp_listener: socket");
int one = 1;
(void)::setsockopt(fd_, SOL_SOCKET, SO_REUSEADDR, &one, sizeof(one));
detail::sctp_configure(fd_, nstreams, o);
struct sockaddr_in addr;
std::memset(&addr, 0, sizeof(addr));
addr.sin_family = AF_INET;
addr.sin_addr.s_addr = htonl(INADDR_ANY);
addr.sin_port = htons(port);
if (::bind(fd_, reinterpret_cast<struct sockaddr *>(&addr), sizeof(addr)) < 0
|| ::listen(fd_, 16) < 0)
{
const int e = errno;
::close(fd_);
fd_ = -1;
throw std::system_error(e, std::generic_category(),
"sctp_listener: bind/listen on port " + std::to_string(port));
}
socklen_t alen = sizeof(addr);
if (::getsockname(fd_, reinterpret_cast<struct sockaddr *>(&addr), &alen) == 0)
port_ = ntohs(addr.sin_port);
}
sctp_listener(const sctp_listener &) = delete;
sctp_listener & operator=(const sctp_listener &) = delete;
~sctp_listener()
{
if (fd_ >= 0)
::close(fd_);
}
unsigned short port() const noexcept { return port_; }
int native_handle() const noexcept { return fd_; }
/// @brief Accept one association within `budget`; returns the fd.
int accept(std::chrono::milliseconds budget)
{
const auto deadline = setup_clock::now() + budget;
detail::wait_fd(fd_, POLLIN, deadline,
"sctp accept on port " + std::to_string(port_));
const int cfd = ::accept(fd_, nullptr, nullptr);
if (cfd < 0)
throw std::system_error(errno, std::generic_category(), "sctp accept");
detail::sctp_configure(cfd, nstreams_, opts_);
return cfd;
}
private:
int fd_ = -1;
unsigned short port_ = 0;
std::size_t nstreams_ = 0;
socket_options opts_{};
};
/// @brief Connect one SCTP association, retrying refusals until `budget`.
HEDLEY_WARN_UNUSED_RESULT
inline int connect_sctp(const std::string & host, unsigned short port,
std::size_t nstreams, std::chrono::milliseconds budget,
const socket_options & o = {})
{
const auto deadline = setup_clock::now() + budget;
struct sockaddr_in addr;
std::memset(&addr, 0, sizeof(addr));
addr.sin_family = AF_INET;
addr.sin_port = htons(port);
const std::string h = host == "localhost" ? "127.0.0.1" : host;
if (::inet_pton(AF_INET, h.c_str(), &addr.sin_addr) != 1)
throw std::invalid_argument("connect_sctp: bad host " + host);
int last = ETIMEDOUT;
for (;;)
{
const int fd = ::socket(AF_INET, SOCK_STREAM, IPPROTO_SCTP);
if (fd < 0)
throw std::system_error(errno, std::generic_category(),
"connect_sctp: socket");
detail::sctp_configure(fd, nstreams, o);
if (::connect(fd, reinterpret_cast<struct sockaddr *>(&addr),
sizeof(addr)) == 0)
return fd;
last = errno;
::close(fd);
if (setup_clock::now() >= deadline)
throw std::system_error(last, std::generic_category(),
"connect_sctp " + host + ":" + std::to_string(port));
std::this_thread::sleep_for(std::chrono::milliseconds(20));
}
}
/// @brief Accept one association on an ephemeral (or given) port.
HEDLEY_WARN_UNUSED_RESULT
inline int accept_sctp_association(asio::io_context &,
std::atomic<unsigned short> & port, std::size_t nstreams,
std::chrono::milliseconds budget = deadlines{}.accept)
{
sctp_listener lst(port.load(), nstreams);
port.store(lst.port());
return lst.accept(budget);
}
/// @brief Connect one association to `host:port`.
HEDLEY_WARN_UNUSED_RESULT
inline int connect_sctp_association(asio::io_context &, const std::string & host,
std::atomic<unsigned short> & port, std::size_t nstreams,
std::chrono::milliseconds budget = deadlines{}.connect)
{
return connect_sctp(host, port.load(), nstreams, budget);
}
#else // !DPF_HAS_LIBSCTP
class async_sctp_stream_array final : public async_stream_array
{
public:
async_sctp_stream_array(asio::io_context &, int, std::size_t,
const wire_policy & = {})
{
throw std::logic_error(
"async_sctp_stream_array: real SCTP requires Linux + libsctp "
"(<netinet/sctp.h>, link -lsctp)");
}
std::size_t size() const noexcept override { return 0; }
asio::io_context & context() noexcept override
{
std::terminate();
}
void async_write(std::size_t, const void *, std::size_t,
async_handler) override
{
throw std::logic_error("async_sctp_stream_array: not available");
}
void async_read(std::size_t, void *, std::size_t, async_handler) override
{
throw std::logic_error("async_sctp_stream_array: not available");
}
};
class sctp_listener
{
public:
sctp_listener(unsigned short, std::size_t, const socket_options & = {})
{
throw std::logic_error("sctp_listener: SCTP requires Linux + libsctp");
}
unsigned short port() const noexcept { return 0; }
int native_handle() const noexcept { return -1; }
int accept(std::chrono::milliseconds) { return -1; }
};
HEDLEY_WARN_UNUSED_RESULT
inline int connect_sctp(const std::string &, unsigned short, std::size_t,
std::chrono::milliseconds, const socket_options & = {})
{
throw std::logic_error("connect_sctp: SCTP requires Linux + libsctp");
}
HEDLEY_WARN_UNUSED_RESULT
inline int accept_sctp_association(asio::io_context &,
std::atomic<unsigned short> &, std::size_t,
std::chrono::milliseconds = deadlines{}.accept)
{
throw std::logic_error(
"accept_sctp_association: SCTP requires Linux + libsctp");
}
HEDLEY_WARN_UNUSED_RESULT
inline int connect_sctp_association(asio::io_context &, const std::string &,
std::atomic<unsigned short> &, std::size_t,
std::chrono::milliseconds = deadlines{}.connect)
{
throw std::logic_error(
"connect_sctp_association: SCTP requires Linux + libsctp");
}
#endif // DPF_HAS_LIBSCTP
} // namespace net
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_NET_ASYNC_SCTP_STREAM_ARRAY_HPP__