libdpf/include/dpf/ot_pack.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
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>
2026-09-28 05:59:19 -06:00

404 lines
12 KiB
C++
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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