libdpf/include/dpf/net/stream_mesh.hpp

124 lines
3.7 KiB
C++
Raw Permalink Normal View History

/// @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 <algorithm>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <memory>
#include <stdexcept>
#include <utility>
#include <vector>
#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<std::shared_ptr<memory_stream_hub>> hubs;
std::vector<std::pair<memory_stream_array, memory_stream_array>> 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<std::uint8_t> 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<std::size_t>(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<memory_stream_hub>(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