/// @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 #include #include #include #include #include "hedley/hedley.h" #include "dpf/compose.hpp" #include "dpf/random.hpp" #include "dpf/trunc.hpp" namespace dpf { template 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 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(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(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(a[i] * b[i]); return out; } /// @brief Beaver product of two additive share vectors (both parties). static std::pair 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 declassify(const share_vec & a, const share_vec & b) { if (a.size() != b.size()) throw std::invalid_argument("share_vec declassify size"); std::vector out(a.size()); for (std::size_t i = 0; i < a.size(); ++i) out[i] = static_cast(a[i] + b[i]); return out; } static std::pair share(const std::vector & 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(); p1[i] = static_cast(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 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 & bytes() noexcept { return data_; } const std::vector & 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((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(i % 8); if (b & 1u) data_[byte] = static_cast(data_[byte] | (1u << off)); else data_[byte] = static_cast(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(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 data_; }; } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_SHARE_VEC_HPP__