libdpf/include/dpf/party_run.hpp

369 lines
11 KiB
C++
Raw Normal View History

/// @file dpf/party_run.hpp
/// @brief Run one protocol as threads on a memory clique, or as TCP peers.
#ifndef LIBDPF_INCLUDE_DPF_PARTY_RUN_HPP__
#define LIBDPF_INCLUDE_DPF_PARTY_RUN_HPP__
#include <atomic>
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <functional>
#include <mutex>
#include <stdexcept>
#include <string>
#include <thread>
#include <utility>
#include <vector>
#include "dpf/net/asio_ns.hpp"
#include "dpf/net/channel.hpp"
#include "dpf/net/connect.hpp"
#include "dpf/net/edge_mesh.hpp"
#include "dpf/net/memory_sink.hpp"
#include "dpf/net/party_session.hpp"
#include "dpf/net/secure_channel.hpp"
#include "dpf/net/security.hpp"
#include "dpf/net/socket_tune.hpp"
#include "dpf/net/stream_array.hpp"
#include "dpf/net/stream_mesh.hpp"
#include "dpf/net/sync_stream_array.hpp"
#include "dpf/net/tcp_mesh.hpp"
namespace dpf
{
namespace run
{
/// @brief Two parties, two threads, one memory pair. `fn(party, sink)`.
template <typename Fn>
void threads_2(const std::vector<std::size_t> & slots, Fn fn)
{
auto sinks = net::make_memory_sink_pair(1, slots);
std::exception_ptr err;
std::mutex err_mu;
auto note = [&](std::exception_ptr e) {
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::move(e);
};
std::thread t0([&] {
try
{
fn(0u, std::ref(sinks.first));
}
catch (...)
{
note(std::current_exception());
}
});
std::thread t1([&] {
try
{
fn(1u, std::ref(sinks.second));
}
catch (...)
{
note(std::current_exception());
}
});
t0.join();
t1.join();
if (err)
std::rethrow_exception(err);
}
/// @brief `n` parties. `fn(party, clique)` so a party can `end(me, peer)`.
template <typename Fn>
void threads_on_clique(std::size_t n, const std::vector<std::size_t> & slots,
Fn fn)
{
if (n < 2)
throw std::invalid_argument("threads_on_clique needs >= 2");
auto clique = net::make_memory_clique(n, 1, slots);
std::vector<std::thread> ts;
std::mutex err_mu;
std::exception_ptr err;
for (std::size_t p = 0; p < n; ++p)
{
ts.emplace_back([&, p] {
try
{
fn(static_cast<unsigned>(p), clique);
}
catch (...)
{
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::current_exception();
}
});
}
for (auto & t : ts)
t.join();
if (err)
std::rethrow_exception(err);
}
/// @brief `n` parties on a memory clique. `fn(party, mesh)`.
template <typename Fn>
void threads_clique(std::size_t n, const std::vector<std::size_t> & slots, Fn fn)
{
if (n < 2)
throw std::invalid_argument("threads_clique needs >= 2");
auto clique = net::make_memory_clique(n, 1, slots);
std::vector<std::thread> ts;
std::mutex err_mu;
std::exception_ptr err;
for (std::size_t p = 0; p < n; ++p)
{
ts.emplace_back([&, p] {
try
{
auto mesh = clique.mesh_for(p);
fn(static_cast<unsigned>(p), mesh);
}
catch (...)
{
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::current_exception();
}
});
}
for (auto & t : ts)
t.join();
if (err)
std::rethrow_exception(err);
}
/// @brief Sum additive shares on a 3-party ring. Each party sends its share to
/// the other two and adds what it receives.
inline void declassify_ring(net::memory_clique & clique, unsigned me,
const std::uint8_t * mine, std::size_t n, std::uint8_t * sum_out,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
if (clique.roles != 3)
throw std::invalid_argument("declassify_ring is a 3-party open");
for (unsigned peer = 0; peer < 3; ++peer)
{
if (peer == me)
continue;
clique.end(me, peer).submit(0, 0, mine, n);
clique.end(me, peer).flush();
}
std::memcpy(sum_out, mine, n);
for (unsigned peer = 0; peer < 3; ++peer)
{
if (peer == me)
continue;
auto & sink = clique.end(me, peer);
net::wait_peer_ready(sink, 0, 0, budget, "declassify_ring");
std::vector<std::uint8_t> got(n);
sink.read_peer(0, 0, got.data(), n);
for (std::size_t i = 0; i < n; i += 8)
{
std::uint64_t a = 0, b = 0;
const std::size_t k = std::min<std::size_t>(8, n - i);
std::memcpy(&a, sum_out + i, k);
std::memcpy(&b, got.data() + i, k);
a += b;
std::memcpy(sum_out + i, &a, k);
}
}
}
/// @brief Listen or connect, then hand the caller a framed `net::channel`.
/// @details Party 0 binds `port` (0 picks an ephemeral port written back).
/// The other party connects to `host:port`. The link uses the same
/// `peer_security` defaults as `party_session` (TLS 1.3 unless
/// encryption is off). Usable from threads or separate processes.
template <typename Fn>
void tcp_pair(unsigned party, std::string host, std::atomic<unsigned short> & port,
Fn fn, const net::deadlines & lim = {},
const net::peer_security & sec = {}, const net::socket_options & so = {})
{
asio::io_context io;
const unsigned peer = party == 0 ? 1u : 0u;
if (party == 0)
{
asio::ip::tcp::acceptor acc(io);
net::open_listener(acc, port.load());
port.store(acc.local_endpoint().port());
asio::ip::tcp::socket sock(io);
net::accept_until(acc, sock, lim.accept);
auto ch = net::secure_tcp_channel(io, std::move(sock), true, peer, sec, so,
lim.handshake, "p" + std::to_string(peer), "accept");
fn(ch);
}
else
{
const auto deadline = net::setup_clock::now() + lim.join;
while (port.load() == 0)
{
if (net::setup_clock::now() >= deadline)
throw std::runtime_error("tcp_pair: party 0 never published a port");
std::this_thread::sleep_for(std::chrono::milliseconds(1));
}
asio::ip::tcp::socket sock(io);
net::connect_until(sock, host, port.load(), lim.connect);
auto ch = net::secure_tcp_channel(io, std::move(sock), false, peer, sec, so,
lim.handshake, "p" + std::to_string(peer), "connect");
fn(ch);
}
}
/// @brief TCP pair that exposes a `mux_stream_array` with `streams` indexes.
/// @details Joins a 2-party `party_session` (same encryption and wire policy
/// as other mux edges). Party 0 publishes its listen port on `port`;
/// party 1 waits for it. `fn(party, mux_stream_array &)`.
template <typename Fn>
void tcp_pair_mux(unsigned party, std::string host,
std::atomic<unsigned short> & port, std::size_t streams, Fn fn,
const net::session_options & opt = {})
{
if (streams == 0)
throw std::invalid_argument("tcp_pair_mux: streams");
asio::io_context io;
net::session_options so = opt;
so.n_lanes = streams; // `streams` is authoritative; opt sets security/policy
net::party_session sess(io, party, 2, so);
if (party == 0)
port.store(sess.listen(port.load()));
else
{
const auto deadline = net::setup_clock::now() + so.limits.join;
while (port.load() == 0)
{
if (net::setup_clock::now() >= deadline)
throw std::runtime_error(
"tcp_pair_mux: party 0 never published a port");
std::this_thread::sleep_for(std::chrono::milliseconds(1));
}
}
const unsigned short p0 = port.load();
std::vector<net::peer_address> table = {
{host, p0},
{host, party == 0 ? p0 : static_cast<unsigned short>(0)}};
sess.join(table);
net::mux_stream_array mux(sess.peer(1 - party));
fn(party, mux);
}
/// @brief Two threads, each with one end of a memory stream array.
template <typename Fn>
void threads_2_streams(std::size_t streams, Fn fn)
{
auto ends = net::make_memory_stream_pair(streams);
std::exception_ptr err;
std::mutex err_mu;
auto note = [&](std::exception_ptr e) {
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::move(e);
};
std::thread t0([&] {
try
{
fn(0u, std::ref(ends.first));
}
catch (...)
{
note(std::current_exception());
}
});
std::thread t1([&] {
try
{
fn(1u, std::ref(ends.second));
}
catch (...)
{
note(std::current_exception());
}
});
t0.join();
t1.join();
if (err)
std::rethrow_exception(err);
}
/// @brief N-party open on a pairwise stream clique (one stream per edge).
inline void declassify_streams(net::memory_stream_clique & clique, unsigned me,
const std::uint8_t * mine, std::size_t n, std::uint8_t * sum_out)
{
clique.declassify(me, mine, n, sum_out);
}
/// @brief `n` threads join a localhost TCP mux clique; `fn(party, tcp_mesh &)`.
template <typename Fn>
void threads_on_tcp_mesh(unsigned n, std::size_t streams, Fn fn)
{
if (n < 2)
throw std::invalid_argument("threads_on_tcp_mesh needs >= 2");
auto ports = net::make_mesh_ports(n);
std::vector<std::thread> ts;
std::mutex err_mu;
std::exception_ptr err;
for (unsigned p = 0; p < n; ++p)
{
ts.emplace_back([&, p] {
try
{
auto mesh = net::join_tcp_mesh(p, n, "127.0.0.1", ports, streams);
fn(p, mesh);
}
catch (...)
{
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::current_exception();
}
});
}
for (auto & t : ts)
t.join();
if (err)
std::rethrow_exception(err);
}
/// @brief `n` threads join an async TCP mux clique; `fn(party, io, async_tcp_mesh &)`.
template <typename Fn>
void threads_on_async_tcp_mesh(unsigned n, std::size_t streams, Fn fn)
{
if (n < 2)
throw std::invalid_argument("threads_on_async_tcp_mesh needs >= 2");
auto ports = net::make_mesh_ports(n);
std::vector<std::thread> ts;
std::mutex err_mu;
std::exception_ptr err;
for (unsigned p = 0; p < n; ++p)
{
ts.emplace_back([&, p] {
try
{
asio::io_context io;
auto mesh =
net::join_async_tcp_mesh(io, p, n, "127.0.0.1", ports, streams);
fn(p, io, mesh);
}
catch (...)
{
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::current_exception();
}
});
}
for (auto & t : ts)
t.join();
if (err)
std::rethrow_exception(err);
}
} // namespace run
} // namespace dpf
#endif