libdpf/include/dpf/async_protocol.hpp

235 lines
8.2 KiB
C++
Raw Permalink Normal View History

/// @file dpf/async_protocol.hpp
/// @brief Real overlapped, event-driven byte-round protocol runner.
/// @details `overlapped_byte_protocol` replaces the blocking loop of
/// `factory::async_byte_protocol` with true asynchronous I/O over an
/// `async_stream_array`. Each round: read the blind from the dealer
/// (async, skipped when `blind_bytes == 0`), `produce` the outbound
/// message, then overlap the peer write and peer read as one
/// `async_exchange`, and on completion `finish`, fire the round
/// callback, and continue to the next round or the done handler.
/// Nothing spins on a `peer_ready` flag and nothing blocks the calling
/// thread — every continuation runs on an `io_context` thread.
#ifndef LIBDPF_INCLUDE_DPF_ASYNC_PROTOCOL_HPP__
#define LIBDPF_INCLUDE_DPF_ASYNC_PROTOCOL_HPP__
#include <cstddef>
#include <cstdint>
#include <functional>
#include <memory>
#include <mutex>
#include <stdexcept>
#include <system_error>
#include <utility>
#include <vector>
#include "dpf/net/asio_ns.hpp"
#include "dpf/net/async_stream_array.hpp"
#include "dpf/protocol_factory.hpp" // factory::async_byte_round
namespace dpf
{
namespace async
{
/// @brief Overlap a write and a read on stream `i`; fire `done` once both end.
/// @details The two operations are issued concurrently (full duplex): the
/// handler runs after both complete, carrying the first error seen.
/// `out` and `in` must stay valid until `done` fires.
inline void async_exchange(net::async_stream_array & s, std::size_t i,
const void * out, std::size_t out_n, void * in, std::size_t in_n,
net::async_handler done)
{
struct state
{
net::async_handler done;
int remaining = 2;
std::error_code ec;
std::mutex mu;
};
auto st = std::make_shared<state>();
st->done = std::move(done);
auto complete = [st](const std::error_code & ec) {
bool fire = false;
std::error_code final_ec;
{
std::lock_guard<std::mutex> lock(st->mu);
if (ec && !st->ec)
st->ec = ec;
if (--st->remaining == 0)
fire = true;
final_ec = st->ec;
}
if (fire && st->done)
st->done(final_ec);
};
s.async_write(i, out, out_n, complete);
s.async_read(i, in, in_n, complete);
}
/// @brief Fully overlapped runner for a list of `async_byte_round`s.
/// @details Constructed per party per run and driven through a `shared_ptr`
/// (kept alive by its own continuations). `start` kicks round 0.
class overlapped_byte_protocol
: public std::enable_shared_from_this<overlapped_byte_protocol>
{
public:
/// @param ec error (empty on success); @param state final protocol state.
using done_handler =
std::function<void(const std::error_code & ec,
std::vector<std::uint8_t> state)>;
using on_round_complete_fn = std::function<void(std::size_t round)>;
/// @param peer overlapped peer array; lane = `r % peer.size()` so a small
/// pool can carry many sequential rounds (known-size exchange).
/// @param dealer optional blind source (lane `r % dealer.size()`); may be
/// null when every round has `blind_bytes == 0`.
overlapped_byte_protocol(net::async_stream_array & peer,
net::async_stream_array * dealer,
std::vector<factory::async_byte_round> rounds,
on_round_complete_fn on_round_complete = {})
: peer_(&peer),
dealer_(dealer),
rounds_(std::move(rounds)),
on_round_complete_(std::move(on_round_complete))
{
if (peer_->size() == 0)
throw std::invalid_argument("overlapped_byte_protocol: empty peer");
for (const auto & r : rounds_)
{
if (!r.produce || !r.finish)
throw std::invalid_argument("overlapped_byte_protocol: missing callbacks");
if (r.blind_bytes != 0 && dealer_ == nullptr)
throw std::invalid_argument("overlapped_byte_protocol: blind needs a dealer");
}
if (dealer_ != nullptr && dealer_->size() == 0)
throw std::invalid_argument("overlapped_byte_protocol: empty dealer");
}
std::size_t rounds() const noexcept { return rounds_.size(); }
/// @brief Begin the protocol. `done` fires once, after the last round.
void start(std::size_t index, std::vector<std::uint8_t> state,
done_handler done)
{
(void)index;
state_ = std::move(state);
done_ = std::move(done);
run_round(0);
}
private:
void finish_with(const std::error_code & ec)
{
if (done_)
{
auto d = std::move(done_);
done_ = nullptr;
d(ec, std::move(state_));
}
}
void run_round(std::size_t r)
{
if (r >= rounds_.size())
{
finish_with(std::error_code{});
return;
}
const auto & spec = rounds_[r];
blind_.assign(spec.blind_bytes, 0);
auto self = shared_from_this();
if (spec.blind_bytes != 0)
{
const std::size_t dlane = r % dealer_->size();
dealer_->async_read(dlane, blind_.data(), blind_.size(),
[self, r](const std::error_code & ec) {
self->after_blind(r, ec);
});
}
else
{
// No blind: hop through the io_context so we never recurse deeply.
asio::post(peer_->context(),
[self, r]() { self->after_blind(r, std::error_code{}); });
}
}
void after_blind(std::size_t r, const std::error_code & ec)
{
if (ec)
{
finish_with(ec);
return;
}
const auto & spec = rounds_[r];
outbound_ = spec.produce(state_, blind_.data(), blind_.size());
if (outbound_.size() != spec.msg_bytes)
{
finish_with(std::make_error_code(std::errc::message_size));
return;
}
inbound_.assign(spec.msg_bytes, 0);
auto self = shared_from_this();
const std::size_t lane = r % peer_->size();
async_exchange(*peer_, lane, outbound_.data(), outbound_.size(),
inbound_.data(), inbound_.size(),
[self, r](const std::error_code & xec) {
self->after_exchange(r, xec);
});
}
void after_exchange(std::size_t r, const std::error_code & ec)
{
if (ec)
{
finish_with(ec);
return;
}
const auto & spec = rounds_[r];
spec.finish(state_, inbound_.data(), inbound_.size(), blind_.data(),
blind_.size());
if (on_round_complete_)
on_round_complete_(r);
run_round(r + 1);
}
net::async_stream_array * peer_ = nullptr;
net::async_stream_array * dealer_ = nullptr;
std::vector<factory::async_byte_round> rounds_;
on_round_complete_fn on_round_complete_;
std::vector<std::uint8_t> state_;
done_handler done_;
std::vector<std::uint8_t> blind_;
std::vector<std::uint8_t> outbound_;
std::vector<std::uint8_t> inbound_;
};
/// @brief Make an `overlapped_byte_protocol` as a `shared_ptr` (required for
/// the self-owning continuation chain).
HEDLEY_WARN_UNUSED_RESULT
inline std::shared_ptr<overlapped_byte_protocol> make_overlapped_byte_protocol(
net::async_stream_array & peer, net::async_stream_array * dealer,
std::vector<factory::async_byte_round> rounds,
overlapped_byte_protocol::on_round_complete_fn on_round_complete = {})
{
return std::make_shared<overlapped_byte_protocol>(peer, dealer,
std::move(rounds), std::move(on_round_complete));
}
/// @brief Post `start_all`, then run the io_context until all work drains.
/// @details The single entry point for driving overlapped parties: everything
/// the parties initiate is chased to completion by `io.run()`.
template <typename Fn>
void run_overlapped(asio::io_context & io, Fn start_all)
{
asio::post(io, [start_all = std::move(start_all)]() mutable { start_all(); });
io.run();
}
} // namespace async
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_ASYNC_PROTOCOL_HPP__