192 lines
5.7 KiB
C++
192 lines
5.7 KiB
C++
|
|
/// @file dpf/share_vec.hpp
|
||
|
|
/// @brief Lane-block secret-share vectors for arithmetic and boolean domains.
|
||
|
|
#ifndef LIBDPF_INCLUDE_DPF_SHARE_VEC_HPP__
|
||
|
|
#define LIBDPF_INCLUDE_DPF_SHARE_VEC_HPP__
|
||
|
|
|
||
|
|
#include <cstddef>
|
||
|
|
#include <cstdint>
|
||
|
|
#include <stdexcept>
|
||
|
|
#include <utility>
|
||
|
|
#include <vector>
|
||
|
|
|
||
|
|
#include "hedley/hedley.h"
|
||
|
|
|
||
|
|
#include "dpf/compose.hpp"
|
||
|
|
#include "dpf/random.hpp"
|
||
|
|
#include "dpf/trunc.hpp"
|
||
|
|
|
||
|
|
namespace dpf
|
||
|
|
{
|
||
|
|
|
||
|
|
template <typename T>
|
||
|
|
class share_vec
|
||
|
|
{
|
||
|
|
public:
|
||
|
|
using value_type = T;
|
||
|
|
|
||
|
|
share_vec() = default;
|
||
|
|
|
||
|
|
share_vec(std::size_t n, protocol::domain dom, std::size_t party)
|
||
|
|
: dom_(dom), party_(party), data_(n)
|
||
|
|
{
|
||
|
|
}
|
||
|
|
|
||
|
|
share_vec(std::vector<T> data, protocol::domain dom, std::size_t party)
|
||
|
|
: dom_(dom), party_(party), data_(std::move(data))
|
||
|
|
{
|
||
|
|
}
|
||
|
|
|
||
|
|
std::size_t size() const noexcept { return data_.size(); }
|
||
|
|
protocol::domain domain() const noexcept { return dom_; }
|
||
|
|
std::size_t party() const noexcept { return party_; }
|
||
|
|
|
||
|
|
T & operator[](std::size_t i) { return data_.at(i); }
|
||
|
|
const T & operator[](std::size_t i) const { return data_.at(i); }
|
||
|
|
|
||
|
|
T * data() noexcept { return data_.data(); }
|
||
|
|
const T * data() const noexcept { return data_.data(); }
|
||
|
|
|
||
|
|
share_vec & operator+=(const share_vec & o)
|
||
|
|
{
|
||
|
|
check_compatible(o);
|
||
|
|
for (std::size_t i = 0; i < data_.size(); ++i)
|
||
|
|
data_[i] = static_cast<T>(data_[i] + o.data_[i]);
|
||
|
|
return *this;
|
||
|
|
}
|
||
|
|
|
||
|
|
friend share_vec operator+(share_vec a, const share_vec & b)
|
||
|
|
{
|
||
|
|
a += b;
|
||
|
|
return a;
|
||
|
|
}
|
||
|
|
|
||
|
|
share_vec & operator-=(const share_vec & o)
|
||
|
|
{
|
||
|
|
check_compatible(o);
|
||
|
|
for (std::size_t i = 0; i < data_.size(); ++i)
|
||
|
|
data_[i] = static_cast<T>(data_[i] - o.data_[i]);
|
||
|
|
return *this;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Local cross term only (Beaver / RSS correction still required).
|
||
|
|
static share_vec mul_local(const share_vec & a, const share_vec & b)
|
||
|
|
{
|
||
|
|
a.check_compatible(b);
|
||
|
|
share_vec out(a.size(), a.dom_, a.party_);
|
||
|
|
for (std::size_t i = 0; i < a.size(); ++i)
|
||
|
|
out[i] = static_cast<T>(a[i] * b[i]);
|
||
|
|
return out;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Beaver product of two additive share vectors (both parties).
|
||
|
|
static std::pair<share_vec, share_vec> mul(const share_vec & x0,
|
||
|
|
const share_vec & x1, const share_vec & y0, const share_vec & y1)
|
||
|
|
{
|
||
|
|
if (x0.size() != x1.size() || x0.size() != y0.size()
|
||
|
|
|| y0.size() != y1.size())
|
||
|
|
throw std::invalid_argument("share_vec::mul size");
|
||
|
|
share_vec z0(x0.size(), x0.domain(), 0);
|
||
|
|
share_vec z1(x0.size(), x0.domain(), 1);
|
||
|
|
for (std::size_t i = 0; i < x0.size(); ++i)
|
||
|
|
{
|
||
|
|
auto mt = trunc::mul_trunc_pair(x0[i], x1[i], y0[i], y1[i], 0);
|
||
|
|
z0[i] = mt.z0;
|
||
|
|
z1[i] = mt.z1;
|
||
|
|
}
|
||
|
|
return {std::move(z0), std::move(z1)};
|
||
|
|
}
|
||
|
|
|
||
|
|
static std::vector<T> declassify(const share_vec & a, const share_vec & b)
|
||
|
|
{
|
||
|
|
if (a.size() != b.size())
|
||
|
|
throw std::invalid_argument("share_vec declassify size");
|
||
|
|
std::vector<T> out(a.size());
|
||
|
|
for (std::size_t i = 0; i < a.size(); ++i)
|
||
|
|
out[i] = static_cast<T>(a[i] + b[i]);
|
||
|
|
return out;
|
||
|
|
}
|
||
|
|
|
||
|
|
static std::pair<share_vec, share_vec> share(const std::vector<T> & clear,
|
||
|
|
protocol::domain dom = protocol::domain::a)
|
||
|
|
{
|
||
|
|
share_vec p0(clear.size(), dom, 0);
|
||
|
|
share_vec p1(clear.size(), dom, 1);
|
||
|
|
for (std::size_t i = 0; i < clear.size(); ++i)
|
||
|
|
{
|
||
|
|
p0[i] = dpf::uniform_sample<T>();
|
||
|
|
p1[i] = static_cast<T>(clear[i] - p0[i]);
|
||
|
|
}
|
||
|
|
return {std::move(p0), std::move(p1)};
|
||
|
|
}
|
||
|
|
|
||
|
|
private:
|
||
|
|
void check_compatible(const share_vec & o) const
|
||
|
|
{
|
||
|
|
if (data_.size() != o.data_.size() || dom_ != o.dom_
|
||
|
|
|| party_ != o.party_)
|
||
|
|
throw std::invalid_argument("share_vec incompatible");
|
||
|
|
}
|
||
|
|
|
||
|
|
protocol::domain dom_ = protocol::domain::a;
|
||
|
|
std::size_t party_ = 0;
|
||
|
|
std::vector<T> data_;
|
||
|
|
};
|
||
|
|
|
||
|
|
class bool_share_vec
|
||
|
|
{
|
||
|
|
public:
|
||
|
|
bool_share_vec() = default;
|
||
|
|
|
||
|
|
bool_share_vec(std::size_t nbits, protocol::domain dom, std::size_t party)
|
||
|
|
: nbits_(nbits),
|
||
|
|
dom_(dom),
|
||
|
|
party_(party),
|
||
|
|
data_((nbits + 7u) / 8u, 0)
|
||
|
|
{
|
||
|
|
if (dom != protocol::domain::bin && dom != protocol::domain::bin_rss)
|
||
|
|
throw std::invalid_argument("bool_share_vec domain");
|
||
|
|
}
|
||
|
|
|
||
|
|
std::size_t size() const noexcept { return nbits_; }
|
||
|
|
std::vector<std::uint8_t> & bytes() noexcept { return data_; }
|
||
|
|
const std::vector<std::uint8_t> & bytes() const noexcept { return data_; }
|
||
|
|
|
||
|
|
std::uint8_t bit(std::size_t i) const
|
||
|
|
{
|
||
|
|
if (i >= nbits_)
|
||
|
|
throw std::out_of_range("bool_share_vec");
|
||
|
|
return static_cast<std::uint8_t>((data_[i / 8] >> (i % 8)) & 1u);
|
||
|
|
}
|
||
|
|
|
||
|
|
void set_bit(std::size_t i, std::uint8_t b)
|
||
|
|
{
|
||
|
|
if (i >= nbits_)
|
||
|
|
throw std::out_of_range("bool_share_vec");
|
||
|
|
const std::size_t byte = i / 8;
|
||
|
|
const unsigned off = static_cast<unsigned>(i % 8);
|
||
|
|
if (b & 1u)
|
||
|
|
data_[byte] = static_cast<std::uint8_t>(data_[byte] | (1u << off));
|
||
|
|
else
|
||
|
|
data_[byte] = static_cast<std::uint8_t>(data_[byte] & ~(1u << off));
|
||
|
|
}
|
||
|
|
|
||
|
|
bool_share_vec & operator^=(const bool_share_vec & o)
|
||
|
|
{
|
||
|
|
if (nbits_ != o.nbits_ || party_ != o.party_)
|
||
|
|
throw std::invalid_argument("bool_share_vec xor");
|
||
|
|
for (std::size_t i = 0; i < data_.size(); ++i)
|
||
|
|
data_[i] = static_cast<std::uint8_t>(data_[i] ^ o.data_[i]);
|
||
|
|
return *this;
|
||
|
|
}
|
||
|
|
|
||
|
|
private:
|
||
|
|
std::size_t nbits_ = 0;
|
||
|
|
protocol::domain dom_ = protocol::domain::bin;
|
||
|
|
std::size_t party_ = 0;
|
||
|
|
std::vector<std::uint8_t> data_;
|
||
|
|
};
|
||
|
|
|
||
|
|
} // namespace dpf
|
||
|
|
|
||
|
|
#endif // LIBDPF_INCLUDE_DPF_SHARE_VEC_HPP__
|