libdpf/include/dpf/share_expr.hpp

370 lines
12 KiB
C++
Raw Normal View History

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