libdpf/include/dpf/online_session.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

573 lines
20 KiB
C++

/// @file dpf/online_session.hpp
/// @brief Prep delivery and application sessions over any async link.
/// @details Every function here comes in two layers: a per-party function
/// that takes the links it should use (memory, unix, TCP, SCTP,
/// `party_session::dealer()`), and an in-process wrapper that builds
/// links from `run_config` and runs every party on a thread.
/// Prep travels as `u32 length || view` on lane 0.
#ifndef LIBDPF_INCLUDE_DPF_ONLINE_SESSION_HPP__
#define LIBDPF_INCLUDE_DPF_ONLINE_SESSION_HPP__
#include <atomic>
#include <chrono>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <exception>
#include <functional>
#include <memory>
#include <mutex>
#include <stdexcept>
#include <string>
#include <thread>
#include <utility>
#include <vector>
#include "dpf/net/asio_ns.hpp"
#include "dpf/app_plans.hpp"
#include "dpf/mesh_apps.hpp"
#include "dpf/net/async_round_sink.hpp"
#include "dpf/net/async_stream_array.hpp"
#include "dpf/net/edge_mesh.hpp"
#include "dpf/net/party_session.hpp"
#include "dpf/prep_source.hpp"
#include "dpf/protocol.hpp"
#include "dpf/run_config.hpp"
namespace dpf
{
namespace session
{
struct shipped_prep
{
prep::cursor party0;
prep::cursor party1;
std::size_t bytes0 = 0;
std::size_t bytes1 = 0;
};
namespace detail
{
/// @brief Pump `io` until `done` or `budget` passes.
inline void pump_until(asio::io_context & io, const std::atomic<bool> & done,
std::chrono::milliseconds budget, const char * what)
{
const auto deadline = std::chrono::steady_clock::now() + budget;
while (!done.load(std::memory_order_acquire))
{
const auto now = std::chrono::steady_clock::now();
if (now >= deadline)
throw std::system_error(std::make_error_code(std::errc::timed_out),
std::string(what) + ": timed out after "
+ std::to_string(budget.count()) + " ms");
if (io.stopped())
io.restart();
io.run_one_for(std::min<std::chrono::milliseconds>(
std::chrono::duration_cast<std::chrono::milliseconds>(deadline - now),
std::chrono::milliseconds(50)));
}
}
inline void run_session(protocol::schedule_session & sess,
std::chrono::milliseconds budget, const char * what)
{
protocol::drive_options opt;
opt.wait_timeout = budget;
protocol::detail::run_session(sess, opt, what);
}
} // namespace detail
/// @brief Send one prep view on lane 0 of `link` (`u32 length || bytes`).
inline void send_prep(net::async_stream_array & link,
const std::vector<std::uint8_t> & view,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
auto frame = net::acquire_buffer(4 + view.size());
net::detail::put_u32(frame->data(), static_cast<std::uint32_t>(view.size()));
if (!view.empty())
std::memcpy(frame->data() + 4, view.data(), view.size());
auto done = std::make_shared<std::atomic<bool>>(false);
auto err = std::make_shared<std::error_code>();
link.async_write_owned(0, std::move(frame),
[done, err](const std::error_code & ec) {
*err = ec;
done->store(true, std::memory_order_release);
});
detail::pump_until(link.context(), *done, budget, "send_prep");
if (*err)
throw std::system_error(*err, "send_prep");
}
/// @brief Receive one prep view from lane 0 of `link`.
inline std::vector<std::uint8_t> receive_prep(net::async_stream_array & link,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
auto lenb = std::make_shared<std::array<std::uint8_t, 4>>();
auto done = std::make_shared<std::atomic<bool>>(false);
auto err = std::make_shared<std::error_code>();
link.async_read(0, lenb->data(), 4, [done, err, lenb](const std::error_code & ec) {
*err = ec;
done->store(true, std::memory_order_release);
});
detail::pump_until(link.context(), *done, budget, "receive_prep length");
if (*err)
throw std::system_error(*err, "receive_prep length");
auto body = std::make_shared<std::vector<std::uint8_t>>(
net::detail::get_u32(lenb->data()));
if (body->empty())
return {};
done->store(false);
link.async_read(0, body->data(), body->size(),
[done, err, body](const std::error_code & ec) {
*err = ec;
done->store(true, std::memory_order_release);
});
detail::pump_until(link.context(), *done, budget, "receive_prep body");
if (*err)
throw std::system_error(*err, "receive_prep body");
return *body;
}
/// @brief Dealer role: deal `d` and send each party its view.
inline std::pair<std::size_t, std::size_t> deal_and_send(
net::async_stream_array & to_p0, net::async_stream_array & to_p1,
const prep::demand & d,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
auto views = prep::deal_views(d);
send_prep(to_p0, views.first, budget);
send_prep(to_p1, views.second, budget);
return {views.first.size(), views.second.size()};
}
/// @brief Deal `d` and ship both views over links built from `cfg`.
/// @details Socket transports use a `dealer_session` that both parties
/// connect to with `party_session::connect_dealer`; in-process
/// transports use memory pairs.
inline shipped_prep ship_prep(const prep::demand & d,
const app::run_config & cfg = {})
{
std::exception_ptr err;
std::mutex err_mu;
auto note = [&] {
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::current_exception();
};
std::vector<std::uint8_t> got0, got1;
const auto budget = cfg.wait_timeout.count() != 0 ? cfg.wait_timeout
: std::chrono::milliseconds(30000);
if (net::is_socket_transport(cfg.kind) && cfg.kind != net::transport::local)
{
asio::io_context io_d, io0, io1;
net::session_options so;
so.policy = cfg.policy;
so.limits = cfg.limits;
so.security = cfg.security;
std::atomic<unsigned short> port{0};
std::thread dealer([&] {
try
{
net::dealer_session ds(io_d, 2, so);
port.store(ds.listen());
ds.accept_parties();
deal_and_send(ds.party(0), ds.party(1), d, budget);
}
catch (...)
{
note();
}
});
auto party = [&](unsigned me, asio::io_context & io,
std::vector<std::uint8_t> & out) {
try
{
const auto deadline = std::chrono::steady_clock::now() + cfg.limits.join;
while (port.load() == 0)
{
if (std::chrono::steady_clock::now() > deadline)
throw std::runtime_error("ship_prep: dealer never listened");
std::this_thread::sleep_for(std::chrono::milliseconds(1));
}
net::party_session ps(io, me, 2, so);
ps.connect_dealer(cfg.host, port.load());
out = receive_prep(ps.dealer(), budget);
}
catch (...)
{
note();
}
};
std::thread p0([&] { party(0, io0, got0); });
std::thread p1([&] { party(1, io1, got1); });
dealer.join();
p0.join();
p1.join();
}
else
{
asio::io_context io_d, io0, io1;
auto link0 = net::make_async_dual_memory_stream_pair(io_d, io0, 1);
auto link1 = net::make_async_dual_memory_stream_pair(io_d, io1, 1);
std::thread dealer([&] {
try
{
deal_and_send(link0.first, link1.first, d, budget);
}
catch (...)
{
note();
}
});
std::thread p0([&] {
try
{
got0 = receive_prep(link0.second, budget);
}
catch (...)
{
note();
}
});
std::thread p1([&] {
try
{
got1 = receive_prep(link1.second, budget);
}
catch (...)
{
note();
}
});
dealer.join();
p0.join();
p1.join();
}
if (err)
std::rethrow_exception(err);
shipped_prep out;
out.bytes0 = got0.size();
out.bytes1 = got1.size();
out.party0 = prep::cursor(std::move(got0));
out.party1 = prep::cursor(std::move(got1));
return out;
}
/// @brief One party of hushmap ADD: pad tape on `pad_link`, opens on `peer_link`.
inline void drive_hushmap_add_party(std::size_t layers,
net::async_stream_array & peer_link, net::async_stream_array & pad_link,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
auto tape = std::make_shared<std::vector<std::uint8_t>>();
auto rounds = protocol::hushmap_add_schedule(layers, tape);
net::async_round_sink peer_sink(peer_link, {8u, 8u}, 1);
net::async_round_sink pad_sink(pad_link, std::vector<std::size_t>(layers, 24u), 1);
protocol::edge_mesh mesh;
mesh.sinks = {&peer_sink, nullptr, &pad_sink};
protocol::schedule_session sess(1, mesh, std::move(rounds), false);
sess.submit(0);
detail::run_session(sess, budget, "hushmap_add");
}
/// @brief Both hushmap parties in one process over links from `cfg`.
inline void drive_hushmap_add(std::size_t layers, const app::run_config & cfg = {})
{
if (net::is_socket_transport(cfg.kind) && cfg.kind != net::transport::local)
{
auto ports = net::make_mesh_ports(2);
std::exception_ptr err;
std::mutex mu;
auto run = [&](unsigned me) {
try
{
asio::io_context io;
net::session_options so;
so.n_lanes = std::max<std::size_t>(2, layers) + 2;
so.kind = cfg.kind;
so.policy = cfg.policy;
so.limits = cfg.limits;
so.security = cfg.security;
net::party_session ps(io, me, 2, so);
ps.join(cfg.host, ports);
net::async_stream_view peer(ps.peer(1 - me), 0, 2);
net::async_stream_view pad(ps.peer(1 - me), 2,
std::max<std::size_t>(1, layers));
drive_hushmap_add_party(layers, peer, pad, cfg.wait_timeout);
}
catch (...)
{
std::lock_guard<std::mutex> lock(mu);
if (!err)
err = std::current_exception();
}
};
std::thread t0([&] { run(0); });
std::thread t1([&] { run(1); });
t0.join();
t1.join();
if (err)
std::rethrow_exception(err);
return;
}
asio::io_context io0, io1;
auto peer = net::make_async_dual_memory_stream_pair(io0, io1, 2);
auto pad = net::make_async_dual_memory_stream_pair(io0, io1,
std::max<std::size_t>(1, layers));
std::exception_ptr err;
std::mutex mu;
auto run = [&](net::async_stream_array & pl, net::async_stream_array & dl) {
try
{
drive_hushmap_add_party(layers, pl, dl, cfg.wait_timeout);
}
catch (...)
{
std::lock_guard<std::mutex> lock(mu);
if (!err)
err = std::current_exception();
pl.close();
dl.close();
}
};
std::thread t0([&] { run(peer.first, pad.first); });
std::thread t1([&] { run(peer.second, pad.second); });
t0.join();
t1.join();
if (err)
std::rethrow_exception(err);
}
/// @brief Star client over `links[i]` (edge `i` → server `i`).
inline void drive_star_client(const std::vector<net::async_stream_array *> & links,
const std::vector<std::size_t> & slots,
std::vector<protocol::schedule_round> rounds,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000),
const net::sink_options & so = {})
{
std::vector<std::unique_ptr<net::async_round_sink>> sinks;
protocol::edge_mesh mesh;
mesh.sinks.resize(links.size());
for (std::size_t i = 0; i < links.size(); ++i)
{
sinks.push_back(
std::make_unique<net::async_round_sink>(*links[i], slots, 1, so));
mesh.sinks[i] = sinks.back().get();
}
protocol::schedule_session sess(1, mesh, std::move(rounds), false);
sess.submit(0);
detail::run_session(sess, budget, "star client");
}
/// @brief One star server over `link`.
inline void drive_star_server(net::async_stream_array & link,
const std::vector<std::size_t> & slots,
std::vector<protocol::schedule_round> rounds,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000),
const net::sink_options & so = {})
{
net::async_round_sink sink(link, slots, 1, so);
protocol::edge_mesh mesh;
mesh.sinks = {&sink};
protocol::schedule_session sess(1, mesh, std::move(rounds), false);
sess.submit(0);
detail::run_session(sess, budget, "star server");
}
/// @brief Client and `n_servers` servers in one process on split-io memory.
inline void drive_async_star(std::size_t n_servers,
const std::vector<std::size_t> & slots,
std::vector<protocol::schedule_round> client_rounds,
const std::function<std::vector<protocol::schedule_round>(std::size_t)> &
server_rounds_for,
std::chrono::milliseconds budget = std::chrono::milliseconds(30000))
{
if (n_servers < 1)
throw std::invalid_argument("drive_async_star servers");
asio::io_context io_client;
std::vector<asio::io_context> io_server(n_servers);
std::vector<std::pair<net::async_dual_memory_stream_array,
net::async_dual_memory_stream_array>>
links;
links.reserve(n_servers);
for (std::size_t i = 0; i < n_servers; ++i)
links.push_back(net::make_async_dual_memory_stream_pair(io_client,
io_server[i], slots.size()));
std::exception_ptr err;
std::mutex err_mu;
auto note = [&] {
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::current_exception();
};
std::thread client([&] {
std::vector<net::async_stream_array *> ls;
for (auto & l : links)
ls.push_back(&l.first);
try
{
drive_star_client(ls, slots, std::move(client_rounds), budget);
}
catch (...)
{
note();
for (auto * l : ls)
l->close();
}
});
std::vector<std::thread> servers;
for (std::size_t i = 0; i < n_servers; ++i)
{
servers.emplace_back([&, i] {
try
{
drive_star_server(links[i].second, slots, server_rounds_for(i),
budget);
}
catch (...)
{
note();
links[i].second.close();
}
});
}
client.join();
for (auto & t : servers)
t.join();
if (err)
std::rethrow_exception(err);
}
/// @brief Client and `n_servers` servers, one thread each, over `cfg.kind`.
/// @details Every client-server edge is its own link and servers never connect
/// to each other: split-io memory for `async`, unix sockets for
/// `local`, and a two-party `party_session` per edge for `mux`,
/// `parallel`, and `sctp`. Lanes, framing, wire policy, deadlines,
/// and the per-round wait budget come from `cfg`.
inline void drive_async_star(std::size_t n_servers,
const std::vector<std::size_t> & slots,
std::vector<protocol::schedule_round> client_rounds,
const std::function<std::vector<protocol::schedule_round>(std::size_t)> &
server_rounds_for,
const app::run_config & cfg)
{
if (n_servers < 1)
throw std::invalid_argument("drive_async_star servers");
const bool sockets = cfg.kind == net::transport::mux
|| cfg.kind == net::transport::parallel || cfg.kind == net::transport::sctp;
if (!sockets && cfg.kind != net::transport::async_memory
&& cfg.kind != net::transport::local)
throw std::invalid_argument(std::string("drive_async_star: transport ")
+ net::transport_name(cfg.kind)
+ " is not an async link (use async|local|mux|parallel|sctp)");
const std::size_t lanes = net::lane_count_for_rounds(slots.size(), cfg.n_lanes);
net::sink_options so;
so.framing = cfg.framing;
const auto budget = cfg.wait_timeout;
asio::io_context io_client;
std::vector<asio::io_context> io_server(n_servers);
std::vector<std::unique_ptr<net::async_stream_array>> client_end(n_servers);
std::vector<std::unique_ptr<net::async_stream_array>> server_end(n_servers);
for (std::size_t i = 0; i < n_servers && !sockets; ++i)
{
if (cfg.kind == net::transport::async_memory)
{
auto pr = net::make_async_dual_memory_stream_pair(io_client, io_server[i],
lanes, cfg.policy.window_bytes);
client_end[i] =
std::make_unique<net::async_dual_memory_stream_array>(std::move(pr.first));
server_end[i] =
std::make_unique<net::async_dual_memory_stream_array>(std::move(pr.second));
}
else
{
auto pr = net::make_local_socket_pairs(io_client, io_server[i], lanes);
client_end[i] = std::make_unique<net::async_local_parallel_stream_array>(
io_client, std::move(pr.first), cfg.policy);
server_end[i] = std::make_unique<net::async_local_parallel_stream_array>(
io_server[i], std::move(pr.second), cfg.policy);
}
}
net::session_options sopt;
sopt.n_lanes = lanes;
sopt.kind = cfg.kind;
sopt.policy = cfg.policy;
sopt.limits = cfg.limits;
sopt.security = cfg.security;
std::vector<net::mesh_ports> ports;
ports.reserve(n_servers);
for (std::size_t i = 0; i < n_servers; ++i)
ports.push_back(net::make_mesh_ports(2));
std::exception_ptr err;
std::mutex err_mu;
auto note = [&] {
std::lock_guard<std::mutex> lock(err_mu);
if (!err)
err = std::current_exception();
};
std::thread client([&] {
std::vector<std::unique_ptr<net::party_session>> sess;
std::vector<net::async_stream_array *> ls;
try
{
for (std::size_t i = 0; i < n_servers; ++i)
{
if (sockets)
{
sess.push_back(
std::make_unique<net::party_session>(io_client, 0, 2, sopt));
sess.back()->join(cfg.host, ports[i]);
ls.push_back(&sess.back()->peer(1));
}
else
ls.push_back(client_end[i].get());
}
drive_star_client(ls, slots, std::move(client_rounds), budget, so);
}
catch (...)
{
note();
for (auto * l : ls)
l->close();
}
});
std::vector<std::thread> servers;
for (std::size_t i = 0; i < n_servers; ++i)
{
servers.emplace_back([&, i] {
std::unique_ptr<net::party_session> s;
net::async_stream_array * link = server_end[i].get();
try
{
if (sockets)
{
s = std::make_unique<net::party_session>(io_server[i], 1, 2, sopt);
s->join(cfg.host, ports[i]);
link = &s->peer(0);
}
drive_star_server(*link, slots, server_rounds_for(i), budget, so);
}
catch (...)
{
note();
if (link != nullptr)
link->close();
}
});
}
client.join();
for (auto & t : servers)
t.join();
if (err)
std::rethrow_exception(err);
}
} // namespace session
} // namespace dpf
#endif