libdpf/include/dpf/ot_pack.hpp

405 lines
12 KiB
C++
Raw Normal View History

/// @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__