124 lines
3.7 KiB
C++
124 lines
3.7 KiB
C++
|
|
/// @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
|