libdpf/include/dpf/net/channel.hpp

485 lines
15 KiB
C++
Raw Normal View History

/// @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__