405 lines
12 KiB
C++
405 lines
12 KiB
C++
|
|
/// @file dpf/ot_pack.hpp
|
|||
|
|
/// @brief Consuming cursor over correlated B2A / bit / bit×ring pads.
|
|||
|
|
/// @details Dealer mode builds one shared tape and splits it into party views.
|
|||
|
|
/// Two parties each calling `dealer()` independently is unsupported;
|
|||
|
|
/// use `sample_dealer_pair`.
|
|||
|
|
#ifndef LIBDPF_INCLUDE_DPF_OT_PACK_HPP__
|
|||
|
|
#define LIBDPF_INCLUDE_DPF_OT_PACK_HPP__
|
|||
|
|
|
|||
|
|
#include <cstddef>
|
|||
|
|
#include <cstdint>
|
|||
|
|
#include <cstring>
|
|||
|
|
#include <stdexcept>
|
|||
|
|
#include <utility>
|
|||
|
|
#include <vector>
|
|||
|
|
|
|||
|
|
#include "hedley/hedley.h"
|
|||
|
|
#include "simde/simde/x86/avx2.h"
|
|||
|
|
|
|||
|
|
#include "dpf/beaver.hpp"
|
|||
|
|
#include "dpf/random.hpp"
|
|||
|
|
|
|||
|
|
namespace dpf
|
|||
|
|
{
|
|||
|
|
namespace ot
|
|||
|
|
{
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
struct dabit
|
|||
|
|
{
|
|||
|
|
std::uint8_t bit = 0;
|
|||
|
|
Ring arith{};
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct bit_triple
|
|||
|
|
{
|
|||
|
|
std::uint8_t a = 0;
|
|||
|
|
std::uint8_t b = 0;
|
|||
|
|
std::uint8_t c = 0;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
/// @brief Bit × ring triple: `(a0⊕a1)·(b0+b1) = c0+c1`.
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
struct bit_ring_triple
|
|||
|
|
{
|
|||
|
|
std::uint8_t a = 0; ///< XOR share of the bit
|
|||
|
|
Ring b{}; ///< Additive share of the scalar
|
|||
|
|
Ring c{}; ///< Additive share of the product
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct b2a_slot
|
|||
|
|
{
|
|||
|
|
std::uint8_t r = 0;
|
|||
|
|
std::uint64_t add = 0;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
struct dabit_pair
|
|||
|
|
{
|
|||
|
|
dabit<Ring> p0;
|
|||
|
|
dabit<Ring> p1;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct bit_triple_pair
|
|||
|
|
{
|
|||
|
|
bit_triple p0;
|
|||
|
|
bit_triple p1;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
struct bit_ring_triple_pair
|
|||
|
|
{
|
|||
|
|
bit_ring_triple<Ring> p0;
|
|||
|
|
bit_ring_triple<Ring> p1;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
dabit_pair<Ring> sample_dabit_pair()
|
|||
|
|
{
|
|||
|
|
struct pad
|
|||
|
|
{
|
|||
|
|
std::uint8_t bit()
|
|||
|
|
{
|
|||
|
|
return static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<std::uint8_t>() & 1u);
|
|||
|
|
}
|
|||
|
|
simde__m128i block()
|
|||
|
|
{
|
|||
|
|
return dpf::uniform_sample<simde__m128i>();
|
|||
|
|
}
|
|||
|
|
} p;
|
|||
|
|
auto ba = beavers::sample_bit_arith(p);
|
|||
|
|
dabit_pair<Ring> out;
|
|||
|
|
out.p0.bit = ba.xor0;
|
|||
|
|
out.p0.arith = static_cast<Ring>(ba.add0);
|
|||
|
|
out.p1.bit = ba.xor1;
|
|||
|
|
out.p1.arith = static_cast<Ring>(ba.add1);
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
inline bit_triple_pair sample_bit_triple_pair()
|
|||
|
|
{
|
|||
|
|
const std::uint8_t a = static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<std::uint8_t>() & 1u);
|
|||
|
|
const std::uint8_t b = static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<std::uint8_t>() & 1u);
|
|||
|
|
const std::uint8_t c = static_cast<std::uint8_t>(a & b);
|
|||
|
|
const std::uint8_t a0 = static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<std::uint8_t>() & 1u);
|
|||
|
|
const std::uint8_t b0 = static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<std::uint8_t>() & 1u);
|
|||
|
|
const std::uint8_t c0 = static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<std::uint8_t>() & 1u);
|
|||
|
|
bit_triple_pair out;
|
|||
|
|
out.p0 = {a0, b0, c0};
|
|||
|
|
out.p1 = {static_cast<std::uint8_t>(a ^ a0),
|
|||
|
|
static_cast<std::uint8_t>(b ^ b0),
|
|||
|
|
static_cast<std::uint8_t>(c ^ c0)};
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
bit_ring_triple_pair<Ring> sample_bit_ring_triple_pair()
|
|||
|
|
{
|
|||
|
|
// a is XOR-shared with party 1 holding 0 so a0+a1 = a0⊕a1 in the ring.
|
|||
|
|
const std::uint8_t a = static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<std::uint8_t>() & 1u);
|
|||
|
|
const Ring b = dpf::uniform_sample<Ring>();
|
|||
|
|
const Ring c = static_cast<Ring>(Ring{a} * b);
|
|||
|
|
const Ring b0 = dpf::uniform_sample<Ring>();
|
|||
|
|
const Ring c0 = dpf::uniform_sample<Ring>();
|
|||
|
|
bit_ring_triple_pair<Ring> out;
|
|||
|
|
out.p0 = {a, b0, c0};
|
|||
|
|
out.p1 = {0, static_cast<Ring>(b - b0), static_cast<Ring>(c - c0)};
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
class pack
|
|||
|
|
{
|
|||
|
|
public:
|
|||
|
|
pack() = default;
|
|||
|
|
|
|||
|
|
void load_b2a(std::vector<b2a_slot> slots)
|
|||
|
|
{
|
|||
|
|
mode_ = mode::pads;
|
|||
|
|
b2a_ = std::move(slots);
|
|||
|
|
b2a_pos_ = 0;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
void load_bits(std::vector<bit_triple> bits)
|
|||
|
|
{
|
|||
|
|
mode_ = mode::pads;
|
|||
|
|
bits_ = std::move(bits);
|
|||
|
|
bit_pos_ = 0;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
void load_bit_ring(std::vector<bit_ring_triple<Ring>> triples)
|
|||
|
|
{
|
|||
|
|
mode_ = mode::pads;
|
|||
|
|
bit_ring_.clear();
|
|||
|
|
bit_ring_.reserve(triples.size() * sizeof(bit_ring_triple<Ring>));
|
|||
|
|
for (const auto & t : triples)
|
|||
|
|
{
|
|||
|
|
const auto * p = reinterpret_cast<const std::uint8_t *>(&t);
|
|||
|
|
bit_ring_.insert(bit_ring_.end(), p, p + sizeof(t));
|
|||
|
|
}
|
|||
|
|
bit_ring_elem_ = sizeof(bit_ring_triple<Ring>);
|
|||
|
|
bit_ring_pos_ = 0;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief One party's view of a pre-split correlated tape.
|
|||
|
|
static pack from_party_view(int me, std::vector<b2a_slot> b2a,
|
|||
|
|
std::vector<bit_triple> bits,
|
|||
|
|
std::vector<std::uint8_t> bit_ring_bytes = {},
|
|||
|
|
std::size_t bit_ring_elem = 0)
|
|||
|
|
{
|
|||
|
|
pack p;
|
|||
|
|
p.mode_ = mode::pads;
|
|||
|
|
p.me_ = me;
|
|||
|
|
p.b2a_ = std::move(b2a);
|
|||
|
|
p.bits_ = std::move(bits);
|
|||
|
|
p.bit_ring_ = std::move(bit_ring_bytes);
|
|||
|
|
p.bit_ring_elem_ = bit_ring_elem;
|
|||
|
|
return p;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
std::size_t remaining_b2a() const noexcept
|
|||
|
|
{
|
|||
|
|
return b2a_.size() > b2a_pos_ ? b2a_.size() - b2a_pos_ : 0;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
std::size_t remaining_bits() const noexcept
|
|||
|
|
{
|
|||
|
|
return bits_.size() > bit_pos_ ? bits_.size() - bit_pos_ : 0;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
std::size_t remaining_bit_ring() const noexcept
|
|||
|
|
{
|
|||
|
|
if (bit_ring_elem_ == 0)
|
|||
|
|
return 0;
|
|||
|
|
const std::size_t n = bit_ring_.size() / bit_ring_elem_;
|
|||
|
|
return n > bit_ring_pos_ ? n - bit_ring_pos_ : 0;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
dabit<Ring> take_dabit()
|
|||
|
|
{
|
|||
|
|
if (b2a_pos_ >= b2a_.size())
|
|||
|
|
throw std::runtime_error("ot_pack: no B2A slots left");
|
|||
|
|
const auto & s = b2a_[b2a_pos_++];
|
|||
|
|
dabit<Ring> out;
|
|||
|
|
out.bit = s.r;
|
|||
|
|
out.arith = static_cast<Ring>(s.add);
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
bit_triple take_bit_triple()
|
|||
|
|
{
|
|||
|
|
if (bit_pos_ >= bits_.size())
|
|||
|
|
throw std::runtime_error("ot_pack: no bit triples left");
|
|||
|
|
return bits_[bit_pos_++];
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
bit_ring_triple<Ring> take_bit_ring()
|
|||
|
|
{
|
|||
|
|
if (bit_ring_elem_ != sizeof(bit_ring_triple<Ring>))
|
|||
|
|
throw std::runtime_error("ot_pack: bit_ring element size mismatch");
|
|||
|
|
if (remaining_bit_ring() == 0)
|
|||
|
|
throw std::runtime_error("ot_pack: no bit×ring triples left");
|
|||
|
|
bit_ring_triple<Ring> t{};
|
|||
|
|
std::memcpy(&t, bit_ring_.data() + bit_ring_pos_ * bit_ring_elem_,
|
|||
|
|
sizeof(t));
|
|||
|
|
++bit_ring_pos_;
|
|||
|
|
return t;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
int party() const noexcept { return me_; }
|
|||
|
|
|
|||
|
|
/// @brief Remember a party tape written by `fill_tape_ot_pair`.
|
|||
|
|
template <typename Ring>
|
|||
|
|
void stash_party_tape(const beavers::party_tape<Ring> & tape)
|
|||
|
|
{
|
|||
|
|
tape_elem_ = sizeof(Ring);
|
|||
|
|
tape_lambda_ready_ = tape.lambda_ready;
|
|||
|
|
tape_mono_ready_ = tape.monomial_ready;
|
|||
|
|
tape_bundle_ready_ = tape.bundles_ready;
|
|||
|
|
tape_dot_ready_ = tape.dot_ready;
|
|||
|
|
tape_lambda_ = bytes_of(tape.lambda);
|
|||
|
|
tape_mono_ = bytes_of(tape.monomial);
|
|||
|
|
tape_bundle_ = bytes_of(tape.bundles);
|
|||
|
|
tape_dot_ = bytes_of(tape.dot_cross);
|
|||
|
|
has_tape_ = true;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
bool has_party_tape() const noexcept { return has_tape_; }
|
|||
|
|
|
|||
|
|
/// @brief Serializable OT pad view for dealer / online streams.
|
|||
|
|
struct wire
|
|||
|
|
{
|
|||
|
|
int me = 0;
|
|||
|
|
std::vector<b2a_slot> b2a;
|
|||
|
|
std::vector<bit_triple> bits;
|
|||
|
|
std::vector<std::uint8_t> bit_ring;
|
|||
|
|
std::size_t bit_ring_elem = 0;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
wire export_wire() const
|
|||
|
|
{
|
|||
|
|
wire w;
|
|||
|
|
w.me = me_;
|
|||
|
|
w.b2a = b2a_;
|
|||
|
|
w.bits = bits_;
|
|||
|
|
w.bit_ring = bit_ring_;
|
|||
|
|
w.bit_ring_elem = bit_ring_elem_;
|
|||
|
|
return w;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
static pack from_wire(wire w)
|
|||
|
|
{
|
|||
|
|
pack p;
|
|||
|
|
p.mode_ = mode::pads;
|
|||
|
|
p.me_ = w.me;
|
|||
|
|
p.b2a_ = std::move(w.b2a);
|
|||
|
|
p.bits_ = std::move(w.bits);
|
|||
|
|
p.bit_ring_ = std::move(w.bit_ring);
|
|||
|
|
p.bit_ring_elem_ = w.bit_ring_elem;
|
|||
|
|
return p;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
template <typename Ring>
|
|||
|
|
beavers::party_tape<Ring> load_party_tape() const
|
|||
|
|
{
|
|||
|
|
if (!has_tape_ || tape_elem_ != sizeof(Ring))
|
|||
|
|
throw std::runtime_error(
|
|||
|
|
"ot_pack: no stashed party tape for this ring");
|
|||
|
|
beavers::party_tape<Ring> t;
|
|||
|
|
t.lambda_ready = tape_lambda_ready_;
|
|||
|
|
t.monomial_ready = tape_mono_ready_;
|
|||
|
|
t.bundles_ready = tape_bundle_ready_;
|
|||
|
|
t.dot_ready = tape_dot_ready_;
|
|||
|
|
t.lambda = vec_of<Ring>(tape_lambda_);
|
|||
|
|
t.monomial = vec_of<Ring>(tape_mono_);
|
|||
|
|
t.bundles = vec_of<Ring>(tape_bundle_);
|
|||
|
|
t.dot_cross = vec_of<Ring>(tape_dot_);
|
|||
|
|
return t;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
private:
|
|||
|
|
template <typename Ring>
|
|||
|
|
static std::vector<std::uint8_t> bytes_of(const std::vector<Ring> & v)
|
|||
|
|
{
|
|||
|
|
std::vector<std::uint8_t> out(v.size() * sizeof(Ring));
|
|||
|
|
if (!out.empty())
|
|||
|
|
std::memcpy(out.data(), v.data(), out.size());
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
template <typename Ring>
|
|||
|
|
static std::vector<Ring> vec_of(const std::vector<std::uint8_t> & b)
|
|||
|
|
{
|
|||
|
|
if (b.size() % sizeof(Ring) != 0)
|
|||
|
|
throw std::runtime_error("ot_pack: tape byte size");
|
|||
|
|
std::vector<Ring> out(b.size() / sizeof(Ring));
|
|||
|
|
if (!out.empty())
|
|||
|
|
std::memcpy(out.data(), b.data(), b.size());
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
enum class mode : unsigned char { empty, pads };
|
|||
|
|
mode mode_ = mode::empty;
|
|||
|
|
int me_ = 0;
|
|||
|
|
std::vector<b2a_slot> b2a_;
|
|||
|
|
std::vector<bit_triple> bits_;
|
|||
|
|
std::vector<std::uint8_t> bit_ring_;
|
|||
|
|
std::size_t bit_ring_elem_ = 0;
|
|||
|
|
std::size_t b2a_pos_ = 0;
|
|||
|
|
std::size_t bit_pos_ = 0;
|
|||
|
|
std::size_t bit_ring_pos_ = 0;
|
|||
|
|
bool has_tape_ = false;
|
|||
|
|
std::size_t tape_elem_ = 0;
|
|||
|
|
std::vector<std::uint8_t> tape_lambda_ready_;
|
|||
|
|
std::vector<std::uint8_t> tape_mono_ready_;
|
|||
|
|
std::vector<std::uint8_t> tape_bundle_ready_;
|
|||
|
|
std::vector<std::uint8_t> tape_dot_ready_;
|
|||
|
|
std::vector<std::uint8_t> tape_lambda_;
|
|||
|
|
std::vector<std::uint8_t> tape_mono_;
|
|||
|
|
std::vector<std::uint8_t> tape_bundle_;
|
|||
|
|
std::vector<std::uint8_t> tape_dot_;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
/// @brief Sample correlated pads and return both party views.
|
|||
|
|
template <typename Ring = std::uint64_t>
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
std::pair<pack, pack> sample_dealer_pair(std::size_t n_dabit,
|
|||
|
|
std::size_t n_bit_triple = 0, std::size_t n_bit_ring = 0)
|
|||
|
|
{
|
|||
|
|
std::vector<b2a_slot> b0, b1;
|
|||
|
|
b0.reserve(n_dabit);
|
|||
|
|
b1.reserve(n_dabit);
|
|||
|
|
for (std::size_t i = 0; i < n_dabit; ++i)
|
|||
|
|
{
|
|||
|
|
auto d = sample_dabit_pair<Ring>();
|
|||
|
|
b0.push_back({d.p0.bit, static_cast<std::uint64_t>(d.p0.arith)});
|
|||
|
|
b1.push_back({d.p1.bit, static_cast<std::uint64_t>(d.p1.arith)});
|
|||
|
|
}
|
|||
|
|
std::vector<bit_triple> t0, t1;
|
|||
|
|
t0.reserve(n_bit_triple);
|
|||
|
|
t1.reserve(n_bit_triple);
|
|||
|
|
for (std::size_t i = 0; i < n_bit_triple; ++i)
|
|||
|
|
{
|
|||
|
|
auto tp = sample_bit_triple_pair();
|
|||
|
|
t0.push_back(tp.p0);
|
|||
|
|
t1.push_back(tp.p1);
|
|||
|
|
}
|
|||
|
|
std::vector<bit_ring_triple<Ring>> r0, r1;
|
|||
|
|
r0.reserve(n_bit_ring);
|
|||
|
|
r1.reserve(n_bit_ring);
|
|||
|
|
for (std::size_t i = 0; i < n_bit_ring; ++i)
|
|||
|
|
{
|
|||
|
|
auto rp = sample_bit_ring_triple_pair<Ring>();
|
|||
|
|
r0.push_back(rp.p0);
|
|||
|
|
r1.push_back(rp.p1);
|
|||
|
|
}
|
|||
|
|
auto pack0 = pack::from_party_view(0, std::move(b0), std::move(t0));
|
|||
|
|
auto pack1 = pack::from_party_view(1, std::move(b1), std::move(t1));
|
|||
|
|
if (n_bit_ring != 0)
|
|||
|
|
{
|
|||
|
|
pack0.load_bit_ring(std::move(r0));
|
|||
|
|
pack1.load_bit_ring(std::move(r1));
|
|||
|
|
}
|
|||
|
|
return {std::move(pack0), std::move(pack1)};
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
} // namespace ot
|
|||
|
|
} // namespace dpf
|
|||
|
|
|
|||
|
|
#endif // LIBDPF_INCLUDE_DPF_OT_PACK_HPP__
|