libdpf/include/dpf/prep_source.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

349 lines
11 KiB
C++

/// @file dpf/prep_source.hpp
/// @brief Preprocess views as a byte stream: dealer, file, or 2PC channel.
/// @details Online code reads one party's view. It does not sample the other
/// party's share. `deal_views` is the dealer. `setup_2pc_sampled` is
/// the flagged semi-honest 2PC setup (see `revealing.hpp`).
#ifndef LIBDPF_INCLUDE_DPF_PREP_SOURCE_HPP__
#define LIBDPF_INCLUDE_DPF_PREP_SOURCE_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <fstream>
#include <memory>
#include <stdexcept>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/net/channel.hpp"
#include "dpf/net/stream_array.hpp"
#include "dpf/ot_pack.hpp"
#include "dpf/protocol_factory.hpp"
#include "dpf/random.hpp"
namespace dpf
{
namespace prep
{
struct demand
{
std::uint16_t limb = 8;
std::uint32_t ring_triples = 0;
std::uint32_t dabits = 0;
std::uint32_t bit_triples = 0;
};
inline constexpr std::uint32_t k_magic = 0x31525050u; // 'PPR1' LE
/// @brief Cursor over one party's prep view.
class cursor
{
public:
cursor() = default;
explicit cursor(std::vector<std::uint8_t> bytes)
: bytes_(std::move(bytes))
{
if (bytes_.size() < 18)
throw std::runtime_error("prep cursor: short header");
std::uint32_t magic = 0;
read_pod(magic);
if (magic != k_magic)
throw std::runtime_error("prep cursor: bad magic");
std::uint16_t limb = 0;
read_pod(limb);
limb_ = limb;
read_pod(n_ring_);
read_pod(n_dabit_);
read_pod(n_bit_);
if (limb_ == 0 || limb_ > 8)
throw std::runtime_error("prep cursor: limb must be 1..8");
}
std::uint16_t limb() const noexcept { return limb_; }
void take_ring(std::uint8_t * a, std::uint8_t * b, std::uint8_t * c)
{
if (ring_used_ >= n_ring_)
throw std::runtime_error("prep cursor: ring triple exhausted");
std::memcpy(a, bytes_.data() + pos_, limb_);
pos_ += limb_;
std::memcpy(b, bytes_.data() + pos_, limb_);
pos_ += limb_;
std::memcpy(c, bytes_.data() + pos_, limb_);
pos_ += limb_;
++ring_used_;
}
ot::dabit<std::uint64_t> take_dabit()
{
if (dabit_used_ >= n_dabit_)
throw std::runtime_error("prep cursor: dabit exhausted");
ot::dabit<std::uint64_t> d;
d.bit = bytes_[pos_++];
std::uint64_t arith = 0;
std::memcpy(&arith, bytes_.data() + pos_, limb_);
pos_ += limb_;
d.arith = arith;
++dabit_used_;
return d;
}
ot::bit_triple take_bit()
{
if (bit_used_ >= n_bit_)
throw std::runtime_error("prep cursor: bit triple exhausted");
ot::bit_triple t;
t.a = bytes_[pos_++];
t.b = bytes_[pos_++];
t.c = bytes_[pos_++];
++bit_used_;
return t;
}
private:
template <typename T>
void read_pod(T & out)
{
if (pos_ + sizeof(T) > bytes_.size())
throw std::runtime_error("prep cursor: truncated header");
std::memcpy(&out, bytes_.data() + pos_, sizeof(T));
pos_ += sizeof(T);
}
std::vector<std::uint8_t> bytes_;
std::size_t pos_ = 0;
std::uint16_t limb_ = 8;
std::uint32_t n_ring_ = 0;
std::uint32_t n_dabit_ = 0;
std::uint32_t n_bit_ = 0;
std::uint32_t ring_used_ = 0;
std::uint32_t dabit_used_ = 0;
std::uint32_t bit_used_ = 0;
};
namespace detail
{
inline void append(std::vector<std::uint8_t> & o, const void * p, std::size_t n)
{
const auto * b = static_cast<const std::uint8_t *>(p);
o.insert(o.end(), b, b + n);
}
inline std::uint64_t limb_mask(std::uint16_t limb)
{
if (limb >= 8)
return ~std::uint64_t{0};
return (std::uint64_t{1} << (8u * limb)) - 1u;
}
inline void store_limb(std::vector<std::uint8_t> & o, std::uint64_t v,
std::uint16_t limb)
{
v &= limb_mask(limb);
append(o, &v, limb);
}
} // namespace detail
/// @brief Dealer writes both parties' views. Neither view contains the other.
HEDLEY_WARN_UNUSED_RESULT
inline std::pair<std::vector<std::uint8_t>, std::vector<std::uint8_t>>
deal_views(const demand & d)
{
if (d.limb == 0 || d.limb > 8)
throw std::invalid_argument("prep limb must be 1..8");
std::vector<std::uint8_t> p0, p1;
auto header = [&](std::vector<std::uint8_t> & o) {
detail::append(o, &k_magic, 4);
detail::append(o, &d.limb, 2);
detail::append(o, &d.ring_triples, 4);
detail::append(o, &d.dabits, 4);
detail::append(o, &d.bit_triples, 4);
};
header(p0);
header(p1);
for (std::uint32_t i = 0; i < d.ring_triples; ++i)
{
const auto a = dpf::uniform_sample<std::uint64_t>() & detail::limb_mask(d.limb);
const auto b = dpf::uniform_sample<std::uint64_t>() & detail::limb_mask(d.limb);
const auto c = (a * b) & detail::limb_mask(d.limb);
const auto a0 = dpf::uniform_sample<std::uint64_t>() & detail::limb_mask(d.limb);
const auto b0 = dpf::uniform_sample<std::uint64_t>() & detail::limb_mask(d.limb);
const auto c0 = dpf::uniform_sample<std::uint64_t>() & detail::limb_mask(d.limb);
detail::store_limb(p0, a0, d.limb);
detail::store_limb(p0, b0, d.limb);
detail::store_limb(p0, c0, d.limb);
detail::store_limb(p1, a - a0, d.limb);
detail::store_limb(p1, b - b0, d.limb);
detail::store_limb(p1, c - c0, d.limb);
}
for (std::uint32_t i = 0; i < d.dabits; ++i)
{
auto bit = ot::sample_dabit_pair<std::uint64_t>();
p0.push_back(bit.p0.bit);
detail::store_limb(p0, bit.p0.arith, d.limb);
p1.push_back(bit.p1.bit);
detail::store_limb(p1, bit.p1.arith, d.limb);
}
for (std::uint32_t i = 0; i < d.bit_triples; ++i)
{
auto t = ot::sample_bit_triple_pair();
p0.push_back(t.p0.a);
p0.push_back(t.p0.b);
p0.push_back(t.p0.c);
p1.push_back(t.p1.a);
p1.push_back(t.p1.b);
p1.push_back(t.p1.c);
}
return {std::move(p0), std::move(p1)};
}
inline void write_file(const std::string & path, const std::vector<std::uint8_t> & bytes)
{
std::ofstream out(path, std::ios::binary);
if (!out)
throw std::runtime_error("prep write_file: " + path);
out.write(reinterpret_cast<const char *>(bytes.data()),
static_cast<std::streamsize>(bytes.size()));
if (!out)
throw std::runtime_error("prep write_file failed");
}
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> read_file(const std::string & path)
{
std::ifstream in(path, std::ios::binary);
if (!in)
throw std::runtime_error("prep read_file: " + path);
return std::vector<std::uint8_t>(std::istreambuf_iterator<char>(in),
std::istreambuf_iterator<char>());
}
/// @brief Party 0 samples both views and sends party 1's. See `revealing.hpp`.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> setup_2pc_sampled(net::channel & ch,
unsigned party, const demand & d)
{
if (party > 1)
throw std::invalid_argument("setup_2pc_sampled party");
if (party == 0)
{
auto views = deal_views(d);
ch.send_vec(views.second, net::msg::beaver_tape);
return std::move(views.first);
}
return ch.recv_vec<std::uint8_t>(net::msg::beaver_tape);
}
/// @brief Write dealer views onto two stream arrays (stream 0 = full blob).
/// @details Convenience over `deal_views`. Online parties read stream 0 into a
/// `cursor`. Prefer `factory::make_dealer` with per-kind functors when
/// the protocol is assembled from round factories.
inline void deal_to_streams(const demand & d, net::stream_array & party0,
net::stream_array & party1)
{
if (party0.size() < 1 || party1.size() < 1)
throw std::invalid_argument("deal_to_streams: need stream 0");
auto views = deal_views(d);
party0.write(0, views.first.data(), views.first.size());
party0.flush(0);
party1.write(0, views.second.data(), views.second.size());
party1.flush(0);
}
/// @brief Fill a memory pair with dealer views; return both read ends.
HEDLEY_WARN_UNUSED_RESULT
inline std::pair<net::memory_stream_array, net::memory_stream_array>
deal_memory_pair(const demand & d)
{
auto hubs0 = std::make_shared<net::memory_stream_hub>(1);
auto hubs1 = std::make_shared<net::memory_stream_hub>(1);
net::memory_stream_array write0(hubs0, true);
net::memory_stream_array write1(hubs1, true);
deal_to_streams(d, write0, write1);
return {net::memory_stream_array(hubs0, false),
net::memory_stream_array(hubs1, false)};
}
/// @brief Read a full prep blob from stream `index` into a cursor.
HEDLEY_WARN_UNUSED_RESULT
inline cursor cursor_from_stream(net::stream_array & streams, std::size_t index,
std::size_t nbytes)
{
std::vector<std::uint8_t> bytes(nbytes);
streams.read(index, bytes.data(), nbytes);
return cursor(std::move(bytes));
}
/// @brief Deal atom-by-atom (ring / dabit / bit) onto successive stream indexes.
/// @details Stream 0 holds all ring triples when `ring_triples > 0`, then the
/// next used index holds dabits, then bit triples. Online code that
/// expects a packed `cursor` should keep using `deal_to_streams`.
inline void deal_atoms(const demand & d, net::stream_array & party0,
net::stream_array & party1)
{
std::size_t need = (d.ring_triples ? 1u : 0u) + (d.dabits ? 1u : 0u)
+ (d.bit_triples ? 1u : 0u);
if (need == 0)
return;
if (party0.size() < need || party1.size() < need)
throw std::invalid_argument("deal_atoms: not enough streams");
std::size_t idx = 0;
if (d.ring_triples)
{
factory::make_dealer(party0, party1, d.ring_triples,
[limb = d.limb] { return factory::deal_ring_triple(limb); });
++idx;
}
if (d.dabits)
{
// make_dealer always starts at stream 0; shift by writing via a
// one-stream alias when dabits are not the first kind.
if (idx == 0)
{
factory::make_dealer(party0, party1, d.dabits,
[limb = d.limb] { return factory::deal_dabit(limb); });
}
else
{
for (std::uint32_t i = 0; i < d.dabits; ++i)
{
auto v = factory::deal_dabit(d.limb);
party0.write(idx, &v.first, sizeof(v.first));
party1.write(idx, &v.second, sizeof(v.second));
}
party0.flush(idx);
party1.flush(idx);
}
++idx;
}
if (d.bit_triples)
{
if (idx == 0)
{
factory::make_dealer(party0, party1, d.bit_triples,
[] { return factory::deal_bit_triple(); });
}
else
{
for (std::uint32_t i = 0; i < d.bit_triples; ++i)
{
auto v = factory::deal_bit_triple();
party0.write(idx, &v.first, sizeof(v.first));
party1.write(idx, &v.second, sizeof(v.second));
}
party0.flush(idx);
party1.flush(idx);
}
}
}
} // namespace prep
} // namespace dpf
#endif