libdpf/include/dpf/net/trio.hpp

564 lines
19 KiB
C++
Raw Normal View History

/// @file dpf/net/trio.hpp
/// @brief Dial/accept three-party (2+1) socket mesh for libdpf protocols.
/// @details Socket paths live under a driver-created directory:
/// `p0-p1`, `p0-p2`, `p1-p2`. The lower role accepts; the higher
/// dials. Helpers `deal` / `accept_deal` / `open_with` name the
/// protocol steps without exposing ASIO.
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details.
#ifndef LIBDPF_INCLUDE_DPF_NET_TRIO_HPP__
#define LIBDPF_INCLUDE_DPF_NET_TRIO_HPP__
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <memory>
#include <stdexcept>
#include <string>
#include <thread>
#include <type_traits>
#include <utility>
#include <vector>
#include <unistd.h>
#include "dpf/net/asio_ns.hpp"
#include "dpf/net/channel.hpp"
#include "dpf/net/comm_hook.hpp"
#include "dpf/net/policy.hpp"
#include "dpf/net/round_sink.hpp"
#include "dpf/net/secure_channel.hpp"
#include "dpf/net/security.hpp"
#include "dpf/net/socket_tune.hpp"
#include "hedley/hedley.h"
namespace dpf
{
namespace net
{
/// @brief Party roles in a (2+1) run. p2 is the dealer.
enum class role : unsigned
{
p0 = 0,
p1 = 1,
p2 = 2,
};
HEDLEY_CONST
HEDLEY_NO_THROW
inline constexpr unsigned to_u(role r) noexcept
{
return static_cast<unsigned>(r);
}
HEDLEY_CONST
HEDLEY_NO_THROW
inline const char * role_name(role r) noexcept
{
switch (r)
{
case role::p0: return "p0";
case role::p1: return "p1";
case role::p2: return "p2";
}
return "?";
}
inline std::string link_path(const std::string & dir, role a, role b)
{
if (to_u(a) > to_u(b))
std::swap(a, b);
return dir + "/" + role_name(a) + "-" + role_name(b);
}
/// @brief Connected view of the other two parties.
class trio
{
public:
trio() = default;
/// @brief Connect local unix-domain sockets under `dir`.
/// @param self this process's role
/// @param dir directory holding the three socket paths
/// @param retries dial/accept attempts before giving up
/// @param sec same peer TLS policy as `party_session` (default encrypt)
static trio connect_local(role self, const std::string & dir,
unsigned retries = 200, const peer_security & sec = {},
const deadlines & lim = {})
{
trio t;
t.self_ = self;
t.io_ = std::make_unique<asio::io_context>();
for (unsigned other = 0; other < 3; ++other)
{
if (other == to_u(self))
continue;
auto peer = static_cast<role>(other);
t.link_[other] = connect_one(self, peer, dir, *t.io_, retries, sec,
lim);
}
return t;
}
/// @brief Connect the p0–p1 socket only. The dealer link stays closed
/// until `install_inbox`.
static trio connect_pair(role self, const std::string & dir,
unsigned retries = 200, const peer_security & sec = {},
const deadlines & lim = {})
{
if (self == role::p2)
throw std::invalid_argument("connect_pair is p0 and p1");
trio t;
t.self_ = self;
t.io_ = std::make_unique<asio::io_context>();
const role peer = self == role::p0 ? role::p1 : role::p0;
t.link_[to_u(peer)] = connect_one(self, peer, dir, *t.io_, retries, sec,
lim);
return t;
}
/// @brief Serve later `to(p2).recv` calls from a local frame buffer.
void install_inbox(std::vector<std::uint8_t> frames)
{
if (self_ == role::p2)
throw std::logic_error("p2 has no dealer inbox");
link_[to_u(role::p2)] = channel::from_inbox(std::move(frames));
}
/// @brief True when the installed dealer tape has been fully consumed.
HEDLEY_NO_THROW
bool dealer_inbox_done() const noexcept
{
if (self_ == role::p2)
return false;
return link_[to_u(role::p2)].inbox_done();
}
/// @brief Connect over TCP loopback. Ports are `base + 10*lo + hi`.
static trio connect_tcp(role self, std::uint16_t base_port,
const std::string & host = "127.0.0.1", unsigned retries = 200,
const peer_security & sec = {}, const deadlines & lim = {},
const socket_options & so = {})
{
trio t;
t.self_ = self;
t.io_ = std::make_unique<asio::io_context>();
for (unsigned other = 0; other < 3; ++other)
{
if (other == to_u(self))
continue;
auto peer = static_cast<role>(other);
t.link_[other] = connect_one_tcp(self, peer, base_port, host,
*t.io_, retries, sec, lim, so);
}
return t;
}
role self() const noexcept { return self_; }
/// @brief Sum of byte and frame counters on the open links.
HEDLEY_NO_THROW
io_tally tally() const noexcept
{
io_tally sum;
for (const auto & link : link_)
{
if (link.open())
sum += link.tally();
}
return sum;
}
/// @brief Counters on the link to `peer` only.
HEDLEY_NO_THROW
io_tally tally_to(role peer) const noexcept
{
if (peer == self_ || !link_[to_u(peer)].open())
return {};
return link_[to_u(peer)].tally();
}
/// @brief Bytes received from the dealer (p2). Empty for the dealer itself.
HEDLEY_NO_THROW
io_tally tally_from_p2() const noexcept
{
if (self_ == role::p2)
return {};
return tally_to(role::p2);
}
/// @brief Bytes with the other computing party (p0↔p1).
/// @details For p2 this is the sum of the links to p0 and p1.
HEDLEY_NO_THROW
io_tally tally_from_peer() const noexcept
{
if (self_ == role::p0)
return tally_to(role::p1);
if (self_ == role::p1)
return tally_to(role::p0);
io_tally sum;
sum += tally_to(role::p0);
sum += tally_to(role::p1);
return sum;
}
/// @brief Zero counters on every open link.
HEDLEY_NO_THROW
void reset_tally() noexcept
{
for (auto & link : link_)
{
if (link.open())
link.reset_tally();
}
}
/// @brief The open channel to `peer` (raw bypass; does not use the hook).
/// @param peer the other party
/// @return that channel
/// @throws std::invalid_argument if `peer` is this process
/// @throws std::logic_error if that link was not connected
channel & to(role peer)
{
if (peer == self_)
throw std::invalid_argument("trio::to(self)");
auto & c = link_[to_u(peer)];
if (!c.open())
throw std::logic_error("trio link is not open");
return c;
}
/// @brief Install a replaceable transport for helpers and `batch`.
/// @details Non-owning. Null restores the framed mesh defaults.
void set_hook(comm_hook * hook) noexcept { hook_ = hook; }
/// @brief Current hook, or null when helpers use the mesh directly.
HEDLEY_NO_THROW
comm_hook * hook() const noexcept { return hook_; }
/// @brief One-way send through the hook (or the mesh when unset).
template <typename T>
void send_to(role peer, msg tag, const T & value)
{
static_assert(std::is_trivially_copyable_v<T>,
"trio::send_to requires a trivially copyable type");
if (hook_)
{
hook_->send_bytes(*this, to_u(peer), tag, &value, sizeof(T));
return;
}
to(peer).send(tag, value);
}
/// @brief Homogeneous vector send through the hook (or the mesh).
template <typename T>
void send_vec_to(role peer, msg tag, const std::vector<T> & values)
{
static_assert(std::is_trivially_copyable_v<T>,
"trio::send_vec_to requires a trivially copyable type");
if (hook_)
{
const std::size_t nbytes = values.size() * sizeof(T);
hook_->send_bytes(*this, to_u(peer), tag, values.data(), nbytes);
return;
}
to(peer).send_vec(values, tag);
}
/// @brief Untyped one-way send through the hook (or the mesh).
void send_bytes_to(role peer, msg tag, const void * data, std::size_t n)
{
if (hook_)
{
hook_->send_bytes(*this, to_u(peer), tag, data, n);
return;
}
to(peer).send_bytes(tag, data, n);
}
/// @brief Untyped one-way receive through the hook (or the mesh).
HEDLEY_WARN_UNUSED_RESULT
std::vector<std::uint8_t> recv_bytes_from(role peer, msg tag)
{
if (hook_)
return hook_->recv_bytes(*this, to_u(peer), tag);
return to(peer).recv_bytes(tag);
}
/// @brief One-way receive through the hook (or the mesh when unset).
template <typename T>
HEDLEY_WARN_UNUSED_RESULT
T recv_from(role peer, msg tag)
{
static_assert(std::is_trivially_copyable_v<T>,
"trio::recv_from requires a trivially copyable type");
if (hook_)
{
auto bytes = hook_->recv_bytes(*this, to_u(peer), tag);
if (bytes.size() != sizeof(T))
throw std::runtime_error("trio::recv_from size mismatch");
T value{};
if constexpr (sizeof(T) != 0)
std::memcpy(&value, bytes.data(), sizeof(T));
return value;
}
return to(peer).template recv<T>(tag);
}
/// @brief Homogeneous vector receive through the hook (or the mesh).
template <typename T>
HEDLEY_WARN_UNUSED_RESULT
std::vector<T> recv_vec_from(role peer, msg tag = msg::beaver_tape)
{
static_assert(std::is_trivially_copyable_v<T>,
"trio::recv_vec_from requires a trivially copyable type");
if (hook_)
{
auto bytes = hook_->recv_bytes(*this, to_u(peer), tag);
if constexpr (sizeof(T) == 0)
{
if (!bytes.empty())
throw std::runtime_error("trio::recv_vec_from size");
return {};
}
if (bytes.size() % sizeof(T) != 0)
throw std::runtime_error("trio::recv_vec_from size mismatch");
std::vector<T> out(bytes.size() / sizeof(T));
if (!out.empty())
std::memcpy(out.data(), bytes.data(), bytes.size());
return out;
}
return to(peer).template recv_vec<T>(tag);
}
/// @brief Round-batched peer sink. Default is a mux on the p0–p1 link.
/// @details Defined next to `mux_sink` so the default factory can build one.
std::unique_ptr<RoundSink> batch(std::size_t count,
std::vector<std::size_t> slot_bytes);
/// @brief Dealer sends one side of a split to each computing party.
/// @tparam Split a type with `p0` and `p1` members
/// @param s the split
/// @throws std::logic_error if this process is not p2
template <typename Split>
void deal(const Split & s)
{
if (self_ != role::p2)
throw std::logic_error("only the dealer may deal");
send_to(role::p0, msg::beaver_tape, s.p0);
send_to(role::p1, msg::beaver_tape, s.p1);
}
/// @brief Computing party receives its share of a dealt value.
/// @tparam Ring the dealt type
/// @return this party's share
/// @throws std::logic_error if this process is p2
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring accept_deal()
{
if (self_ == role::p2)
throw std::logic_error("dealer does not accept_deal");
return recv_from<Ring>(role::p2, msg::beaver_tape);
}
/// @brief Open an additive share with the peer computing party.
/// @tparam Ring additive ring
/// @param peer the other computing party
/// @param mine this party's share
/// @return reconstructed public value (`mine + peer`)
/// @throws std::logic_error if either side is p2
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring open_with(role peer, const Ring & mine)
{
if (self_ == role::p2 || peer == role::p2)
throw std::logic_error("open_with is between p0 and p1");
Ring theirs = exchange_with(peer, mine, msg::delta);
return static_cast<Ring>(mine + theirs);
}
/// @brief Exchange raw values (no algebra) with a peer.
/// @tparam T trivially copyable payload
/// @param peer the other party
/// @param mine this party's value
/// @param tag the message tag
/// @return the peer's value
template <typename T>
HEDLEY_WARN_UNUSED_RESULT
T exchange_with(role peer, const T & mine, msg tag = msg::delta)
{
static_assert(std::is_trivially_copyable_v<T>,
"trio::exchange_with requires a trivially copyable type");
if (hook_)
{
auto bytes = hook_->exchange_bytes(*this, to_u(peer), tag, &mine,
sizeof(T));
if (bytes.size() != sizeof(T))
throw std::runtime_error("trio::exchange_with size mismatch");
T theirs{};
if constexpr (sizeof(T) != 0)
std::memcpy(&theirs, bytes.data(), sizeof(T));
return theirs;
}
return to(peer).template exchange<T>(to_u(self_), to_u(peer), mine, tag);
}
/// @brief Exchange a homogeneous vector with a peer (one barrier).
/// @tparam T trivially copyable element
/// @param peer the other party
/// @param mine this party's values
/// @param tag the message tag
/// @return the peer's values
template <typename T>
HEDLEY_WARN_UNUSED_RESULT
std::vector<T> exchange_vec_with(role peer, const std::vector<T> & mine,
msg tag = msg::ring_vector)
{
static_assert(std::is_trivially_copyable_v<T>,
"trio::exchange_vec_with requires a trivially copyable type");
if (mine.empty())
return {};
if (hook_)
{
const std::size_t nbytes = mine.size() * sizeof(T);
auto bytes = hook_->exchange_vec_bytes(*this, to_u(peer), tag,
mine.data(), nbytes);
if (bytes.size() != nbytes)
throw std::runtime_error("trio::exchange_vec_with size mismatch");
std::vector<T> theirs(mine.size());
std::memcpy(theirs.data(), bytes.data(), nbytes);
return theirs;
}
return to(peer).template exchange_vec<T>(
to_u(self_), to_u(peer), mine, tag);
}
/// @brief Open additive shares packed into one vector exchange.
/// @tparam Ring additive ring
/// @param peer the other computing party
/// @param mine this party's shares
/// @return reconstructed public values (`mine[i] + peer[i]`)
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
std::vector<Ring> open_vec_with(role peer, const std::vector<Ring> & mine)
{
if (self_ == role::p2 || peer == role::p2)
throw std::logic_error("open_vec_with is between p0 and p1");
auto theirs = exchange_vec_with(peer, mine);
std::vector<Ring> out(mine.size());
for (std::size_t i = 0; i < mine.size(); ++i)
out[i] = static_cast<Ring>(mine[i] + theirs[i]);
return out;
}
private:
role self_ = role::p0;
std::unique_ptr<asio::io_context> io_;
channel link_[3];
comm_hook * hook_ = nullptr;
static channel connect_one(role self, role peer, const std::string & dir,
asio::io_context & io, unsigned retries, const peer_security & sec,
const deadlines & lim)
{
using proto = asio::local::stream_protocol;
const bool accept = to_u(self) < to_u(peer);
const auto path = link_path(dir, self, peer);
const std::string who = role_name(peer);
if (accept)
{
::unlink(path.c_str());
proto::endpoint ep(path);
proto::acceptor acc(io, ep);
for (unsigned i = 0; i < retries; ++i)
{
asio::error_code ec;
proto::socket sock(io);
acc.accept(sock, ec);
if (!ec)
return secure_local_channel(io, std::move(sock), true,
to_u(peer), sec, lim.handshake, who, "accept");
std::this_thread::sleep_for(std::chrono::milliseconds(10));
}
throw std::runtime_error("trio accept failed: " + path);
}
proto::endpoint ep(path);
for (unsigned i = 0; i < retries; ++i)
{
asio::error_code ec;
proto::socket sock(io);
sock.connect(ep, ec);
if (!ec)
return secure_local_channel(io, std::move(sock), false,
to_u(peer), sec, lim.handshake, who, "connect");
std::this_thread::sleep_for(std::chrono::milliseconds(10));
}
throw std::runtime_error("trio connect failed: " + path);
}
static std::uint16_t tcp_port(std::uint16_t base, role a, role b)
{
if (to_u(a) > to_u(b))
std::swap(a, b);
return static_cast<std::uint16_t>(base + 10u * to_u(a) + to_u(b));
}
static channel connect_one_tcp(role self, role peer, std::uint16_t base,
const std::string & host, asio::io_context & io, unsigned retries,
const peer_security & sec, const deadlines & lim,
const socket_options & so)
{
using tcp = asio::ip::tcp;
const bool accept = to_u(self) < to_u(peer);
const auto port = tcp_port(base, self, peer);
const std::string who = role_name(peer);
if (accept)
{
tcp::endpoint ep(tcp::v4(), port);
tcp::acceptor acc(io);
asio::error_code ec;
acc.open(ep.protocol(), ec);
acc.set_option(tcp::acceptor::reuse_address(true), ec);
acc.bind(ep, ec);
if (ec)
throw std::runtime_error("trio tcp bind: " + ec.message());
acc.listen(1, ec);
for (unsigned i = 0; i < retries; ++i)
{
tcp::socket sock(io);
acc.accept(sock, ec);
if (!ec)
return secure_tcp_channel(io, std::move(sock), true,
to_u(peer), sec, so, lim.handshake, who, "accept");
std::this_thread::sleep_for(std::chrono::milliseconds(10));
}
throw std::runtime_error("trio tcp accept failed");
}
tcp::resolver resolver(io);
auto endpoints = resolver.resolve(host, std::to_string(port));
for (unsigned i = 0; i < retries; ++i)
{
asio::error_code ec;
tcp::socket sock(io);
asio::connect(sock, endpoints, ec);
if (!ec)
return secure_tcp_channel(io, std::move(sock), false,
to_u(peer), sec, so, lim.handshake, who, "connect");
std::this_thread::sleep_for(std::chrono::milliseconds(10));
}
throw std::runtime_error("trio tcp connect failed");
}
};
} // namespace net
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_NET_TRIO_HPP__