485 lines
15 KiB
C++
485 lines
15 KiB
C++
|
|
/// @file dpf/net/channel.hpp
|
||
|
|
/// @brief Framed duplex stream for party message exchange.
|
||
|
|
/// @details Every message is `u32` little-endian length, `u16` type tag, then
|
||
|
|
/// payload. `exchange` orders send/recv by role so a single stream
|
||
|
|
/// cannot deadlock. Call sites name `send` / `recv` / `exchange`;
|
||
|
|
/// they do not touch ASIO buffers.
|
||
|
|
/// @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_CHANNEL_HPP__
|
||
|
|
#define LIBDPF_INCLUDE_DPF_NET_CHANNEL_HPP__
|
||
|
|
|
||
|
|
#include <array>
|
||
|
|
#include <cstddef>
|
||
|
|
#include <cstdint>
|
||
|
|
#include <cstring>
|
||
|
|
#include <memory>
|
||
|
|
#include <stdexcept>
|
||
|
|
#include <string>
|
||
|
|
#include <type_traits>
|
||
|
|
#include <utility>
|
||
|
|
#include <vector>
|
||
|
|
|
||
|
|
#include "dpf/net/asio_ns.hpp"
|
||
|
|
#include "dpf/net/tls.hpp"
|
||
|
|
|
||
|
|
#include "hedley/hedley.h"
|
||
|
|
|
||
|
|
namespace dpf
|
||
|
|
{
|
||
|
|
namespace net
|
||
|
|
{
|
||
|
|
|
||
|
|
/// @brief Bytes and frames observed on one channel since the last reset.
|
||
|
|
/// @details Counts include the 6-byte frame header. A failed read or write
|
||
|
|
/// does not add to the tally. `payload_*` is the body alone;
|
||
|
|
/// `exchanges` counts `exchange()` calls (one logical round each).
|
||
|
|
struct io_tally
|
||
|
|
{
|
||
|
|
std::uint64_t bytes_sent = 0;
|
||
|
|
std::uint64_t bytes_recv = 0;
|
||
|
|
std::uint64_t frames_sent = 0;
|
||
|
|
std::uint64_t frames_recv = 0;
|
||
|
|
std::uint64_t payload_sent = 0;
|
||
|
|
std::uint64_t payload_recv = 0;
|
||
|
|
std::uint64_t exchanges = 0;
|
||
|
|
|
||
|
|
io_tally & operator+=(const io_tally & other) noexcept
|
||
|
|
{
|
||
|
|
bytes_sent += other.bytes_sent;
|
||
|
|
bytes_recv += other.bytes_recv;
|
||
|
|
frames_sent += other.frames_sent;
|
||
|
|
frames_recv += other.frames_recv;
|
||
|
|
payload_sent += other.payload_sent;
|
||
|
|
payload_recv += other.payload_recv;
|
||
|
|
exchanges += other.exchanges;
|
||
|
|
return *this;
|
||
|
|
}
|
||
|
|
};
|
||
|
|
|
||
|
|
enum class msg : std::uint16_t
|
||
|
|
{
|
||
|
|
hangup = 0,
|
||
|
|
beaver_tape = 1,
|
||
|
|
ring_vector = 2,
|
||
|
|
delta = 3,
|
||
|
|
dpf_key = 4,
|
||
|
|
proof_token = 5,
|
||
|
|
case_ok = 6,
|
||
|
|
case_fail = 7,
|
||
|
|
bytes = 8,
|
||
|
|
mac_key = 9,
|
||
|
|
mac_share = 10,
|
||
|
|
sketch_share = 11,
|
||
|
|
round_batch = 12,
|
||
|
|
};
|
||
|
|
|
||
|
|
inline constexpr std::uint16_t to_u16(msg t) noexcept
|
||
|
|
{
|
||
|
|
return static_cast<std::uint16_t>(t);
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief One framed duplex byte stream.
|
||
|
|
class channel
|
||
|
|
{
|
||
|
|
public:
|
||
|
|
using tcp_socket = asio::ip::tcp::socket;
|
||
|
|
using local_socket = asio::local::stream_protocol::socket;
|
||
|
|
|
||
|
|
channel() = default;
|
||
|
|
|
||
|
|
explicit channel(tcp_socket sock)
|
||
|
|
: tcp_(std::make_unique<tcp_socket>(std::move(sock)))
|
||
|
|
{ }
|
||
|
|
|
||
|
|
explicit channel(local_socket sock)
|
||
|
|
: local_(std::make_unique<local_socket>(std::move(sock)))
|
||
|
|
{ }
|
||
|
|
|
||
|
|
#if DPF_HAS_OPENSSL
|
||
|
|
/// @brief Framed channel over an established TLS 1.3 TCP stream.
|
||
|
|
static channel from_tls(asio::io_context &, tls_stream s,
|
||
|
|
std::shared_ptr<tls_context> ctx)
|
||
|
|
{
|
||
|
|
channel c;
|
||
|
|
c.tls_ctx_ = std::move(ctx);
|
||
|
|
c.tls_ = std::make_unique<tls_stream>(std::move(s));
|
||
|
|
return c;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Framed channel over an established TLS 1.3 unix-domain stream.
|
||
|
|
static channel from_tls_local(asio::io_context &, tls_local_stream s,
|
||
|
|
std::shared_ptr<tls_context> ctx)
|
||
|
|
{
|
||
|
|
channel c;
|
||
|
|
c.tls_ctx_ = std::move(ctx);
|
||
|
|
c.tls_local_ = std::make_unique<tls_local_stream>(std::move(s));
|
||
|
|
return c;
|
||
|
|
}
|
||
|
|
#endif
|
||
|
|
|
||
|
|
channel(channel &&) noexcept = default;
|
||
|
|
channel & operator=(channel &&) noexcept = default;
|
||
|
|
|
||
|
|
channel(const channel &) = delete;
|
||
|
|
channel & operator=(const channel &) = delete;
|
||
|
|
|
||
|
|
/// @brief Replay `frames` as a recv-only dealer tape. `send` throws.
|
||
|
|
static channel from_inbox(std::vector<std::uint8_t> frames)
|
||
|
|
{
|
||
|
|
channel c;
|
||
|
|
c.inbox_ = std::make_unique<inbox_buf>();
|
||
|
|
c.inbox_->bytes = std::move(frames);
|
||
|
|
return c;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief True when an inbox tape has been read through its last byte.
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
bool inbox_done() const noexcept
|
||
|
|
{
|
||
|
|
return inbox_ != nullptr && inbox_->pos == inbox_->bytes.size();
|
||
|
|
}
|
||
|
|
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
bool open() const noexcept
|
||
|
|
{
|
||
|
|
return tcp_ != nullptr || local_ != nullptr || inbox_ != nullptr
|
||
|
|
#if DPF_HAS_OPENSSL
|
||
|
|
|| tls_ != nullptr || tls_local_ != nullptr
|
||
|
|
#endif
|
||
|
|
;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Underlying socket descriptor, or -1 for an inbox tape / TLS edge.
|
||
|
|
/// @details Prefer framed `send` / `recv` on TLS channels. Raw descriptor
|
||
|
|
/// I/O bypasses TLS and is rejected when the channel is encrypted.
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
int native_handle() noexcept
|
||
|
|
{
|
||
|
|
if (tcp_)
|
||
|
|
return tcp_->native_handle();
|
||
|
|
if (local_)
|
||
|
|
return local_->native_handle();
|
||
|
|
#if DPF_HAS_OPENSSL
|
||
|
|
if (tls_ || tls_local_)
|
||
|
|
return -1;
|
||
|
|
#endif
|
||
|
|
return -1;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief True when application bytes ride TLS 1.3 under the frame header.
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
bool encrypted() const noexcept
|
||
|
|
{
|
||
|
|
#if DPF_HAS_OPENSSL
|
||
|
|
return tls_ != nullptr || tls_local_ != nullptr;
|
||
|
|
#else
|
||
|
|
return false;
|
||
|
|
#endif
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Bytes and frames since construction or the last `reset_tally`.
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
io_tally tally() const noexcept
|
||
|
|
{
|
||
|
|
return tally_;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Zero the byte and frame counters. Does not touch the socket.
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
void reset_tally() noexcept
|
||
|
|
{
|
||
|
|
tally_ = {};
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @name Framed messages
|
||
|
|
/// @brief Each frame is a little-endian length, a tag, then the payload.
|
||
|
|
/// `T` must be trivially copyable. A zero-length payload is a
|
||
|
|
/// header only.
|
||
|
|
/// @throws std::runtime_error if the socket closes or the tag/size disagree
|
||
|
|
/// @throws std::invalid_argument if a frame exceeds 2^32-1 bytes, or
|
||
|
|
/// `exchange` is called with `self_id == peer_id`
|
||
|
|
/// @throws std::logic_error if the channel is closed
|
||
|
|
/// @{
|
||
|
|
|
||
|
|
/// @brief Send one value.
|
||
|
|
/// @tparam T trivially copyable payload
|
||
|
|
/// @param tag the message tag
|
||
|
|
/// @param value the payload
|
||
|
|
template <typename T>
|
||
|
|
void send(msg tag, const T & value)
|
||
|
|
{
|
||
|
|
static_assert(std::is_trivially_copyable_v<T>,
|
||
|
|
"net::channel::send requires a trivially copyable type");
|
||
|
|
write_frame(tag, &value, sizeof(T));
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Receive one value.
|
||
|
|
/// @tparam T trivially copyable payload
|
||
|
|
/// @param tag the expected message tag
|
||
|
|
/// @return the payload
|
||
|
|
template <typename T>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
T recv(msg tag)
|
||
|
|
{
|
||
|
|
static_assert(std::is_trivially_copyable_v<T>,
|
||
|
|
"net::channel::recv requires a trivially copyable type");
|
||
|
|
T value{};
|
||
|
|
read_frame(tag, &value, sizeof(T));
|
||
|
|
return value;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Send `n` values.
|
||
|
|
/// @param data the values
|
||
|
|
/// @param n the value count
|
||
|
|
/// @param tag the message tag
|
||
|
|
template <typename T>
|
||
|
|
void send_vec(const T * data, std::size_t n, msg tag = msg::ring_vector)
|
||
|
|
{
|
||
|
|
static_assert(std::is_trivially_copyable_v<T>,
|
||
|
|
"net::channel::send_vec requires a trivially copyable type");
|
||
|
|
write_frame(tag, data, n * sizeof(T));
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Send a vector of values.
|
||
|
|
/// @param v the values
|
||
|
|
/// @param tag the message tag
|
||
|
|
template <typename T>
|
||
|
|
void send_vec(const std::vector<T> & v, msg tag = msg::ring_vector)
|
||
|
|
{
|
||
|
|
send_vec(v.data(), v.size(), tag);
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Receive a homogeneous vector.
|
||
|
|
/// @tparam T trivially copyable element
|
||
|
|
/// @param tag the expected message tag
|
||
|
|
/// @return the elements. Empty when the payload length is 0.
|
||
|
|
template <typename T>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::vector<T> recv_vec(msg tag = msg::ring_vector)
|
||
|
|
{
|
||
|
|
static_assert(std::is_trivially_copyable_v<T>,
|
||
|
|
"net::channel::recv_vec requires a trivially copyable type");
|
||
|
|
auto bytes = read_frame_bytes(tag);
|
||
|
|
if (bytes.size() % sizeof(T) != 0)
|
||
|
|
throw std::runtime_error("net::channel::recv_vec size mismatch");
|
||
|
|
std::vector<T> out(bytes.size() / sizeof(T));
|
||
|
|
if (!out.empty())
|
||
|
|
std::memcpy(out.data(), bytes.data(), bytes.size());
|
||
|
|
return out;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Send `n` bytes.
|
||
|
|
/// @param tag the message tag
|
||
|
|
/// @param data the bytes
|
||
|
|
/// @param n the byte count
|
||
|
|
void send_bytes(msg tag, const void * data, std::size_t n)
|
||
|
|
{
|
||
|
|
write_frame(tag, data, n);
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Send a byte vector.
|
||
|
|
/// @param tag the message tag
|
||
|
|
/// @param bytes the bytes
|
||
|
|
void send_bytes(msg tag, const std::vector<std::uint8_t> & bytes)
|
||
|
|
{
|
||
|
|
send_bytes(tag, bytes.data(), bytes.size());
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Receive an untyped payload.
|
||
|
|
/// @param tag the expected message tag
|
||
|
|
/// @return the payload bytes
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::vector<std::uint8_t> recv_bytes(msg tag)
|
||
|
|
{
|
||
|
|
return read_frame_bytes(tag);
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Exchange one value. The lower id sends first.
|
||
|
|
/// @tparam T trivially copyable payload
|
||
|
|
/// @param self_id this party's role, as an integer
|
||
|
|
/// @param peer_id the peer's role, as an integer
|
||
|
|
/// @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(unsigned self_id, unsigned peer_id, const T & mine, msg tag = msg::delta)
|
||
|
|
{
|
||
|
|
static_assert(std::is_trivially_copyable_v<T>,
|
||
|
|
"net::channel::exchange requires a trivially copyable type");
|
||
|
|
++tally_.exchanges;
|
||
|
|
if (self_id < peer_id)
|
||
|
|
{
|
||
|
|
send(tag, mine);
|
||
|
|
return recv<T>(tag);
|
||
|
|
}
|
||
|
|
if (self_id > peer_id)
|
||
|
|
{
|
||
|
|
T theirs = recv<T>(tag);
|
||
|
|
send(tag, mine);
|
||
|
|
return theirs;
|
||
|
|
}
|
||
|
|
throw std::invalid_argument("net::channel::exchange with self");
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Exchange a homogeneous vector. One barrier; lower id sends first.
|
||
|
|
/// @tparam T trivially copyable element
|
||
|
|
/// @param self_id this party's role, as an integer
|
||
|
|
/// @param peer_id the peer's role, as an integer
|
||
|
|
/// @param mine this party's values
|
||
|
|
/// @param tag the message tag
|
||
|
|
/// @return the peer's values (same length as `mine`)
|
||
|
|
/// @throws std::runtime_error if the peer's vector length disagrees
|
||
|
|
template <typename T>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::vector<T> exchange_vec(unsigned self_id, unsigned peer_id,
|
||
|
|
const std::vector<T> & mine, msg tag = msg::ring_vector)
|
||
|
|
{
|
||
|
|
static_assert(std::is_trivially_copyable_v<T>,
|
||
|
|
"net::channel::exchange_vec requires a trivially copyable type");
|
||
|
|
if (mine.empty())
|
||
|
|
return {};
|
||
|
|
++tally_.exchanges;
|
||
|
|
std::vector<T> theirs;
|
||
|
|
if (self_id < peer_id)
|
||
|
|
{
|
||
|
|
send_vec(mine, tag);
|
||
|
|
theirs = recv_vec<T>(tag);
|
||
|
|
}
|
||
|
|
else if (self_id > peer_id)
|
||
|
|
{
|
||
|
|
theirs = recv_vec<T>(tag);
|
||
|
|
send_vec(mine, tag);
|
||
|
|
}
|
||
|
|
else
|
||
|
|
throw std::invalid_argument("net::channel::exchange_vec with self");
|
||
|
|
if (theirs.size() != mine.size())
|
||
|
|
throw std::runtime_error("net::channel::exchange_vec size mismatch");
|
||
|
|
return theirs;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @}
|
||
|
|
|
||
|
|
private:
|
||
|
|
struct inbox_buf
|
||
|
|
{
|
||
|
|
std::vector<std::uint8_t> bytes;
|
||
|
|
std::size_t pos = 0;
|
||
|
|
};
|
||
|
|
|
||
|
|
std::unique_ptr<tcp_socket> tcp_;
|
||
|
|
std::unique_ptr<local_socket> local_;
|
||
|
|
#if DPF_HAS_OPENSSL
|
||
|
|
std::unique_ptr<tls_stream> tls_;
|
||
|
|
std::unique_ptr<tls_local_stream> tls_local_;
|
||
|
|
std::shared_ptr<tls_context> tls_ctx_;
|
||
|
|
#endif
|
||
|
|
std::unique_ptr<inbox_buf> inbox_;
|
||
|
|
io_tally tally_{};
|
||
|
|
|
||
|
|
void write_all(const void * data, std::size_t n)
|
||
|
|
{
|
||
|
|
if (inbox_)
|
||
|
|
throw std::logic_error("net::channel inbox is recv-only");
|
||
|
|
asio::const_buffer buf(data, n);
|
||
|
|
asio::error_code ec;
|
||
|
|
if (tcp_)
|
||
|
|
asio::write(*tcp_, buf, ec);
|
||
|
|
else if (local_)
|
||
|
|
asio::write(*local_, buf, ec);
|
||
|
|
#if DPF_HAS_OPENSSL
|
||
|
|
else if (tls_)
|
||
|
|
asio::write(*tls_, buf, ec);
|
||
|
|
else if (tls_local_)
|
||
|
|
asio::write(*tls_local_, buf, ec);
|
||
|
|
#endif
|
||
|
|
else
|
||
|
|
throw std::logic_error("net::channel is closed");
|
||
|
|
if (ec)
|
||
|
|
throw std::runtime_error("net::channel write: " + ec.message());
|
||
|
|
tally_.bytes_sent += n;
|
||
|
|
}
|
||
|
|
|
||
|
|
void read_all(void * data, std::size_t n)
|
||
|
|
{
|
||
|
|
if (inbox_)
|
||
|
|
{
|
||
|
|
if (inbox_->pos + n > inbox_->bytes.size())
|
||
|
|
throw std::runtime_error("net::channel inbox underrun");
|
||
|
|
if (n != 0)
|
||
|
|
std::memcpy(data, inbox_->bytes.data() + inbox_->pos, n);
|
||
|
|
inbox_->pos += n;
|
||
|
|
tally_.bytes_recv += n;
|
||
|
|
return;
|
||
|
|
}
|
||
|
|
asio::mutable_buffer buf(data, n);
|
||
|
|
asio::error_code ec;
|
||
|
|
if (tcp_)
|
||
|
|
asio::read(*tcp_, buf, ec);
|
||
|
|
else if (local_)
|
||
|
|
asio::read(*local_, buf, ec);
|
||
|
|
#if DPF_HAS_OPENSSL
|
||
|
|
else if (tls_)
|
||
|
|
asio::read(*tls_, buf, ec);
|
||
|
|
else if (tls_local_)
|
||
|
|
asio::read(*tls_local_, buf, ec);
|
||
|
|
#endif
|
||
|
|
else
|
||
|
|
throw std::logic_error("net::channel is closed");
|
||
|
|
if (ec)
|
||
|
|
throw std::runtime_error("net::channel read: " + ec.message());
|
||
|
|
tally_.bytes_recv += n;
|
||
|
|
}
|
||
|
|
|
||
|
|
void write_frame(msg tag, const void * payload, std::size_t n)
|
||
|
|
{
|
||
|
|
if (n > 0xffffffffu)
|
||
|
|
throw std::invalid_argument("net::channel frame too large");
|
||
|
|
std::uint32_t len = static_cast<std::uint32_t>(n);
|
||
|
|
std::uint16_t t = to_u16(tag);
|
||
|
|
std::array<std::uint8_t, 6> hdr{};
|
||
|
|
std::memcpy(hdr.data(), &len, 4);
|
||
|
|
std::memcpy(hdr.data() + 4, &t, 2);
|
||
|
|
write_all(hdr.data(), hdr.size());
|
||
|
|
if (n != 0)
|
||
|
|
write_all(payload, n);
|
||
|
|
tally_.payload_sent += n;
|
||
|
|
++tally_.frames_sent;
|
||
|
|
}
|
||
|
|
|
||
|
|
std::vector<std::uint8_t> read_frame_bytes(msg expected)
|
||
|
|
{
|
||
|
|
std::array<std::uint8_t, 6> hdr{};
|
||
|
|
read_all(hdr.data(), hdr.size());
|
||
|
|
std::uint32_t len = 0;
|
||
|
|
std::uint16_t t = 0;
|
||
|
|
std::memcpy(&len, hdr.data(), 4);
|
||
|
|
std::memcpy(&t, hdr.data() + 4, 2);
|
||
|
|
if (t != to_u16(expected))
|
||
|
|
throw std::runtime_error("net::channel unexpected message tag");
|
||
|
|
std::vector<std::uint8_t> body(len);
|
||
|
|
if (len != 0)
|
||
|
|
read_all(body.data(), body.size());
|
||
|
|
tally_.payload_recv += len;
|
||
|
|
++tally_.frames_recv;
|
||
|
|
return body;
|
||
|
|
}
|
||
|
|
|
||
|
|
void read_frame(msg expected, void * dest, std::size_t expect_n)
|
||
|
|
{
|
||
|
|
auto body = read_frame_bytes(expected);
|
||
|
|
if (body.size() != expect_n)
|
||
|
|
throw std::runtime_error("net::channel frame size mismatch");
|
||
|
|
if (expect_n != 0)
|
||
|
|
std::memcpy(dest, body.data(), expect_n);
|
||
|
|
}
|
||
|
|
};
|
||
|
|
|
||
|
|
} // namespace net
|
||
|
|
} // namespace dpf
|
||
|
|
|
||
|
|
#endif // LIBDPF_INCLUDE_DPF_NET_CHANNEL_HPP__
|