Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
404 lines
12 KiB
C++
404 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__
|