/// @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 #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/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 void threads_2(const std::vector & 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 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 void threads_on_clique(std::size_t n, const std::vector & 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 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(p), clique); } catch (...) { std::lock_guard 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 void threads_clique(std::size_t n, const std::vector & 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 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(p), mesh); } catch (...) { std::lock_guard 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 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(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 void tcp_pair(unsigned party, std::string host, std::atomic & 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 void tcp_pair_mux(unsigned party, std::string host, std::atomic & 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 table = { {host, p0}, {host, party == 0 ? p0 : static_cast(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 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 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 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 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 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 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 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 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