libdpf/include/dpf/gilboa.hpp

278 lines
9.3 KiB
C++
Raw Normal View History

/// @file dpf/gilboa.hpp
/// @brief One-off Gilboa multiplication from correlated bit×ring triples.
#ifndef LIBDPF_INCLUDE_DPF_GILBOA_HPP__
#define LIBDPF_INCLUDE_DPF_GILBOA_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/beaver.hpp"
#include "dpf/edabit.hpp"
#include "dpf/ot_pack.hpp"
#include "dpf/random.hpp"
namespace dpf
{
namespace gilboa
{
/// @brief Clear product of additive shares (oracle only).
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring mul_clear(Ring x0, Ring x1, Ring y0, Ring y1)
{
return static_cast<Ring>((x0 + x1) * (y0 + y1));
}
template <typename Ring>
struct product_shares
{
Ring z0{};
Ring z1{};
};
/// @brief Oracle: reconstruct, multiply, re-share. Not a party protocol.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
product_shares<Ring> mul_dealer(Ring x0, Ring x1, Ring y0, Ring y1)
{
const Ring z = mul_clear(x0, x1, y0, y1);
const Ring m = dpf::uniform_sample<Ring>();
return product_shares<Ring>{m, static_cast<Ring>(z - m)};
}
/// @brief Per-bit messages one party produces before the peer exchange.
/// @details Factor `x` must be held as bits (party 1 share 0, or XOR bit
/// shares). Additive `x` shares are bit-sliced locally — correct when
/// one party holds the clear factor.
template <typename Ring>
struct gilboa_round
{
std::vector<Ring> d_share; ///< additive share of x_bit - a
std::vector<Ring> e_share; ///< additive share of y - b
std::vector<ot::bit_ring_triple<Ring>> triples;
Ring x_share{};
Ring y_share{};
unsigned party = 0;
};
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
gilboa_round<Ring> mul_from_ot_begin(ot::pack & pack, Ring x_share, Ring y_share,
unsigned party, unsigned bits = 64)
{
if (bits == 0 || bits > 8u * sizeof(Ring))
throw std::invalid_argument("gilboa bits");
if (pack.remaining_bit_ring() < bits)
throw std::runtime_error("gilboa::mul_from_ot: need bit×ring triples");
gilboa_round<Ring> r;
r.x_share = x_share;
r.y_share = y_share;
r.party = party;
r.d_share.resize(bits);
r.e_share.resize(bits);
r.triples.reserve(bits);
for (unsigned i = 0; i < bits; ++i)
{
auto t = pack.take_bit_ring<Ring>();
const Ring x_bit = static_cast<Ring>(
(static_cast<std::uint64_t>(x_share) >> i) & 1u);
r.d_share[i] = static_cast<Ring>(x_bit - Ring{t.a});
r.e_share[i] = static_cast<Ring>(y_share - t.b);
r.triples.push_back(t);
}
return r;
}
/// @brief Finish after exchanging additive `d` and `e` with the peer.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring mul_from_ot_finish(const gilboa_round<Ring> & r,
const std::vector<Ring> & peer_d, const std::vector<Ring> & peer_e)
{
if (peer_d.size() != r.d_share.size() || peer_e.size() != r.e_share.size())
throw std::invalid_argument("gilboa finish size");
Ring acc{};
for (std::size_t i = 0; i < r.d_share.size(); ++i)
{
const auto & t = r.triples[i];
const Ring d = static_cast<Ring>(r.d_share[i] + peer_d[i]);
const Ring e = static_cast<Ring>(r.e_share[i] + peer_e[i]);
Ring z = t.c;
z = static_cast<Ring>(z + d * t.b);
z = static_cast<Ring>(z + e * Ring{t.a});
if (r.party == 0)
z = static_cast<Ring>(z + d * e);
acc = static_cast<Ring>(acc + (z << i));
}
return acc;
}
/// @brief One bit×ring product of an additive 0/1 value with `y`.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
product_shares<Ring> mul_bit_from_ot(ot::pack & pack0, ot::pack & pack1,
Ring bit0, Ring bit1, Ring y0, Ring y1)
{
auto t0 = pack0.template take_bit_ring<Ring>();
auto t1 = pack1.template take_bit_ring<Ring>();
const Ring d0 = static_cast<Ring>(bit0 - Ring{t0.a});
const Ring d1 = static_cast<Ring>(bit1 - Ring{t1.a});
const Ring e0 = static_cast<Ring>(y0 - t0.b);
const Ring e1 = static_cast<Ring>(y1 - t1.b);
const Ring d = static_cast<Ring>(d0 + d1);
const Ring e = static_cast<Ring>(e0 + e1);
auto acc = [&](const ot::bit_ring_triple<Ring> & t, unsigned party) {
Ring z = t.c;
z = static_cast<Ring>(z + d * t.b);
z = static_cast<Ring>(z + e * Ring{t.a});
if (party == 0)
z = static_cast<Ring>(z + d * e);
return z;
};
return product_shares<Ring>{acc(t0, 0), acc(t1, 1)};
}
/// @brief General additive `x`: A2B, daBit B2A per bit, then bit×ring Gilboa.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
product_shares<Ring> mul_from_ot_pair(ot::pack & pack0, ot::pack & pack1,
Ring x0, Ring x1, Ring y0, Ring y1, unsigned bits = 64)
{
if (pack0.remaining_b2a() < bits || pack1.remaining_b2a() < bits
|| pack0.remaining_bit_ring() < bits || pack1.remaining_bit_ring() < bits)
throw std::runtime_error("gilboa::mul_from_ot_pair: need dabits and bit×ring");
auto eda = edabit::sample_edabit_pair<Ring>(bits);
auto bits_xy = edabit::a2b_gmw_pair(eda, x0, x1);
const auto & bx0 = bits_xy.first;
const auto & bx1 = bits_xy.second;
Ring z0{}, z1{};
for (unsigned i = 0; i < bits; ++i)
{
auto d0 = pack0.template take_dabit<Ring>();
auto d1 = pack1.template take_dabit<Ring>();
const std::uint8_t b0 = edabit::detail::get_bit(bx0, i);
const std::uint8_t b1 = edabit::detail::get_bit(bx1, i);
const std::uint8_t mask = static_cast<std::uint8_t>(
(b0 ^ d0.bit) ^ (b1 ^ d1.bit));
const Ring a0 = edabit::b2a_party_bit(b0, d0, mask, 0, 0);
const Ring a1 = edabit::b2a_party_bit(b1, d1, mask, 1, 0);
auto part = mul_bit_from_ot(pack0, pack1, a0, a1, y0, y1);
z0 = static_cast<Ring>(z0 + (part.z0 << i));
z1 = static_cast<Ring>(z1 + (part.z1 << i));
}
return product_shares<Ring>{z0, z1};
}
/// @brief Convenience: one party after peer messages are known (same as finish).
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring mul_from_ot(ot::pack & pack, Ring x_share, Ring y_share, unsigned party,
const std::vector<Ring> & peer_d, const std::vector<Ring> & peer_e,
unsigned bits = 64)
{
auto r = mul_from_ot_begin(pack, x_share, y_share, party, bits);
return mul_from_ot_finish(r, peer_d, peer_e);
}
/// @brief Honest-dealer tape: `session::sample` then `export_party`.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
beavers::party_tape<Ring> fill_tape_dealer(beavers::session<Ring> & s,
unsigned party)
{
s.sample();
return s.export_party(party);
}
/// @brief Sample λ from correlated bit×ring `b` shares; monomials are Π λ^e.
/// @details One λ per wire (repeated factors share it). Both packs are consumed
/// in lockstep. `sample` is driven by those triples, not overwritten after.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
std::pair<beavers::party_tape<Ring>, beavers::party_tape<Ring>>
fill_tape_ot_pair(beavers::session<Ring> & s, ot::pack & pack0, ot::pack & pack1)
{
std::size_t calls = 0;
s.sample([&]() {
++calls;
return Ring{1};
});
std::size_t nblind = 0;
{
const auto probe = s.export_party(0);
for (auto ready : probe.lambda_ready)
nblind += ready ? 1u : 0u;
}
s.clear_sample();
struct sampler
{
ot::pack * p0;
ot::pack * p1;
std::size_t nblind;
std::size_t call = 0;
Ring saved{};
Ring operator()()
{
if (p0->remaining_bit_ring() == 0 || p1->remaining_bit_ring() == 0)
throw std::runtime_error("fill_tape_ot: need bit×ring triples");
if (call < nblind * 2)
{
if ((call & 1u) == 0)
{
auto t0 = p0->template take_bit_ring<Ring>();
auto t1 = p1->template take_bit_ring<Ring>();
saved = t0.b;
++call;
return static_cast<Ring>(t0.b + t1.b);
}
++call;
return saved;
}
auto t0 = p0->template take_bit_ring<Ring>();
auto t1 = p1->template take_bit_ring<Ring>();
++call;
(void)t1;
return t0.c;
}
};
s.sample(sampler{&pack0, &pack1, nblind});
(void)calls;
auto t0 = s.export_party(0);
auto t1 = s.export_party(1);
pack0.stash_party_tape(t0);
pack1.stash_party_tape(t1);
return {t0, t1};
}
/// @brief Read the party view `fill_tape_ot_pair` stashed in `pack`.
/// @details Does not sample. A pack that was not produced by the pair sampler
/// has no view to read.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
beavers::party_tape<Ring> fill_tape_ot(beavers::session<Ring> & s, ot::pack & pack,
unsigned party)
{
if (party > 1)
throw std::invalid_argument("fill_tape_ot party");
if (!pack.has_party_tape())
throw std::runtime_error(
"fill_tape_ot: pack has no party view; call fill_tape_ot_pair");
auto tape = pack.template load_party_tape<Ring>();
if (tape.lambda.size() != s.wire_count())
throw std::invalid_argument("fill_tape_ot: tape does not match session");
(void)party;
return tape;
}
} // namespace gilboa
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_GILBOA_HPP__