/// @file dpf/net/stream_mesh.hpp /// @brief Fully connected clique of `stream_array` duplex links. #ifndef LIBDPF_INCLUDE_DPF_NET_STREAM_MESH_HPP__ #define LIBDPF_INCLUDE_DPF_NET_STREAM_MESH_HPP__ #include #include #include #include #include #include #include #include #include "hedley/hedley.h" #include "dpf/net/stream_array.hpp" namespace dpf { namespace net { /// @brief Pairwise memory stream arrays for `n` parties (one hub per edge). struct memory_stream_clique { std::size_t parties = 0; std::size_t streams_per_edge = 0; std::vector> hubs; std::vector> ends; static std::size_t pair_edge(std::size_t a, std::size_t b, std::size_t n) { if (a == b || a >= n || b >= n) throw std::invalid_argument("memory_stream_clique pair"); if (a > b) std::swap(a, b); std::size_t e = 0; for (std::size_t i = 0; i < a; ++i) e += n - 1 - i; e += b - a - 1; return e; } /// @brief Duplex link from `from` toward `to` (same hub as `to`/`from`). memory_stream_array & end(std::size_t from, std::size_t to) { const auto e = pair_edge(from, to, parties); if (from < to) return ends[e].first; return ends[e].second; } const memory_stream_array & end(std::size_t from, std::size_t to) const { const auto e = pair_edge(from, to, parties); if (from < to) return ends[e].first; return ends[e].second; } /// @brief Sum additive shares: send `mine` to every peer, add incoming. /// @details Works for any `parties >= 2` (not only 3-party rings). void declassify(unsigned me, const std::uint8_t * mine, std::size_t n, std::uint8_t * sum_out) { if (parties < 2 || me >= parties) throw std::invalid_argument("memory_stream_clique declassify party"); for (unsigned peer = 0; peer < parties; ++peer) { if (peer == me) continue; end(me, peer).write(0, mine, n); end(me, peer).flush(0); } std::memcpy(sum_out, mine, n); for (unsigned peer = 0; peer < parties; ++peer) { if (peer == me) continue; std::vector got(n); end(me, peer).read(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); } } } }; HEDLEY_WARN_UNUSED_RESULT inline memory_stream_clique make_memory_stream_clique(std::size_t n, std::size_t streams_per_edge) { if (n < 2) throw std::invalid_argument("make_memory_stream_clique needs >= 2"); if (streams_per_edge == 0) throw std::invalid_argument("make_memory_stream_clique streams"); memory_stream_clique c; c.parties = n; c.streams_per_edge = streams_per_edge; const std::size_t n_edges = n * (n - 1) / 2; c.hubs.reserve(n_edges); c.ends.reserve(n_edges); for (std::size_t e = 0; e < n_edges; ++e) { auto hub = std::make_shared(streams_per_edge); c.hubs.push_back(hub); c.ends.emplace_back(memory_stream_array(hub, true), memory_stream_array(hub, false)); } return c; } } // namespace net } // namespace dpf #endif