libdpf/include/dpf/net/party_session.hpp

1172 lines
42 KiB
C++
Raw Normal View History

/// @file dpf/net/party_session.hpp
/// @brief One role's TCP/SCTP session: join, per-edge transport, dealer, reconnect.
/// @details Socket edges run TLS 1.3 unless `session_options::security.encrypt`
/// is off (see `dpf/net/security.hpp`): each side presents its
/// identity key and authenticates the peer only if it holds that
/// peer's key, so authentication can run in one direction, both, or
/// neither; every unauthenticated direction is logged. After the TLS
/// handshake (or directly, with encryption off) both ends exchange a
/// fixed record (party id, epoch, transport, lane count, lane index,
/// and whether the sender authenticated the receiver) under a
/// deadline, so two ends that disagree fail at connect time with a
/// message instead of corrupting frames later. Lower ids accept,
/// higher ids connect; connects retry until the peer listens, so
/// processes may start in any order. `join(table)` takes a static
/// `host:port` table (separate processes); `join(host, ports)` is the
/// in-process rendezvous. `reconnect(peer)` replaces one edge and
/// returns the new link; `reconnector(peer)` adapts it for
/// `sink_options::reconnect`. The dealer link has its own policy,
/// `io_context`, epoch, and reconnect. SCTP edges cannot be encrypted
/// and are refused while encryption is on.
#ifndef LIBDPF_INCLUDE_DPF_NET_PARTY_SESSION_HPP__
#define LIBDPF_INCLUDE_DPF_NET_PARTY_SESSION_HPP__
#include <algorithm>
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <functional>
#include <map>
#include <memory>
#include <stdexcept>
#include <string>
#include <system_error>
#include <thread>
#include <utility>
#include <vector>
#include <poll.h>
#include "dpf/net/asio_ns.hpp"
#include "dpf/log.hpp"
#include "dpf/net/async_sctp_stream_array.hpp"
#include "dpf/net/async_stream_array.hpp"
#include "dpf/net/connect.hpp"
#include "dpf/net/identity.hpp"
#include "dpf/net/link_log.hpp"
#include "dpf/net/policy.hpp"
#include "dpf/net/security.hpp"
#include "dpf/net/mesh_rendezvous.hpp"
#include "dpf/net/socket_tune.hpp"
#include "dpf/net/tls.hpp"
namespace dpf
{
namespace net
{
/// @brief Where one party listens.
struct peer_address
{
std::string host = "127.0.0.1";
unsigned short port = 0;
};
/// @brief Session-wide defaults. Per-edge overrides live on `party_session`.
struct session_options
{
std::size_t n_lanes = 1;
transport kind = transport::mux;
wire_policy policy{};
deadlines limits{};
peer_security security{};
};
namespace detail
{
struct session_hello
{
static constexpr std::uint32_t magic = 0x53535044u; // 'DPSS'
static constexpr std::size_t size = 28;
/// The sender authenticated the receiver's key.
static constexpr std::uint32_t verified_you = 1;
std::uint32_t id = 0;
std::uint32_t epoch = 0;
std::uint32_t kind = 0;
std::uint32_t lanes = 0;
std::uint32_t lane = 0;
std::uint32_t flags = 0;
void pack(std::uint8_t * p) const
{
put_u32(p, magic);
put_u32(p + 4, id);
put_u32(p + 8, epoch);
put_u32(p + 12, kind);
put_u32(p + 16, lanes);
put_u32(p + 20, lane);
put_u32(p + 24, flags);
}
static session_hello unpack(const std::uint8_t * p)
{
if (p[0] == 0x16 && p[1] == 0x03)
throw std::runtime_error("party_session: the peer started a TLS handshake "
"but this side has encryption off (set encryption the same at both "
"ends)");
if (get_u32(p) != magic)
throw std::runtime_error(
"party_session: peer is not a party_session (bad handshake magic)");
session_hello h;
h.id = get_u32(p + 4);
h.epoch = get_u32(p + 8);
h.kind = get_u32(p + 12);
h.lanes = get_u32(p + 16);
h.lane = get_u32(p + 20);
h.flags = get_u32(p + 24);
return h;
}
};
inline std::string id_name(std::uint32_t id)
{
return id == 0xfffffffeu ? std::string("dealer") : "party " + std::to_string(id);
}
inline std::string role_name(std::uint32_t id)
{
return id == 0xfffffffeu ? std::string("dealer") : "p" + std::to_string(id);
}
/// @brief One TCP connection on its way into a link: plain, or TLS once the
/// handshake is done.
struct peer_conn
{
explicit peer_conn(asio::ip::tcp::socket s)
: plain(std::move(s)), fd(plain.native_handle())
{
}
asio::ip::tcp::socket plain;
#if DPF_HAS_OPENSSL
std::unique_ptr<tls_stream> tls;
#endif
int fd = -1;
link_security sec;
bool encrypted() const noexcept
{
#if DPF_HAS_OPENSSL
return static_cast<bool>(tls);
#else
return false;
#endif
}
};
inline std::chrono::milliseconds left(setup_clock::time_point deadline)
{
return std::chrono::milliseconds(std::max(1, remaining_ms(deadline)));
}
#if DPF_HAS_OPENSSL
using peer_tls_context = tls_context;
#else
struct peer_tls_context
{
};
#endif
/// @brief Tune `c` (TCP_NODELAY first: the handshake is several round trips),
/// then run the TLS handshake when `ctx` is set.
inline void secure(peer_conn & c, asio::io_context & io, peer_tls_context * ctx,
bool server, const socket_options & so, std::chrono::milliseconds budget,
const std::string & who)
{
tune_tcp(c.plain, so);
if (ctx == nullptr)
return;
#if DPF_HAS_OPENSSL
c.tls = std::make_unique<tls_stream>(std::move(c.plain), *ctx);
try
{
tls_handshake(io, *c.tls, server, budget, who);
}
catch (const std::system_error & e)
{
throw std::runtime_error(std::string(e.what())
+ " (if the peer has encryption off, set it the same at both ends)");
}
c.sec = tls_describe(*c.tls);
#else
(void)io;
(void)server;
(void)budget;
(void)who;
#endif
}
inline void send_hello(peer_conn & c, asio::io_context & io, const session_hello & h,
setup_clock::time_point deadline, const std::string & what)
{
std::uint8_t out[session_hello::size];
h.pack(out);
#if DPF_HAS_OPENSSL
if (c.tls)
{
tls_write(io, *c.tls, out, sizeof(out), left(deadline), what);
return;
}
#endif
(void)io;
send_all_until(c.plain.native_handle(), out, sizeof(out), deadline, what);
}
inline session_hello recv_hello(peer_conn & c, asio::io_context & io,
setup_clock::time_point deadline, const std::string & what)
{
std::uint8_t in[session_hello::size];
#if DPF_HAS_OPENSSL
if (c.tls)
{
tls_read(io, *c.tls, in, sizeof(in), left(deadline), what);
return session_hello::unpack(in);
}
#endif
(void)io;
recv_all_until(c.plain.native_handle(), in, sizeof(in), deadline, what);
return session_hello::unpack(in);
}
/// @brief Authenticate `party` on an encrypted connection (no-op when plain).
inline void check_party(peer_conn & c, const peer_security & policy,
std::uint32_t party, const std::string & who)
{
#if DPF_HAS_OPENSSL
if (c.encrypted())
{
try
{
check_peer(c.sec, policy, party, who);
}
catch (...)
{
std::error_code e;
c.tls->lowest_layer().close(e);
throw;
}
}
#else
(void)c;
(void)policy;
(void)party;
(void)who;
#endif
}
inline std::unique_ptr<async_stream_array> mux_link(asio::io_context & io,
peer_conn & c, unsigned self, unsigned peer, std::size_t lanes,
const wire_policy & pol)
{
#if DPF_HAS_OPENSSL
if (c.tls)
return std::make_unique<async_tls_mux_stream_array>(io, std::move(*c.tls), self,
peer, lanes, pol);
#endif
return std::make_unique<async_mux_stream_array>(io, std::move(c.plain), self, peer,
lanes, pol);
}
inline std::unique_ptr<async_stream_array> parallel_link(asio::io_context & io,
std::vector<std::unique_ptr<peer_conn>> & cs, const wire_policy & pol)
{
#if DPF_HAS_OPENSSL
if (cs.front()->tls)
{
std::vector<tls_stream> v;
v.reserve(cs.size());
for (auto & c : cs)
v.push_back(std::move(*c->tls));
return std::make_unique<async_tls_parallel_stream_array>(io, std::move(v), pol);
}
#endif
std::vector<asio::ip::tcp::socket> v;
v.reserve(cs.size());
for (auto & c : cs)
v.push_back(std::move(c->plain));
return std::make_unique<async_parallel_stream_array>(io, std::move(v), pol);
}
/// @brief This process's key for party links, or a fresh one (logged).
inline std::shared_ptr<const identity> link_identity(const peer_security & sec,
const std::string & who)
{
if (sec.self)
{
DPF_LOG(info, "security.identity").kv("who", who)
.kv("key", sec.self->key().base64()).kv("ephemeral", false)
.kv("trusted", sec.trusted.size());
return sec.self;
}
auto id = std::make_shared<identity>(identity::generate());
DPF_LOG(info, "security.identity").kv("who", who).kv("key", id->key().base64())
.kv("ephemeral", true).kv("trusted", sec.trusted.size());
if (log::first_time("security.no_identity." + log::role() + "." + who))
DPF_LOG(warning, "security.no_identity").kv("who", who)
.kv("detail", "no identity key configured: links are encrypted to a fresh "
"key for this run, so peers cannot authenticate " + who);
return id;
}
/// @brief Log a link direction this side could not authenticate.
inline void note_link(const std::string & peer_role, const link_security & sec)
{
if (!sec.encrypted)
return;
const std::string me = log::role().empty() ? std::string("this party") : log::role();
if (sec.peer_auth == "none"
&& log::first_time("security.unauthenticated." + me + "." + peer_role))
DPF_LOG(warning, "security.unauthenticated").kv("peer", peer_role)
.kv("peer_key", sec.peer_key ? sec.peer_key->base64() : std::string("none"))
.kv("detail", "no key configured for " + peer_role + ": the link is "
"encrypted but " + peer_role + " is not authenticated");
if (!sec.peer_verified_us
&& log::first_time("security.not_authenticated_by." + me + "." + peer_role))
DPF_LOG(info, "security.not_authenticated_by_peer").kv("peer", peer_role)
.kv("detail", peer_role + " holds no key for " + me);
}
} // namespace detail
class party_session
{
public:
static constexpr unsigned k_dealer_id = dealer_id;
party_session(asio::io_context & io, unsigned me, unsigned parties,
session_options opt)
: io_(&io),
dealer_io_(&io),
me_(me),
parties_(parties),
opt_(std::move(opt))
{
if (parties_ < 2 || me_ >= parties_)
throw std::invalid_argument("party_session: party/n");
if (opt_.n_lanes == 0)
throw std::invalid_argument("party_session: n_lanes 0");
opt_.policy.validate();
links_.resize(parties_);
sec_.resize(parties_);
table_.resize(parties_);
epoch_.assign(parties_, 0);
kind_.assign(parties_, opt_.kind);
policy_.assign(parties_, opt_.policy);
dealer_policy_ = opt_.policy;
if (opt_.security.encrypt)
{
refuse_sctp(opt_.kind);
#if DPF_HAS_OPENSSL
self_ = detail::link_identity(opt_.security, "p" + std::to_string(me_));
tls_ = make_peer_tls_context(*self_);
#else
throw std::logic_error("party_session: built without OpenSSL; set "
"encryption=off for plaintext links");
#endif
}
}
party_session(asio::io_context & io, unsigned me, unsigned parties,
std::size_t nstreams)
: party_session(io, me, parties, make_opts(nstreams))
{
}
unsigned me() const noexcept { return me_; }
unsigned parties() const noexcept { return parties_; }
std::size_t streams() const noexcept { return opt_.n_lanes; }
asio::io_context & context() noexcept { return *io_; }
const session_options & options() const noexcept { return opt_; }
/// @brief This party's link key (null with encryption off).
std::shared_ptr<const identity> identity_key() const noexcept { return self_; }
// --- configuration (before join, both ends must agree) ---------------
void set_edge_transport(unsigned peer, transport t)
{
check_peer(peer);
if (t != transport::mux && t != transport::parallel && t != transport::sctp)
throw std::invalid_argument(std::string("party_session: transport ")
+ transport_name(t) + " is not a socket edge");
if (t == transport::sctp && !sctp_available())
throw std::logic_error("party_session: SCTP needs Linux + libsctp");
if (opt_.security.encrypt)
refuse_sctp(t);
kind_[peer] = t;
}
transport edge_transport(unsigned peer) const
{
check_peer(peer);
return kind_[peer];
}
void set_edge_policy(unsigned peer, const wire_policy & pol)
{
check_peer(peer);
pol.validate();
policy_[peer] = pol;
if (links_[peer])
links_[peer]->set_window_bytes(pol.window_bytes);
}
/// @brief Window for every edge and the dealer link, now and later.
void set_window_bytes(std::size_t bytes)
{
for (auto & p : policy_)
p.window_bytes = bytes;
dealer_policy_.window_bytes = bytes;
for (auto & link : links_)
if (link)
link->set_window_bytes(bytes);
if (dealer_)
dealer_->set_window_bytes(bytes);
}
void set_dealer_policy(const wire_policy & pol)
{
pol.validate();
dealer_policy_ = pol;
if (dealer_)
dealer_->set_window_bytes(pol.window_bytes);
}
/// @brief Run the dealer link on its own `io_context` (own failure domain).
void set_dealer_context(asio::io_context & io) { dealer_io_ = &io; }
void set_deadlines(const deadlines & d) { opt_.limits = d; }
// --- listen / join ---------------------------------------------------
/// @brief Bind TCP (and SCTP if any edge uses it) on `port`; 0 = ephemeral.
unsigned short listen(unsigned short port = 0)
{
if (!acceptor_)
{
acceptor_ = std::make_unique<asio::ip::tcp::acceptor>(*io_);
open_listener(*acceptor_, port);
published_ = acceptor_->local_endpoint().port();
log_listen(published_, wants_sctp(), opt_.security.encrypt);
}
if (!sctp_ && wants_sctp())
sctp_ = std::make_unique<sctp_listener>(published_, opt_.n_lanes,
opt_.policy.socket);
return published_;
}
unsigned short port() const noexcept { return published_; }
/// @brief Join with a static table: `table[i]` is party `i`'s address.
void join(const std::vector<peer_address> & table)
{
if (table.size() != parties_)
throw std::invalid_argument("party_session::join: table has "
+ std::to_string(table.size()) + " entries, expected "
+ std::to_string(parties_));
table_ = table;
listen(table_[me_].port);
accept_higher();
for (unsigned peer = 0; peer < me_; ++peer)
connect_edge(peer);
}
/// @brief In-process rendezvous through shared `ports` (tests, one binary).
void join(const std::string & host, mesh_ports & ports, std::size_t n_lanes = 0)
{
if (ports.size() != parties_)
throw std::invalid_argument("party_session::join port table");
if (n_lanes != 0)
opt_.n_lanes = n_lanes;
ports[me_].store(listen(ports[me_].load()));
const auto deadline = setup_clock::now() + opt_.limits.join;
std::vector<peer_address> table(parties_);
for (unsigned i = 0; i < parties_; ++i)
{
while (ports[i].load() == 0)
{
if (setup_clock::now() >= deadline)
throw std::runtime_error("party_session: join timed out after "
+ std::to_string(opt_.limits.join.count())
+ " ms waiting for party " + std::to_string(i));
std::this_thread::sleep_for(std::chrono::milliseconds(1));
}
table[i] = peer_address{host, ports[i].load()};
}
join(table);
}
// --- edges -----------------------------------------------------------
async_stream_array & peer(unsigned p)
{
check_peer(p);
if (!links_[p])
throw std::logic_error("party_session: no link to party "
+ std::to_string(p) + " (join first)");
return *links_[p];
}
stream_stats edge_stats(unsigned p) const
{
check_peer(p);
return links_[p] ? links_[p]->stats() : stream_stats{};
}
/// @brief How the edge to `p` was secured (encryption, cipher, who
/// authenticated whom).
const link_security & edge_security(unsigned p) const
{
check_peer(p);
return sec_[p];
}
std::uint32_t epoch(unsigned p) const
{
check_peer(p);
return epoch_[p];
}
/// @brief Replace the edge to `peer` and return it. Both ends must call.
async_stream_array & reconnect(unsigned peer)
{
check_peer(peer);
if (table_[peer].port == 0 && me_ > peer)
throw std::logic_error("party_session::reconnect before join");
links_[peer].reset();
++epoch_[peer];
if (me_ < peer)
accept_one_edge(peer);
else
connect_edge(peer);
return *links_[peer];
}
/// @brief `sink_options::reconnect` for the edge to `peer`.
std::function<async_stream_array *(const std::error_code &)> reconnector(
unsigned peer)
{
return [this, peer](const std::error_code &) { return &reconnect(peer); };
}
// --- dealer ----------------------------------------------------------
/// @brief Party side: connect to the dealer at `host:port`.
void connect_dealer(const std::string & host, unsigned short port)
{
dealer_addr_ = peer_address{host, port};
const std::string who = "dealer at " + host + ":" + std::to_string(port);
asio::ip::tcp::socket sock(*dealer_io_);
connect_until(sock, host, port, opt_.limits.connect);
detail::peer_conn c(std::move(sock));
const auto deadline = setup_clock::now() + opt_.limits.handshake;
detail::secure(c, *dealer_io_, tls_ptr(), false, dealer_policy_.socket,
opt_.limits.handshake, "handshake with " + who);
detail::check_party(c, opt_.security, k_dealer_id, who);
auto mine = dealer_hello(me_);
mine.flags = c.sec.peer_auth == "key" ? detail::session_hello::verified_you : 0;
detail::send_hello(c, *dealer_io_, mine, deadline, "dealer handshake");
const auto got = detail::recv_hello(c, *dealer_io_, deadline, "dealer handshake");
if (got.id != k_dealer_id)
throw std::runtime_error("party_session: " + host + ":"
+ std::to_string(port) + " is " + detail::id_name(got.id)
+ ", not the dealer");
check_match(got, mine, "dealer");
c.sec.peer_verified_us = (got.flags & detail::session_hello::verified_you) != 0;
dealer_sec_ = c.sec;
dealer_ = detail::mux_link(*dealer_io_, c, me_, k_dealer_id, opt_.n_lanes,
dealer_policy_);
log_link_up("connect", "dealer", transport::mux, opt_.n_lanes, 0,
dealer_epoch_, c.fd, dealer_policy_.socket, &dealer_sec_);
detail::note_link("dealer", dealer_sec_);
}
/// @brief Dealer side of a one-party dealer link (see `dealer_session`
/// for a dealer serving several parties).
void accept_dealer()
{
listen(published_);
asio::ip::tcp::socket sock(*dealer_io_);
accept_until(*acceptor_, sock, opt_.limits.accept);
detail::peer_conn c(std::move(sock));
const auto deadline = setup_clock::now() + opt_.limits.handshake;
detail::secure(c, *dealer_io_, tls_ptr(), true, dealer_policy_.socket,
opt_.limits.handshake, "dealer handshake");
const auto got = detail::recv_hello(c, *dealer_io_, deadline, "dealer handshake");
detail::check_party(c, opt_.security, got.id, detail::id_name(got.id));
auto mine = dealer_hello(k_dealer_id);
mine.flags = c.sec.peer_auth == "key" ? detail::session_hello::verified_you : 0;
detail::send_hello(c, *dealer_io_, mine, deadline, "dealer handshake");
check_match(got, mine, "dealer");
c.sec.peer_verified_us = (got.flags & detail::session_hello::verified_you) != 0;
dealer_is_server_ = true;
dealer_sec_ = c.sec;
dealer_ = detail::mux_link(*dealer_io_, c, k_dealer_id, got.id, opt_.n_lanes,
dealer_policy_);
log_link_up("accept", detail::role_name(got.id), transport::mux, opt_.n_lanes,
0, dealer_epoch_, c.fd, dealer_policy_.socket, &dealer_sec_);
detail::note_link(detail::role_name(got.id), dealer_sec_);
}
/// @brief Replace the dealer link (both ends must call).
async_stream_array & reconnect_dealer()
{
dealer_.reset();
++dealer_epoch_;
if (dealer_is_server_)
accept_dealer();
else
connect_dealer(dealer_addr_.host, dealer_addr_.port);
return *dealer_;
}
std::function<async_stream_array *(const std::error_code &)>
dealer_reconnector()
{
return [this](const std::error_code &) { return &reconnect_dealer(); };
}
bool has_dealer() const noexcept { return static_cast<bool>(dealer_); }
async_stream_array & dealer()
{
if (!dealer_)
throw std::logic_error("party_session: no dealer link");
return *dealer_;
}
stream_stats dealer_stats() const
{
return dealer_ ? dealer_->stats() : stream_stats{};
}
const link_security & dealer_security() const noexcept { return dealer_sec_; }
private:
static session_options make_opts(std::size_t nstreams)
{
session_options o;
o.n_lanes = nstreams;
return o;
}
static void refuse_sctp(transport t)
{
if (t == transport::sctp)
throw std::invalid_argument("party_session: SCTP links cannot be "
"encrypted (OpenSSL has no TLS over SCTP streams); use mux or "
"parallel, or set encryption=off");
}
detail::peer_tls_context * tls_ptr() const noexcept
{
#if DPF_HAS_OPENSSL
return tls_.get();
#else
return nullptr;
#endif
}
detail::session_hello dealer_hello(std::uint32_t id) const
{
detail::session_hello h;
h.id = id;
h.epoch = dealer_epoch_;
h.kind = static_cast<std::uint32_t>(transport::mux);
h.lanes = static_cast<std::uint32_t>(opt_.n_lanes);
return h;
}
void check_peer(unsigned p) const
{
if (p == me_ || p >= parties_)
throw std::invalid_argument("party_session: bad peer "
+ std::to_string(p));
}
bool wants_sctp() const
{
for (unsigned p = 0; p < parties_; ++p)
if (p != me_ && kind_[p] == transport::sctp)
return true;
return false;
}
std::size_t sockets_for(unsigned peer) const
{
return kind_[peer] == transport::parallel ? opt_.n_lanes : 1;
}
detail::session_hello my_hello(unsigned peer, std::uint32_t lane) const
{
detail::session_hello h;
h.id = me_;
h.epoch = epoch_[peer];
h.kind = static_cast<std::uint32_t>(kind_[peer]);
h.lanes = static_cast<std::uint32_t>(opt_.n_lanes);
h.lane = lane;
return h;
}
void check_match(const detail::session_hello & got,
const detail::session_hello & mine, const std::string & who) const
{
std::string why;
if (got.kind != mine.kind)
why += std::string(" transport ")
+ transport_name(static_cast<transport>(got.kind)) + " vs "
+ transport_name(static_cast<transport>(mine.kind)) + ";";
if (got.lanes != mine.lanes)
why += " lanes " + std::to_string(got.lanes) + " vs "
+ std::to_string(mine.lanes) + ";";
if (got.epoch != mine.epoch)
why += " epoch " + std::to_string(got.epoch) + " vs "
+ std::to_string(mine.epoch) + ";";
if (!why.empty())
throw std::runtime_error("party_session: " + who
+ " disagrees (peer vs this side):" + why);
}
struct parallel_in_progress
{
std::vector<std::unique_ptr<detail::peer_conn>> conns;
std::size_t got = 0;
};
/// @brief Accept every edge from higher-numbered parties, in any order,
/// on the TCP and SCTP listeners at once.
void accept_higher()
{
std::size_t tcp_left = 0;
std::size_t sctp_left = 0;
for (unsigned p = me_ + 1; p < parties_; ++p)
{
if (kind_[p] == transport::sctp)
++sctp_left;
else
tcp_left += sockets_for(p);
}
std::map<unsigned, parallel_in_progress> partial;
const auto deadline = setup_clock::now() + opt_.limits.accept;
while (tcp_left + sctp_left > 0)
{
pollfd pfds[2]{};
int nfds = 0;
if (tcp_left > 0)
{
pfds[nfds].fd = acceptor_->native_handle();
pfds[nfds].events = POLLIN;
++nfds;
}
if (sctp_left > 0)
{
pfds[nfds].fd = sctp_fd();
pfds[nfds].events = POLLIN;
++nfds;
}
const int ms = detail::remaining_ms(deadline);
if (ms == 0)
throw std::system_error(std::make_error_code(std::errc::timed_out),
"party_session " + std::to_string(me_) + ": accept timed out "
"after " + std::to_string(opt_.limits.accept.count())
+ " ms (" + std::to_string(tcp_left + sctp_left)
+ " connections missing)");
const int rc = ::poll(pfds, static_cast<nfds_t>(nfds), ms);
if (rc <= 0)
continue;
for (int k = 0; k < nfds; ++k)
{
if ((pfds[k].revents & POLLIN) == 0)
continue;
if (tcp_left > 0 && pfds[k].fd == acceptor_->native_handle())
{
asio::ip::tcp::socket sock(*io_);
acceptor_->accept(sock);
accept_tcp(std::move(sock), partial, false, 0);
--tcp_left;
}
else if (sctp_left > 0)
{
accept_sctp(0, false);
--sctp_left;
}
}
}
}
int sctp_fd() const { return sctp_ ? sctp_->native_handle() : -1; }
void accept_tcp(asio::ip::tcp::socket sock,
std::map<unsigned, parallel_in_progress> & partial, bool expect_one,
unsigned expect)
{
auto c = std::make_unique<detail::peer_conn>(std::move(sock));
const auto deadline = setup_clock::now() + opt_.limits.handshake;
detail::secure(*c, *io_, tls_ptr(), true, opt_.policy.socket,
opt_.limits.handshake, "party_session " + std::to_string(me_) + " handshake");
// The accepting side answers with the record for the announced peer.
const auto theirs = detail::recv_hello(*c, *io_, deadline,
"party_session handshake");
const unsigned p = theirs.id;
if (p <= me_ || p >= parties_ || (expect_one && p != expect))
throw std::runtime_error("party_session " + std::to_string(me_)
+ ": unexpected connection from " + detail::id_name(p));
detail::check_party(*c, opt_.security, p, "party " + std::to_string(p));
auto mine = my_hello(p, theirs.lane);
mine.flags = c->sec.peer_auth == "key" ? detail::session_hello::verified_you : 0;
detail::send_hello(*c, *io_, mine, deadline, "party_session handshake");
check_match(theirs, mine, "party " + std::to_string(p));
c->sec.peer_verified_us = (theirs.flags & detail::session_hello::verified_you) != 0;
const std::string role = "p" + std::to_string(p);
if (kind_[p] == transport::mux)
{
if (links_[p])
throw std::runtime_error("party_session: duplicate edge from party "
+ std::to_string(p));
sec_[p] = c->sec;
links_[p] = detail::mux_link(*io_, *c, me_, p, opt_.n_lanes, policy_[p]);
log_link_up("accept", role, kind_[p], opt_.n_lanes, 0, epoch_[p], c->fd,
policy_[p].socket, &sec_[p]);
detail::note_link(role, sec_[p]);
return;
}
auto & pp = partial[p];
if (pp.conns.empty())
pp.conns.resize(opt_.n_lanes);
if (theirs.lane >= opt_.n_lanes || pp.conns[theirs.lane])
throw std::runtime_error("party_session: bad parallel lane "
+ std::to_string(theirs.lane) + " from party " + std::to_string(p));
pp.conns[theirs.lane] = std::move(c);
if (++pp.got == opt_.n_lanes)
adopt_parallel(p, pp.conns, "accept");
if (pp.got == opt_.n_lanes)
partial.erase(p);
}
/// @brief Build a parallel link from one connection per lane.
void adopt_parallel(unsigned p, std::vector<std::unique_ptr<detail::peer_conn>> & cs,
const char * how)
{
std::vector<int> fds;
std::vector<link_security> secs;
for (auto & c : cs)
{
fds.push_back(c->fd);
secs.push_back(c->sec);
}
sec_[p] = secs.front();
for (const auto & s : secs)
sec_[p].peer_verified_us = sec_[p].peer_verified_us && s.peer_verified_us;
links_[p] = detail::parallel_link(*io_, cs, policy_[p]);
const std::string role = "p" + std::to_string(p);
for (std::size_t k = 0; k < fds.size(); ++k)
log_link_up(how, role, kind_[p], opt_.n_lanes, static_cast<std::uint32_t>(k),
epoch_[p], fds[k], policy_[p].socket, &secs[k]);
detail::note_link(role, sec_[p]);
}
void accept_sctp(unsigned expect, bool expect_one)
{
const int fd = sctp_->accept(opt_.limits.accept);
std::uint8_t in[detail::session_hello::size];
const auto deadline = setup_clock::now() + opt_.limits.handshake;
try
{
detail::recv_all_until(fd, in, sizeof(in), deadline,
"party_session sctp handshake");
const auto theirs = detail::session_hello::unpack(in);
const unsigned p = theirs.id;
if (p <= me_ || p >= parties_ || (expect_one && p != expect))
throw std::runtime_error("party_session: unexpected SCTP "
"association from " + detail::id_name(p));
const auto mine = my_hello(p, 0);
std::uint8_t out[detail::session_hello::size];
mine.pack(out);
detail::send_all_until(fd, out, sizeof(out), deadline,
"party_session sctp handshake");
check_match(theirs, mine, "party " + std::to_string(p));
sec_[p] = link_security{};
links_[p] = std::make_unique<async_sctp_stream_array>(*io_, fd,
opt_.n_lanes, policy_[p]);
log_link_up("accept", "p" + std::to_string(p), transport::sctp,
opt_.n_lanes, 0, epoch_[p], fd, policy_[p].socket);
}
catch (...)
{
::close(fd);
throw;
}
}
void accept_one_edge(unsigned peer)
{
listen(published_);
if (kind_[peer] == transport::sctp)
{
accept_sctp(peer, true);
return;
}
std::map<unsigned, parallel_in_progress> partial;
for (std::size_t k = 0; k < sockets_for(peer); ++k)
{
asio::ip::tcp::socket sock(*io_);
accept_until(*acceptor_, sock, opt_.limits.accept);
accept_tcp(std::move(sock), partial, true, peer);
}
}
void connect_edge(unsigned peer)
{
const auto & addr = table_[peer];
if (addr.port == 0)
throw std::logic_error("party_session: no address for party "
+ std::to_string(peer));
const std::string who = "party " + std::to_string(peer) + " at "
+ addr.host + ":" + std::to_string(addr.port);
if (kind_[peer] == transport::sctp)
{
const int fd = connect_sctp(addr.host, addr.port, opt_.n_lanes,
opt_.limits.connect, policy_[peer].socket);
try
{
const auto mine = my_hello(peer, 0);
std::uint8_t out[detail::session_hello::size];
std::uint8_t in[detail::session_hello::size];
mine.pack(out);
exchange_record(fd, out, in, detail::session_hello::size,
opt_.limits.handshake, "sctp handshake with " + who);
const auto got = detail::session_hello::unpack(in);
if (got.id != peer)
throw std::runtime_error("party_session: " + who + " is "
+ detail::id_name(got.id));
check_match(got, mine, who);
sec_[peer] = link_security{};
links_[peer] = std::make_unique<async_sctp_stream_array>(*io_, fd,
opt_.n_lanes, policy_[peer]);
log_link_up("connect", "p" + std::to_string(peer), transport::sctp,
opt_.n_lanes, 0, epoch_[peer], fd, policy_[peer].socket);
}
catch (...)
{
::close(fd);
throw;
}
return;
}
std::vector<std::unique_ptr<detail::peer_conn>> conns;
for (std::size_t k = 0; k < sockets_for(peer); ++k)
{
asio::ip::tcp::socket sock(*io_);
connect_until(sock, addr.host, addr.port, opt_.limits.connect);
auto c = std::make_unique<detail::peer_conn>(std::move(sock));
const auto deadline = setup_clock::now() + opt_.limits.handshake;
detail::secure(*c, *io_, tls_ptr(), false, policy_[peer].socket,
opt_.limits.handshake, "handshake with " + who);
detail::check_party(*c, opt_.security, peer, who);
auto mine = my_hello(peer, static_cast<std::uint32_t>(k));
mine.flags =
c->sec.peer_auth == "key" ? detail::session_hello::verified_you : 0;
detail::send_hello(*c, *io_, mine, deadline, "handshake with " + who);
const auto got = detail::recv_hello(*c, *io_, deadline,
"handshake with " + who);
if (got.id != peer)
throw std::runtime_error("party_session: " + who + " is "
+ detail::id_name(got.id));
check_match(got, mine, who);
c->sec.peer_verified_us = (got.flags & detail::session_hello::verified_you) != 0;
conns.push_back(std::move(c));
}
if (kind_[peer] == transport::mux)
{
auto & c = *conns.front();
sec_[peer] = c.sec;
links_[peer] = detail::mux_link(*io_, c, me_, peer, opt_.n_lanes,
policy_[peer]);
const std::string role = "p" + std::to_string(peer);
log_link_up("connect", role, kind_[peer], opt_.n_lanes, 0, epoch_[peer],
c.fd, policy_[peer].socket, &sec_[peer]);
detail::note_link(role, sec_[peer]);
}
else
adopt_parallel(peer, conns, "connect");
}
asio::io_context * io_ = nullptr;
asio::io_context * dealer_io_ = nullptr;
unsigned me_ = 0;
unsigned parties_ = 0;
session_options opt_;
std::shared_ptr<const identity> self_;
#if DPF_HAS_OPENSSL
std::shared_ptr<tls_context> tls_;
#endif
unsigned short published_ = 0;
std::vector<peer_address> table_;
std::vector<std::uint32_t> epoch_;
std::vector<transport> kind_;
std::vector<wire_policy> policy_;
std::unique_ptr<asio::ip::tcp::acceptor> acceptor_;
std::unique_ptr<sctp_listener> sctp_;
std::vector<link_security> sec_;
std::vector<std::unique_ptr<async_stream_array>> links_;
std::unique_ptr<async_stream_array> dealer_;
link_security dealer_sec_{};
wire_policy dealer_policy_{};
peer_address dealer_addr_{};
std::uint32_t dealer_epoch_ = 0;
bool dealer_is_server_ = false;
};
/// @brief Dealer serving several parties, each on its own link and epoch.
/// @details Links are encrypted like party links; the dealer authenticates a
/// party whose key is in `opt.security.trusted[p]`.
class dealer_session
{
public:
dealer_session(asio::io_context & io, unsigned served, session_options opt = {})
: io_(&io), served_(served), opt_(std::move(opt))
{
if (served == 0)
throw std::invalid_argument("dealer_session: serves no party");
opt_.policy.validate();
links_.resize(served);
sec_.resize(served);
epoch_.assign(served, 0);
if (opt_.security.encrypt)
{
#if DPF_HAS_OPENSSL
self_ = detail::link_identity(opt_.security, "dealer");
tls_ = make_peer_tls_context(*self_);
#else
throw std::logic_error("dealer_session: built without OpenSSL; set "
"encryption=off for plaintext links");
#endif
}
}
unsigned short listen(unsigned short port = 0)
{
if (!acceptor_)
{
acceptor_ = std::make_unique<asio::ip::tcp::acceptor>(*io_);
open_listener(*acceptor_, port);
log_listen(acceptor_->local_endpoint().port(), false, opt_.security.encrypt);
}
return acceptor_->local_endpoint().port();
}
/// @brief Accept one link from every served party (any order).
void accept_parties()
{
listen(0);
for (unsigned k = 0; k < served_; ++k)
accept_any(false, 0);
}
async_stream_array & party(unsigned p)
{
check(p);
if (!links_[p])
throw std::logic_error("dealer_session: no link to party "
+ std::to_string(p));
return *links_[p];
}
stream_stats party_stats(unsigned p) const
{
check(p);
return links_[p] ? links_[p]->stats() : stream_stats{};
}
const link_security & party_security(unsigned p) const
{
check(p);
return sec_[p];
}
std::shared_ptr<const identity> identity_key() const noexcept { return self_; }
async_stream_array & reconnect(unsigned p)
{
check(p);
links_[p].reset();
++epoch_[p];
accept_any(true, p);
return *links_[p];
}
std::function<async_stream_array *(const std::error_code &)> reconnector(
unsigned p)
{
return [this, p](const std::error_code &) { return &reconnect(p); };
}
private:
void check(unsigned p) const
{
if (p >= served_)
throw std::invalid_argument("dealer_session: bad party "
+ std::to_string(p));
}
void accept_any(bool expect_one, unsigned expect)
{
asio::ip::tcp::socket sock(*io_);
accept_until(*acceptor_, sock, opt_.limits.accept);
detail::peer_conn c(std::move(sock));
const auto deadline = setup_clock::now() + opt_.limits.handshake;
#if DPF_HAS_OPENSSL
detail::secure(c, *io_, tls_.get(), true, opt_.policy.socket,
opt_.limits.handshake, "dealer handshake");
#else
detail::secure(c, *io_, nullptr, true, opt_.policy.socket,
opt_.limits.handshake, "dealer handshake");
#endif
const auto theirs = detail::recv_hello(c, *io_, deadline, "dealer handshake");
const unsigned p = theirs.id;
if (p >= served_ || (expect_one && p != expect) || (!expect_one && links_[p]))
throw std::runtime_error("dealer_session: unexpected "
+ detail::id_name(p));
detail::check_party(c, opt_.security, p, "party " + std::to_string(p));
detail::session_hello mine;
mine.id = party_session::k_dealer_id;
mine.epoch = epoch_[p];
mine.kind = static_cast<std::uint32_t>(transport::mux);
mine.lanes = static_cast<std::uint32_t>(opt_.n_lanes);
mine.flags = c.sec.peer_auth == "key" ? detail::session_hello::verified_you : 0;
detail::send_hello(c, *io_, mine, deadline, "dealer handshake");
if (theirs.epoch != mine.epoch || theirs.lanes != mine.lanes
|| theirs.kind != mine.kind)
throw std::runtime_error("dealer_session: party " + std::to_string(p)
+ " disagrees on epoch, lanes, or transport");
c.sec.peer_verified_us = (theirs.flags & detail::session_hello::verified_you) != 0;
sec_[p] = c.sec;
links_[p] = detail::mux_link(*io_, c, party_session::k_dealer_id, p,
opt_.n_lanes, opt_.policy);
const std::string role = "p" + std::to_string(p);
log_link_up("accept", role, transport::mux, opt_.n_lanes, 0, epoch_[p], c.fd,
opt_.policy.socket, &sec_[p]);
detail::note_link(role, sec_[p]);
}
asio::io_context * io_ = nullptr;
unsigned served_ = 0;
session_options opt_;
std::shared_ptr<const identity> self_;
#if DPF_HAS_OPENSSL
std::shared_ptr<tls_context> tls_;
#endif
std::unique_ptr<asio::ip::tcp::acceptor> acceptor_;
std::vector<link_security> sec_;
std::vector<std::unique_ptr<async_stream_array>> links_;
std::vector<std::uint32_t> epoch_;
};
} // namespace net
} // namespace dpf
#endif