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>
This commit is contained in:
parent
695f8e84f7
commit
0d22946a0e
1835 changed files with 170291 additions and 2849 deletions
369
include/dpf/share_expr.hpp
Normal file
369
include/dpf/share_expr.hpp
Normal file
|
|
@ -0,0 +1,369 @@
|
|||
/// @file dpf/share_expr.hpp
|
||||
/// @brief Imperative expression recorder over composer + beaver sessions.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_SHARE_EXPR_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_SHARE_EXPR_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <exception>
|
||||
#include <map>
|
||||
#include <stdexcept>
|
||||
#include <thread>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/beaver.hpp"
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/net/edge_mesh.hpp"
|
||||
#include "dpf/net/memory_sink.hpp"
|
||||
#include "dpf/share_cmp.hpp"
|
||||
#include "dpf/trunc.hpp"
|
||||
#include "dpf/xor_wrapper.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace expr
|
||||
{
|
||||
|
||||
struct handle
|
||||
{
|
||||
std::uint32_t id = 0;
|
||||
};
|
||||
|
||||
class recorder
|
||||
{
|
||||
public:
|
||||
explicit recorder(std::size_t party, std::size_t parties = 2)
|
||||
: party_(party), parties_(parties), composer_(party)
|
||||
{
|
||||
if (parties_ != 2 && parties_ != 3)
|
||||
throw std::invalid_argument("recorder parties must be 2 or 3");
|
||||
if (party_ >= parties_)
|
||||
throw std::invalid_argument("recorder party index");
|
||||
}
|
||||
|
||||
std::size_t party() const noexcept { return party_; }
|
||||
std::size_t parties() const noexcept { return parties_; }
|
||||
protocol::composer & composer() noexcept { return composer_; }
|
||||
const protocol::composer & composer() const noexcept { return composer_; }
|
||||
|
||||
void bind_mesh(net::edge_mesh mesh) { mesh_ = std::move(mesh); }
|
||||
bool has_mesh() const noexcept { return mesh_.size() != 0; }
|
||||
|
||||
template <typename Ring>
|
||||
handle input(Ring clear, unsigned owner)
|
||||
{
|
||||
Ring mine{};
|
||||
if (parties_ == 2)
|
||||
{
|
||||
auto [s0, s1] = share_cmp::share_input(clear, owner);
|
||||
mine = (party_ == 0) ? s0 : s1;
|
||||
}
|
||||
else
|
||||
mine = (party_ == owner) ? clear : Ring{};
|
||||
values_[next_] = encode(mine);
|
||||
kinds_[next_] = kind::value;
|
||||
return handle{next_++};
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
handle bind(Ring mine)
|
||||
{
|
||||
values_[next_] = encode(mine);
|
||||
kinds_[next_] = kind::value;
|
||||
return handle{next_++};
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
handle add(handle a, handle b)
|
||||
{
|
||||
check(a);
|
||||
check(b);
|
||||
Ring va = decode<Ring>(values_.at(a.id));
|
||||
Ring vb = decode<Ring>(values_.at(b.id));
|
||||
values_[next_] = encode(static_cast<Ring>(va + vb));
|
||||
kinds_[next_] = kind::value;
|
||||
auto na = composer_.input(protocol::domain::a, sizeof(Ring));
|
||||
auto nb = composer_.input(protocol::domain::a, sizeof(Ring));
|
||||
composer_.compute(protocol::opcodes::user_base, {na, nb},
|
||||
protocol::domain::a, sizeof(Ring));
|
||||
return handle{next_++};
|
||||
}
|
||||
|
||||
/// @brief Record a product; local share filled by `complete_mul` with peer.
|
||||
template <typename Ring>
|
||||
handle mul(handle a, handle b)
|
||||
{
|
||||
check(a);
|
||||
check(b);
|
||||
pending_mul pm;
|
||||
pm.a = a.id;
|
||||
pm.b = b.id;
|
||||
pm.out = next_;
|
||||
pending_.push_back(pm);
|
||||
values_[next_] = encode(Ring{}); // filled by complete_mul
|
||||
kinds_[next_] = kind::pending_mul;
|
||||
auto na = composer_.input(protocol::domain::a, sizeof(Ring));
|
||||
auto nb = composer_.input(protocol::domain::a, sizeof(Ring));
|
||||
composer_.compute(protocol::opcodes::user_base + 2, {na, nb},
|
||||
protocol::domain::a, sizeof(Ring));
|
||||
pending_muls_++;
|
||||
return handle{next_++};
|
||||
}
|
||||
|
||||
/// @brief Finish one pending mul with the peer's operand shares (Beaver).
|
||||
template <typename Ring>
|
||||
static void complete_mul(recorder & r0, recorder & r1, handle out0,
|
||||
handle out1)
|
||||
{
|
||||
if (r0.kinds_.at(out0.id) != kind::pending_mul
|
||||
|| r1.kinds_.at(out1.id) != kind::pending_mul)
|
||||
throw std::invalid_argument("complete_mul: not pending");
|
||||
pending_mul pm0{}, pm1{};
|
||||
for (const auto & p : r0.pending_)
|
||||
if (p.out == out0.id)
|
||||
pm0 = p;
|
||||
for (const auto & p : r1.pending_)
|
||||
if (p.out == out1.id)
|
||||
pm1 = p;
|
||||
const Ring x0 = r0.local_value<Ring>(handle{pm0.a});
|
||||
const Ring y0 = r0.local_value<Ring>(handle{pm0.b});
|
||||
const Ring x1 = r1.local_value<Ring>(handle{pm1.a});
|
||||
const Ring y1 = r1.local_value<Ring>(handle{pm1.b});
|
||||
// Honest-dealer triple; each party opens only its masked limb.
|
||||
const Ring a0 = dpf::uniform_sample<Ring>();
|
||||
const Ring b0 = dpf::uniform_sample<Ring>();
|
||||
const Ring a1 = dpf::uniform_sample<Ring>();
|
||||
const Ring b1 = dpf::uniform_sample<Ring>();
|
||||
const Ring a = static_cast<Ring>(a0 + a1);
|
||||
const Ring b = static_cast<Ring>(b0 + b1);
|
||||
const Ring c = static_cast<Ring>(a * b);
|
||||
const Ring c0 = dpf::uniform_sample<Ring>();
|
||||
const Ring c1 = static_cast<Ring>(c - c0);
|
||||
const Ring d = static_cast<Ring>((x0 - a0) + (x1 - a1));
|
||||
const Ring e = static_cast<Ring>((y0 - b0) + (y1 - b1));
|
||||
const Ring z0 = trunc::mul_trunc_party(x0, y0, a0, b0, c0, d, e, 0, 0);
|
||||
const Ring z1 = trunc::mul_trunc_party(x1, y1, a1, b1, c1, d, e, 0, 1);
|
||||
r0.values_[out0.id] = encode(z0);
|
||||
r1.values_[out1.id] = encode(z1);
|
||||
r0.kinds_[out0.id] = kind::value;
|
||||
r1.kinds_[out1.id] = kind::value;
|
||||
}
|
||||
|
||||
/// @brief Multi-input AND via a bit-wire beaver monomial (2PC).
|
||||
handle band(const std::vector<handle> & bits)
|
||||
{
|
||||
if (bits.empty())
|
||||
throw std::invalid_argument("band empty");
|
||||
and_width_ = bits.size();
|
||||
pending_and pa;
|
||||
pa.ins.reserve(bits.size());
|
||||
for (auto h : bits)
|
||||
{
|
||||
check(h);
|
||||
pa.ins.push_back(h.id);
|
||||
}
|
||||
pa.out = next_;
|
||||
pending_ands_.push_back(pa);
|
||||
values_[next_] = {0};
|
||||
kinds_[next_] = kind::pending_and;
|
||||
return handle{next_++};
|
||||
}
|
||||
|
||||
static void complete_band(recorder & r0, recorder & r1, handle out0,
|
||||
handle out1)
|
||||
{
|
||||
pending_and pa0{}, pa1{};
|
||||
for (const auto & p : r0.pending_ands_)
|
||||
if (p.out == out0.id)
|
||||
pa0 = p;
|
||||
for (const auto & p : r1.pending_ands_)
|
||||
if (p.out == out1.id)
|
||||
pa1 = p;
|
||||
if (pa0.ins.size() != pa1.ins.size() || pa0.ins.empty())
|
||||
throw std::invalid_argument("complete_band");
|
||||
using bit_ring = dpf::xor_wrapper<std::uint8_t>;
|
||||
using btraits = beavers::ring_traits<bit_ring>;
|
||||
beavers::session<bit_ring> s;
|
||||
std::vector<beavers::wire<bit_ring>> ws;
|
||||
ws.reserve(pa0.ins.size());
|
||||
auto to_xw = [](std::uint8_t b) {
|
||||
return (b & 1u) ? btraits::one() : btraits::zero();
|
||||
};
|
||||
for (std::size_t i = 0; i < pa0.ins.size(); ++i)
|
||||
{
|
||||
auto w = s.bit();
|
||||
ws.push_back(w);
|
||||
const std::uint8_t b0 = r0.values_.at(pa0.ins[i]).empty()
|
||||
? 0
|
||||
: (r0.values_.at(pa0.ins[i])[0] & 1u);
|
||||
const std::uint8_t b1 = r1.values_.at(pa1.ins[i]).empty()
|
||||
? 0
|
||||
: (r1.values_.at(pa1.ins[i])[0] & 1u);
|
||||
s.bind_shares(w, to_xw(b0), to_xw(b1));
|
||||
}
|
||||
beavers::wire<bit_ring> acc = ws[0];
|
||||
for (std::size_t i = 1; i < ws.size(); ++i)
|
||||
acc = s.product(acc, ws[i]);
|
||||
s.pin(acc);
|
||||
s.sample();
|
||||
s.evaluate();
|
||||
const auto v = s.value(acc);
|
||||
r0.values_[out0.id] = {
|
||||
static_cast<std::uint8_t>(static_cast<std::uint8_t>(v.p0) & 1u)};
|
||||
r1.values_[out1.id] = {
|
||||
static_cast<std::uint8_t>(static_cast<std::uint8_t>(v.p1) & 1u)};
|
||||
r0.kinds_[out0.id] = kind::value;
|
||||
r1.kinds_[out1.id] = kind::value;
|
||||
}
|
||||
|
||||
std::size_t last_and_arity() const noexcept { return and_width_; }
|
||||
std::size_t pending_muls() const noexcept { return pending_muls_; }
|
||||
|
||||
template <typename Ring>
|
||||
Ring local_value(handle h) const
|
||||
{
|
||||
check(h);
|
||||
if (kinds_.at(h.id) != kind::value)
|
||||
throw std::logic_error("local_value: pending op");
|
||||
return decode<Ring>(values_.at(h.id));
|
||||
}
|
||||
|
||||
/// @brief Declassify after flushing any pending ops; sums opened shares.
|
||||
template <typename Ring>
|
||||
static Ring declassify(recorder & a, recorder & b, handle ha, handle hb)
|
||||
{
|
||||
if (a.kinds_.at(ha.id) == kind::pending_mul)
|
||||
complete_mul<Ring>(a, b, ha, hb);
|
||||
if (a.kinds_.at(ha.id) == kind::pending_and)
|
||||
complete_band(a, b, ha, hb);
|
||||
if (!a.has_mesh() || !b.has_mesh())
|
||||
throw std::logic_error("expr declassify requires bind_mesh");
|
||||
protocol::composer c0(a.party_);
|
||||
protocol::composer c1(b.party_);
|
||||
auto in0 = c0.input(protocol::domain::a, sizeof(Ring));
|
||||
auto ex0 = c0.exchange(in0);
|
||||
auto in1 = c1.input(protocol::domain::a, sizeof(Ring));
|
||||
auto ex1 = c1.exchange(in1);
|
||||
auto p0 = c0.default_plan();
|
||||
auto p1 = c1.default_plan();
|
||||
std::vector<std::vector<std::uint8_t>> v0(p0.nodes().size());
|
||||
std::vector<std::vector<std::uint8_t>> v1(p1.nodes().size());
|
||||
const Ring mine_a = a.local_value<Ring>(ha);
|
||||
const Ring mine_b = b.local_value<Ring>(hb);
|
||||
v0[in0.id].assign(sizeof(Ring), 0);
|
||||
v1[in1.id].assign(sizeof(Ring), 0);
|
||||
std::memcpy(v0[in0.id].data(), &mine_a, sizeof(Ring));
|
||||
std::memcpy(v1[in1.id].data(), &mine_b, sizeof(Ring));
|
||||
std::map<std::uint32_t, protocol::kernel_fn> kernels;
|
||||
std::exception_ptr err;
|
||||
std::thread t0([&] {
|
||||
try
|
||||
{
|
||||
protocol::drive_via_schedule(p0, a.mesh_.at(net::edge_peer), v0,
|
||||
kernels, a.party_);
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
err = std::current_exception();
|
||||
}
|
||||
});
|
||||
std::thread t1([&] {
|
||||
try
|
||||
{
|
||||
protocol::drive_via_schedule(p1, b.mesh_.at(net::edge_peer), v1,
|
||||
kernels, b.party_);
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
err = std::current_exception();
|
||||
}
|
||||
});
|
||||
t0.join();
|
||||
t1.join();
|
||||
if (err)
|
||||
std::rethrow_exception(err);
|
||||
Ring opened{};
|
||||
std::memcpy(&opened, v0[ex0.id].data(), sizeof(Ring));
|
||||
(void)ex1;
|
||||
return opened;
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
void install_tape(const beavers::party_tape<Ring> & tape)
|
||||
{
|
||||
session_.install_party(static_cast<unsigned>(party_), tape);
|
||||
}
|
||||
|
||||
beavers::session<std::uint64_t> & session_u64() noexcept { return session_; }
|
||||
|
||||
protocol::plan plan() const { return composer_.default_plan(); }
|
||||
|
||||
private:
|
||||
enum class kind : unsigned char { value, pending_mul, pending_and };
|
||||
|
||||
struct pending_mul
|
||||
{
|
||||
std::uint32_t a = 0;
|
||||
std::uint32_t b = 0;
|
||||
std::uint32_t out = 0;
|
||||
};
|
||||
|
||||
struct pending_and
|
||||
{
|
||||
std::vector<std::uint32_t> ins;
|
||||
std::uint32_t out = 0;
|
||||
};
|
||||
|
||||
void check(handle h) const
|
||||
{
|
||||
if (values_.find(h.id) == values_.end())
|
||||
throw std::out_of_range("expr handle");
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
static std::vector<std::uint8_t> encode(Ring v)
|
||||
{
|
||||
std::vector<std::uint8_t> out(sizeof(Ring));
|
||||
std::memcpy(out.data(), &v, sizeof(Ring));
|
||||
return out;
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
static Ring decode(const std::vector<std::uint8_t> & b)
|
||||
{
|
||||
if (b.size() < sizeof(Ring))
|
||||
throw std::invalid_argument("expr decode");
|
||||
Ring v{};
|
||||
std::memcpy(&v, b.data(), sizeof(Ring));
|
||||
return v;
|
||||
}
|
||||
|
||||
std::size_t party_ = 0;
|
||||
std::size_t parties_ = 2;
|
||||
protocol::composer composer_;
|
||||
net::edge_mesh mesh_{};
|
||||
std::uint32_t next_ = 1;
|
||||
std::map<std::uint32_t, std::vector<std::uint8_t>> values_;
|
||||
std::map<std::uint32_t, kind> kinds_;
|
||||
beavers::session<std::uint64_t> session_;
|
||||
std::vector<pending_mul> pending_;
|
||||
std::vector<pending_and> pending_ands_;
|
||||
std::size_t and_width_ = 0;
|
||||
std::size_t pending_muls_ = 0;
|
||||
};
|
||||
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Ring eval_mul_add(Ring x, Ring y, Ring w)
|
||||
{
|
||||
return static_cast<Ring>(x * y + w);
|
||||
}
|
||||
|
||||
} // namespace expr
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_SHARE_EXPR_HPP__
|
||||
Loading…
Add table
Add a link
Reference in a new issue