libdpf/include/dpf/protocol_factory.hpp

757 lines
24 KiB
C++
Raw Normal View History

/// @file dpf/protocol_factory.hpp
/// @brief Dealer and online factories over `net::stream_array`.
/// @details A dealer functor writes one view per party on stream `i`. An online
/// factory reads blinds from a dealer array and swaps messages on a
/// peer array. Both arrays are the same type so the backend can change
/// without rewriting the protocol.
#ifndef LIBDPF_INCLUDE_DPF_PROTOCOL_FACTORY_HPP__
#define LIBDPF_INCLUDE_DPF_PROTOCOL_FACTORY_HPP__
#include <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <functional>
#include <stdexcept>
#include <tuple>
#include <type_traits>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/net/asio_ns.hpp"
#include "dpf/edabit.hpp"
#include "dpf/net/stream_array.hpp"
#include "dpf/ot_pack.hpp"
#include "dpf/random.hpp"
namespace dpf
{
namespace factory
{
namespace detail
{
template <typename T>
void write_pod(net::stream_array & a, std::size_t i, const T & v)
{
static_assert(std::is_trivially_copyable_v<T>,
"factory view must be trivially copyable");
a.write(i, &v, sizeof(T));
}
template <typename T>
T read_pod(net::stream_array & a, std::size_t i)
{
static_assert(std::is_trivially_copyable_v<T>,
"factory view must be trivially copyable");
T v{};
a.read(i, &v, sizeof(T));
return v;
}
template <typename T>
void exchange_pod(net::stream_array & peer, std::size_t i, const T & mine,
T & theirs)
{
static_assert(std::is_trivially_copyable_v<T>,
"factory message must be trivially copyable");
peer.write(i, &mine, sizeof(T));
peer.flush(i);
peer.read(i, &theirs, sizeof(T));
}
} // namespace detail
// ---------------------------------------------------------------------------
// Dealer functors (prep atoms)
// ---------------------------------------------------------------------------
struct ring_triple_view
{
std::uint8_t a[8]{};
std::uint8_t b[8]{};
std::uint8_t c[8]{};
std::uint16_t limb = 8;
};
struct dabit_view
{
std::uint8_t bit = 0;
std::uint8_t arith[8]{};
std::uint16_t limb = 8;
};
struct bit_triple_view
{
std::uint8_t a = 0;
std::uint8_t b = 0;
std::uint8_t c = 0;
};
inline std::pair<ring_triple_view, ring_triple_view>
deal_ring_triple(std::uint16_t limb)
{
if (limb == 0 || limb > 8)
throw std::invalid_argument("deal_ring_triple limb");
const std::uint64_t mask = limb >= 8
? ~std::uint64_t{0}
: ((std::uint64_t{1} << (8u * limb)) - 1u);
const auto a = dpf::uniform_sample<std::uint64_t>() & mask;
const auto b = dpf::uniform_sample<std::uint64_t>() & mask;
const auto c = (a * b) & mask;
const auto a0 = dpf::uniform_sample<std::uint64_t>() & mask;
const auto b0 = dpf::uniform_sample<std::uint64_t>() & mask;
const auto c0 = dpf::uniform_sample<std::uint64_t>() & mask;
ring_triple_view v0{}, v1{};
v0.limb = limb;
v1.limb = limb;
std::memcpy(v0.a, &a0, limb);
std::memcpy(v0.b, &b0, limb);
std::memcpy(v0.c, &c0, limb);
const auto a1 = (a - a0) & mask;
const auto b1 = (b - b0) & mask;
const auto c1 = (c - c0) & mask;
std::memcpy(v1.a, &a1, limb);
std::memcpy(v1.b, &b1, limb);
std::memcpy(v1.c, &c1, limb);
return {v0, v1};
}
inline std::pair<dabit_view, dabit_view> deal_dabit(std::uint16_t limb)
{
if (limb == 0 || limb > 8)
throw std::invalid_argument("deal_dabit limb");
auto pair = ot::sample_dabit_pair<std::uint64_t>();
dabit_view v0{}, v1{};
v0.limb = limb;
v1.limb = limb;
v0.bit = pair.p0.bit;
v1.bit = pair.p1.bit;
std::memcpy(v0.arith, &pair.p0.arith, limb);
std::memcpy(v1.arith, &pair.p1.arith, limb);
return {v0, v1};
}
inline std::pair<bit_triple_view, bit_triple_view> deal_bit_triple()
{
auto t = ot::sample_bit_triple_pair();
return {{t.p0.a, t.p0.b, t.p0.c}, {t.p1.a, t.p1.b, t.p1.c}};
}
// ---------------------------------------------------------------------------
// make_dealer
// ---------------------------------------------------------------------------
/// @brief Run each functor `count` times; write view0/view1 on stream `i`.
template <typename... Fs>
void make_dealer(net::stream_array & party0, net::stream_array & party1,
std::size_t count, Fs &&... fs)
{
if (party0.size() < sizeof...(Fs) || party1.size() < sizeof...(Fs))
throw std::invalid_argument("make_dealer: not enough streams");
std::size_t stream = 0;
auto run_one = [&](auto && f) {
for (std::size_t j = 0; j < count; ++j)
{
auto views = f();
using V0 = std::decay_t<decltype(views.first)>;
using V1 = std::decay_t<decltype(views.second)>;
static_assert(std::is_trivially_copyable_v<V0>
&& std::is_trivially_copyable_v<V1>,
"make_dealer views must be trivially copyable");
if (sizeof(V0) != sizeof(V1))
throw std::logic_error("make_dealer: unequal view sizes");
detail::write_pod(party0, stream, views.first);
detail::write_pod(party1, stream, views.second);
}
party0.flush(stream);
party1.flush(stream);
++stream;
};
(run_one(std::forward<Fs>(fs)), ...);
}
/// @brief Like `make_dealer`, but the functor receives PRG lanes and the index.
template <typename Prg0, typename Prg1, typename F>
void make_dealer_prg(net::stream_array & party0, net::stream_array & party1,
std::size_t count, Prg0 & prg0, Prg1 & prg1, F && f,
std::size_t stream = 0)
{
if (party0.size() <= stream || party1.size() <= stream)
throw std::invalid_argument("make_dealer_prg: stream index");
for (std::size_t j = 0; j < count; ++j)
{
auto views = f(prg0, prg1, j);
using V0 = std::decay_t<decltype(views.first)>;
using V1 = std::decay_t<decltype(views.second)>;
static_assert(std::is_trivially_copyable_v<V0>
&& std::is_trivially_copyable_v<V1>,
"make_dealer_prg views must be trivially copyable");
if (sizeof(V0) != sizeof(V1))
throw std::logic_error("make_dealer_prg: unequal view sizes");
detail::write_pod(party0, stream, views.first);
detail::write_pod(party1, stream, views.second);
}
party0.flush(stream);
party1.flush(stream);
}
/// @brief Three-party dealer. Functors return `std::array<View, 3>`.
template <typename... Fs>
void make_dealer(net::stream_array & party0, net::stream_array & party1,
net::stream_array & party2, std::size_t count, Fs &&... fs)
{
if (party0.size() < sizeof...(Fs) || party1.size() < sizeof...(Fs)
|| party2.size() < sizeof...(Fs))
throw std::invalid_argument("make_dealer 3p: not enough streams");
std::size_t stream = 0;
auto run_one = [&](auto && f) {
for (std::size_t j = 0; j < count; ++j)
{
auto views = f();
using Arr = std::decay_t<decltype(views)>;
static_assert(std::tuple_size_v<Arr> == 3,
"3-party dealer functor must return array/tuple of 3");
using V = std::decay_t<decltype(std::get<0>(views))>;
static_assert(std::is_trivially_copyable_v<V>,
"make_dealer views must be trivially copyable");
detail::write_pod(party0, stream, std::get<0>(views));
detail::write_pod(party1, stream, std::get<1>(views));
detail::write_pod(party2, stream, std::get<2>(views));
}
party0.flush(stream);
party1.flush(stream);
party2.flush(stream);
++stream;
};
(run_one(std::forward<Fs>(fs)), ...);
}
// ---------------------------------------------------------------------------
// Online protocol (blocking round list)
// ---------------------------------------------------------------------------
struct gmw_and_blind
{
std::uint8_t a = 0;
std::uint8_t b = 0;
std::uint8_t c = 0;
};
struct gmw_and_msg
{
std::uint8_t d = 0;
std::uint8_t e = 0;
};
struct gmw_and_keep
{
gmw_and_blind blind{};
gmw_and_msg mine{};
};
/// @brief Round 0: `(p,q), blind -> {msg, keep}`.
struct gmw_and_round0
{
std::pair<gmw_and_msg, gmw_and_keep> operator()(
std::pair<std::uint8_t, std::uint8_t> pq, gmw_and_blind blind) const
{
gmw_and_msg m{};
m.d = static_cast<std::uint8_t>((pq.first ^ blind.a) & 1u);
m.e = static_cast<std::uint8_t>((pq.second ^ blind.b) & 1u);
return {m, gmw_and_keep{blind, m}};
}
};
/// @brief Round 1: finish AND from peer masks.
struct gmw_and_round1
{
unsigned party = 0;
std::uint8_t operator()(gmw_and_keep keep, gmw_and_msg peer,
gmw_and_blind /*unused_blind*/, gmw_and_blind /*corr*/) const
{
const std::uint8_t d_open =
static_cast<std::uint8_t>(keep.mine.d ^ peer.d);
const std::uint8_t e_open =
static_cast<std::uint8_t>(keep.mine.e ^ peer.e);
ot::bit_triple t{keep.blind.a, keep.blind.b, keep.blind.c};
return edabit::and_finish(t, d_open, e_open, party);
}
};
/// @brief Two-round online object: round0 produces a peer message; round1
/// consumes the peer reply. Blinds come from `dealer[0]` (round0 only
/// for GMW AND).
template <typename Round0, typename Round1>
class two_round_protocol
{
public:
two_round_protocol(std::size_t count, net::stream_array & peer,
net::stream_array & dealer, Round0 r0, Round1 r1)
: count_(count),
peer_(&peer),
dealer_(&dealer),
r0_(std::move(r0)),
r1_(std::move(r1))
{
if (peer_->size() < 1 || dealer_->size() < 1)
throw std::invalid_argument("two_round_protocol: empty arrays");
}
template <typename Input>
auto operator()(std::size_t index, const Input & input)
{
if (index >= count_)
throw std::out_of_range("protocol index");
using Blind = gmw_and_blind;
Blind blind = detail::read_pod<Blind>(*dealer_, 0);
auto step0 = r0_(input, blind);
using Msg = std::decay_t<decltype(step0.first)>;
Msg peer_msg{};
detail::exchange_pod(*peer_, 0, step0.first, peer_msg);
Blind unused{};
return r1_(step0.second, peer_msg, unused, unused);
}
private:
std::size_t count_ = 0;
net::stream_array * peer_ = nullptr;
net::stream_array * dealer_ = nullptr;
Round0 r0_;
Round1 r1_;
};
template <typename Round0, typename Round1>
class protocol_factory
{
public:
protocol_factory(Round0 r0, Round1 r1)
: r0_(std::move(r0)), r1_(std::move(r1))
{
}
two_round_protocol<Round0, Round1> create(std::size_t count,
net::stream_array & peer, net::stream_array & dealer) const
{
return two_round_protocol<Round0, Round1>(count, peer, dealer, r0_,
r1_);
}
private:
Round0 r0_;
Round1 r1_;
};
template <typename Round0, typename Round1>
HEDLEY_WARN_UNUSED_RESULT
protocol_factory<std::decay_t<Round0>, std::decay_t<Round1>>
make_protocol_factory(Round0 && r0, Round1 && r1)
{
return protocol_factory<std::decay_t<Round0>, std::decay_t<Round1>>(
std::forward<Round0>(r0), std::forward<Round1>(r1));
}
/// @brief Dealer functor writing `gmw_and_blind` shares on stream 0.
inline auto make_gmw_and_dealer_functor()
{
return [] {
auto t = deal_bit_triple();
gmw_and_blind v0{t.first.a, t.first.b, t.first.c};
gmw_and_blind v1{t.second.a, t.second.b, t.second.c};
return std::pair{v0, v1};
};
}
// ---------------------------------------------------------------------------
// Byte-oriented N-round online runner
// ---------------------------------------------------------------------------
struct round_spec
{
std::size_t blind_bytes = 0;
std::size_t msg_bytes = 0;
};
struct byte_round
{
std::size_t blind_bytes = 0;
std::size_t msg_bytes = 0;
std::function<std::vector<std::uint8_t>(std::vector<std::uint8_t> & state,
const std::uint8_t * blind, std::size_t nb)>
produce;
std::function<void(std::vector<std::uint8_t> & state, const std::uint8_t * peer,
std::size_t nm, const std::uint8_t * blind, std::size_t nb)>
finish;
};
/// @brief Run `rounds` in order: blind from `dealer[r]`, message on `peer[r]`.
class byte_protocol
{
public:
byte_protocol(std::size_t count, net::stream_array & peer,
net::stream_array & dealer, std::vector<byte_round> rounds)
: count_(count),
peer_(&peer),
dealer_(&dealer),
rounds_(std::move(rounds))
{
if (peer_->size() < rounds_.size() || dealer_->size() < rounds_.size())
throw std::invalid_argument("byte_protocol: not enough streams");
for (const auto & r : rounds_)
{
if (!r.produce || !r.finish)
throw std::invalid_argument("byte_protocol: missing callbacks");
}
}
std::vector<std::uint8_t> operator()(std::size_t index,
std::vector<std::uint8_t> state)
{
if (index >= count_)
throw std::out_of_range("byte_protocol index");
(void)index;
for (std::size_t r = 0; r < rounds_.size(); ++r)
{
const auto & spec = rounds_[r];
std::vector<std::uint8_t> blind(spec.blind_bytes);
if (spec.blind_bytes != 0)
dealer_->read(r, blind.data(), spec.blind_bytes);
auto outbound = spec.produce(state, blind.data(), blind.size());
if (outbound.size() != spec.msg_bytes)
throw std::logic_error("byte_protocol: produce size");
peer_->write(r, outbound.data(), outbound.size());
peer_->flush(r);
std::vector<std::uint8_t> inbound(spec.msg_bytes);
peer_->read(r, inbound.data(), inbound.size());
spec.finish(state, inbound.data(), inbound.size(), blind.data(),
blind.size());
}
return state;
}
std::size_t count() const noexcept { return count_; }
std::size_t rounds() const noexcept { return rounds_.size(); }
private:
std::size_t count_ = 0;
net::stream_array * peer_ = nullptr;
net::stream_array * dealer_ = nullptr;
std::vector<byte_round> rounds_;
};
class byte_protocol_factory
{
public:
explicit byte_protocol_factory(std::vector<byte_round> rounds)
: rounds_(std::move(rounds))
{
}
byte_protocol create(std::size_t count, net::stream_array & peer,
net::stream_array & dealer) const
{
return byte_protocol(count, peer, dealer, rounds_);
}
private:
std::vector<byte_round> rounds_;
};
HEDLEY_WARN_UNUSED_RESULT
inline byte_protocol_factory make_byte_protocol_factory(
std::vector<byte_round> rounds)
{
return byte_protocol_factory(std::move(rounds));
}
// ---------------------------------------------------------------------------
// Async-shaped byte protocol (blocking today; completion hook for later I/O)
// ---------------------------------------------------------------------------
/// @brief Same round shape as `byte_round`; used by `async_byte_protocol`.
struct async_byte_round
{
std::size_t blind_bytes = 0;
std::size_t msg_bytes = 0;
std::function<std::vector<std::uint8_t>(std::vector<std::uint8_t> & state,
const std::uint8_t * blind, std::size_t nb)>
produce;
std::function<void(std::vector<std::uint8_t> & state, const std::uint8_t * peer,
std::size_t nm, const std::uint8_t * blind, std::size_t nb)>
finish;
};
/// @brief Blocking runner with an optional per-round completion callback.
/// @details This keeps the legacy synchronous shape: rounds run on the calling
/// thread and the callback is a progress hook. For TRUE overlapped,
/// event-driven I/O with no busy-waiting and no blocking protocol
/// calls, use `dpf::async::overlapped_byte_protocol` in
/// `dpf/async_protocol.hpp`, which drives the same `async_byte_round`
/// list over an `async_stream_array` and completes via an
/// `asio::io_context`.
class async_byte_protocol
{
public:
using on_round_complete_fn = std::function<void(std::size_t round)>;
async_byte_protocol(std::size_t count, net::stream_array & peer,
net::stream_array & dealer, std::vector<async_byte_round> rounds,
on_round_complete_fn on_round_complete = {})
: count_(count),
peer_(&peer),
dealer_(&dealer),
rounds_(std::move(rounds)),
on_round_complete_(std::move(on_round_complete))
{
if (peer_->size() < rounds_.size() || dealer_->size() < rounds_.size())
throw std::invalid_argument("async_byte_protocol: not enough streams");
for (const auto & r : rounds_)
{
if (!r.produce || !r.finish)
throw std::invalid_argument("async_byte_protocol: missing callbacks");
}
}
std::vector<std::uint8_t> operator()(std::size_t index,
std::vector<std::uint8_t> state)
{
if (index >= count_)
throw std::out_of_range("async_byte_protocol index");
(void)index;
for (std::size_t r = 0; r < rounds_.size(); ++r)
{
const auto & spec = rounds_[r];
std::vector<std::uint8_t> blind(spec.blind_bytes);
if (spec.blind_bytes != 0)
dealer_->read(r, blind.data(), spec.blind_bytes);
auto outbound = spec.produce(state, blind.data(), blind.size());
if (outbound.size() != spec.msg_bytes)
throw std::logic_error("async_byte_protocol: produce size");
peer_->write(r, outbound.data(), outbound.size());
peer_->flush(r);
std::vector<std::uint8_t> inbound(spec.msg_bytes);
peer_->read(r, inbound.data(), inbound.size());
spec.finish(state, inbound.data(), inbound.size(), blind.data(),
blind.size());
if (on_round_complete_)
on_round_complete_(r);
}
return state;
}
std::size_t count() const noexcept { return count_; }
std::size_t rounds() const noexcept { return rounds_.size(); }
private:
std::size_t count_ = 0;
net::stream_array * peer_ = nullptr;
net::stream_array * dealer_ = nullptr;
std::vector<async_byte_round> rounds_;
on_round_complete_fn on_round_complete_;
};
class async_byte_protocol_factory
{
public:
explicit async_byte_protocol_factory(std::vector<async_byte_round> rounds)
: rounds_(std::move(rounds))
{
}
async_byte_protocol create(std::size_t count, net::stream_array & peer,
net::stream_array & dealer,
async_byte_protocol::on_round_complete_fn on_round_complete = {}) const
{
return async_byte_protocol(count, peer, dealer, rounds_,
std::move(on_round_complete));
}
private:
std::vector<async_byte_round> rounds_;
};
HEDLEY_WARN_UNUSED_RESULT
inline async_byte_protocol_factory make_async_byte_protocol_factory(
std::vector<async_byte_round> rounds)
{
return async_byte_protocol_factory(std::move(rounds));
}
// ---------------------------------------------------------------------------
// Beaver ring multiply (open d, e; finish locally)
// ---------------------------------------------------------------------------
struct beaver_mul_input
{
std::uint64_t x = 0;
std::uint64_t y = 0;
};
struct beaver_mul_keep
{
std::uint64_t x = 0;
std::uint64_t y = 0;
std::uint64_t a = 0;
std::uint64_t b = 0;
std::uint64_t c = 0;
std::uint64_t d_open = 0;
std::uint64_t d_mine = 0;
std::uint64_t e_mine = 0;
};
/// @brief Round 0: open `d = x - a` on peer stream 0.
struct beaver_mul_round0
{
std::uint16_t limb = 8;
std::pair<std::vector<std::uint8_t>, beaver_mul_keep> operator()(
beaver_mul_input in, ring_triple_view blind) const
{
if (blind.limb != limb)
throw std::invalid_argument("beaver_mul_round0 limb");
std::uint64_t a = 0, b = 0, c = 0;
std::memcpy(&a, blind.a, limb);
std::memcpy(&b, blind.b, limb);
std::memcpy(&c, blind.c, limb);
const std::uint64_t mask = limb >= 8
? ~std::uint64_t{0}
: ((std::uint64_t{1} << (8u * limb)) - 1u);
const std::uint64_t d = (in.x - a) & mask;
std::vector<std::uint8_t> msg(limb);
std::memcpy(msg.data(), &d, limb);
beaver_mul_keep keep{};
keep.x = in.x & mask;
keep.y = in.y & mask;
keep.a = a;
keep.b = b;
keep.c = c;
keep.d_mine = d;
return {std::move(msg), keep};
}
};
/// @brief Round 1: open `e = y - b`, finish product share.
struct beaver_mul_round1
{
std::uint16_t limb = 8;
unsigned party = 0;
std::uint64_t operator()(beaver_mul_keep keep, std::vector<std::uint8_t> peer_e,
ring_triple_view /*unused*/) const
{
if (peer_e.size() != limb)
throw std::invalid_argument("beaver_mul_round1 msg");
const std::uint64_t mask = limb >= 8
? ~std::uint64_t{0}
: ((std::uint64_t{1} << (8u * limb)) - 1u);
std::uint64_t e_peer = 0;
std::memcpy(&e_peer, peer_e.data(), limb);
const std::uint64_t e_open = (keep.e_mine + e_peer) & mask;
std::uint64_t z = (keep.c + keep.d_open * keep.b + e_open * keep.a) & mask;
if (party == 0)
z = (z + keep.d_open * e_open) & mask;
return z;
}
};
/// @brief Two-round beaver multiply: blind triple on dealer stream 0 only.
template <typename Round0, typename Round1>
class beaver_mul_protocol
{
public:
beaver_mul_protocol(std::size_t count, net::stream_array & peer,
net::stream_array & dealer, Round0 r0, Round1 r1, std::uint16_t limb)
: count_(count),
peer_(&peer),
dealer_(&dealer),
r0_(std::move(r0)),
r1_(std::move(r1)),
limb_(limb)
{
if (peer_->size() < 2 || dealer_->size() < 1)
throw std::invalid_argument("beaver_mul_protocol streams");
}
std::uint64_t operator()(std::size_t index, beaver_mul_input input)
{
if (index >= count_)
throw std::out_of_range("beaver_mul_protocol index");
ring_triple_view blind{};
blind.limb = limb_;
dealer_->read(0, &blind, sizeof(blind));
auto step0 = r0_(input, blind);
peer_->write(0, step0.first.data(), step0.first.size());
peer_->flush(0);
std::vector<std::uint8_t> peer_d(limb_);
peer_->read(0, peer_d.data(), limb_);
std::uint64_t d_peer = 0;
std::memcpy(&d_peer, peer_d.data(), limb_);
const std::uint64_t mask = limb_ >= 8
? ~std::uint64_t{0}
: ((std::uint64_t{1} << (8u * limb_)) - 1u);
auto keep = step0.second;
keep.d_open = (keep.d_mine + d_peer) & mask;
const std::uint64_t e = (keep.y - keep.b) & mask;
keep.e_mine = e;
std::vector<std::uint8_t> msg_e(limb_);
std::memcpy(msg_e.data(), &e, limb_);
peer_->write(1, msg_e.data(), limb_);
peer_->flush(1);
std::vector<std::uint8_t> peer_e(limb_);
peer_->read(1, peer_e.data(), limb_);
ring_triple_view unused{};
return r1_(keep, std::move(peer_e), unused);
}
private:
std::size_t count_ = 0;
net::stream_array * peer_ = nullptr;
net::stream_array * dealer_ = nullptr;
Round0 r0_;
Round1 r1_;
std::uint16_t limb_ = 8;
};
class beaver_mul_factory
{
public:
beaver_mul_factory(beaver_mul_round0 r0, beaver_mul_round1 r1)
: r0_(std::move(r0)), r1_(std::move(r1))
{
}
beaver_mul_protocol<beaver_mul_round0, beaver_mul_round1> create(
std::size_t count, net::stream_array & peer, net::stream_array & dealer) const
{
return beaver_mul_protocol<beaver_mul_round0, beaver_mul_round1>(
count, peer, dealer, r0_, r1_, r0_.limb);
}
private:
beaver_mul_round0 r0_;
beaver_mul_round1 r1_;
};
HEDLEY_WARN_UNUSED_RESULT
inline beaver_mul_factory make_beaver_mul_factory(beaver_mul_round0 r0,
beaver_mul_round1 r1)
{
return beaver_mul_factory(std::move(r0), std::move(r1));
}
inline auto make_beaver_mul_dealer_functor()
{
return [limb = std::uint16_t{8}] {
return deal_ring_triple(limb);
};
}
} // namespace factory
} // namespace dpf
#endif