/// @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 #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/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(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(); for (unsigned other = 0; other < 3; ++other) { if (other == to_u(self)) continue; auto peer = static_cast(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(); 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 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(); for (unsigned other = 0; other < 3; ++other) { if (other == to_u(self)) continue; auto peer = static_cast(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 void send_to(role peer, msg tag, const T & value) { static_assert(std::is_trivially_copyable_v, "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 void send_vec_to(role peer, msg tag, const std::vector & values) { static_assert(std::is_trivially_copyable_v, "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 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 HEDLEY_WARN_UNUSED_RESULT T recv_from(role peer, msg tag) { static_assert(std::is_trivially_copyable_v, "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(tag); } /// @brief Homogeneous vector receive through the hook (or the mesh). template HEDLEY_WARN_UNUSED_RESULT std::vector recv_vec_from(role peer, msg tag = msg::beaver_tape) { static_assert(std::is_trivially_copyable_v, "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 out(bytes.size() / sizeof(T)); if (!out.empty()) std::memcpy(out.data(), bytes.data(), bytes.size()); return out; } return to(peer).template recv_vec(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 batch(std::size_t count, std::vector 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 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 HEDLEY_WARN_UNUSED_RESULT Ring accept_deal() { if (self_ == role::p2) throw std::logic_error("dealer does not accept_deal"); return recv_from(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 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(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 HEDLEY_WARN_UNUSED_RESULT T exchange_with(role peer, const T & mine, msg tag = msg::delta) { static_assert(std::is_trivially_copyable_v, "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(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 HEDLEY_WARN_UNUSED_RESULT std::vector exchange_vec_with(role peer, const std::vector & mine, msg tag = msg::ring_vector) { static_assert(std::is_trivially_copyable_v, "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 theirs(mine.size()); std::memcpy(theirs.data(), bytes.data(), nbytes); return theirs; } return to(peer).template exchange_vec( 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 HEDLEY_WARN_UNUSED_RESULT std::vector open_vec_with(role peer, const std::vector & 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 out(mine.size()); for (std::size_t i = 0; i < mine.size(); ++i) out[i] = static_cast(mine[i] + theirs[i]); return out; } private: role self_ = role::p0; std::unique_ptr 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(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__