libdpf/include/dpf/edabit.hpp

347 lines
12 KiB
C++
Raw Permalink Normal View History

/// @file dpf/edabit.hpp
/// @brief daBits and edaBits for arithmetic ↔ boolean conversion.
#ifndef LIBDPF_INCLUDE_DPF_EDABIT_HPP__
#define LIBDPF_INCLUDE_DPF_EDABIT_HPP__
#include <cstddef>
#include <cstdint>
#include <stdexcept>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/ot_pack.hpp"
#include "dpf/rss_seed.hpp"
namespace dpf
{
namespace edabit
{
/// @brief Packed boolean XOR shares + arithmetic share of `r = sum b_i 2^i`.
template <typename Ring = std::uint64_t>
struct edabit_share
{
std::vector<std::uint8_t> bits_packed; ///< ceil(ell/8) XOR / RSS-own bits
std::vector<std::uint8_t> bits_next; ///< RSS next bit component; empty in 2PC
Ring arith{};
Ring arith_next{}; ///< RSS next component; 0 for 2PC
unsigned width = 0;
};
namespace detail
{
inline void set_bit(std::vector<std::uint8_t> & packed, unsigned i,
std::uint8_t bit)
{
const unsigned byte = i / 8u;
const unsigned off = i % 8u;
if (byte >= packed.size())
packed.resize(byte + 1u, 0);
if (bit & 1u)
packed[byte] = static_cast<std::uint8_t>(packed[byte] | (1u << off));
else
packed[byte] = static_cast<std::uint8_t>(packed[byte] & ~(1u << off));
}
inline std::uint8_t get_bit(const std::vector<std::uint8_t> & packed, unsigned i)
{
const unsigned byte = i / 8u;
const unsigned off = i % 8u;
if (byte >= packed.size())
return 0;
return static_cast<std::uint8_t>((packed[byte] >> off) & 1u);
}
} // namespace detail
template <typename Ring = std::uint64_t>
struct edabit_pair
{
edabit_share<Ring> p0;
edabit_share<Ring> p1;
Ring clear_r{};
};
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
edabit_pair<Ring> sample_edabit_pair(unsigned ell)
{
if (ell == 0 || ell > 8u * sizeof(Ring))
throw std::invalid_argument("edabit width");
edabit_pair<Ring> out;
out.p0.width = ell;
out.p1.width = ell;
const std::size_t nbytes = (ell + 7u) / 8u;
out.p0.bits_packed.assign(nbytes, 0);
out.p1.bits_packed.assign(nbytes, 0);
Ring r = Ring{};
for (unsigned i = 0; i < ell; ++i)
{
auto d = ot::sample_dabit_pair<Ring>();
detail::set_bit(out.p0.bits_packed, i, d.p0.bit);
detail::set_bit(out.p1.bits_packed, i, d.p1.bit);
out.p0.arith = static_cast<Ring>(out.p0.arith + (d.p0.arith << i));
out.p1.arith = static_cast<Ring>(out.p1.arith + (d.p1.arith << i));
const Ring bit = static_cast<Ring>((d.p0.bit ^ d.p1.bit) & 1u);
r = static_cast<Ring>(r + (bit << i));
}
out.clear_r = r;
return out;
}
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
edabit_share<Ring> sample_from_pack(ot::pack & pack, unsigned ell)
{
if (ell == 0 || ell > 8u * sizeof(Ring))
throw std::invalid_argument("edabit width");
edabit_share<Ring> out;
out.width = ell;
out.bits_packed.assign((ell + 7u) / 8u, 0);
for (unsigned i = 0; i < ell; ++i)
{
auto d = pack.take_dabit<Ring>();
detail::set_bit(out.bits_packed, i, d.bit);
out.arith = static_cast<Ring>(out.arith + (d.arith << i));
}
return out;
}
/// @brief Three parties' RSS edaBits: bits first, arith = sum bit_comp · 2^i.
template <typename Ring = std::uint64_t>
struct edabit_rss_triple
{
edabit_share<Ring> p0;
edabit_share<Ring> p1;
edabit_share<Ring> p2;
Ring clear_r{};
};
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
edabit_rss_triple<Ring> sample_rss_all(const rss::seed_bundle & bundle,
unsigned ell, std::uint64_t index)
{
if (ell == 0 || ell > 8u * sizeof(Ring))
throw std::invalid_argument("edabit width");
edabit_rss_triple<Ring> out;
out.p0.width = out.p1.width = out.p2.width = ell;
const std::size_t nbytes = (ell + 7u) / 8u;
out.p0.bits_packed.assign(nbytes, 0);
out.p1.bits_packed.assign(nbytes, 0);
out.p2.bits_packed.assign(nbytes, 0);
out.p0.bits_next.assign(nbytes, 0);
out.p1.bits_next.assign(nbytes, 0);
out.p2.bits_next.assign(nbytes, 0);
Ring r = Ring{};
for (unsigned i = 0; i < ell; ++i)
{
// Boolean RSS: own components XOR to the bit; next matches the neighbor.
const auto rnd = rss::random_replicated_all<std::uint8_t>(
bundle, index + 1 + i);
const std::uint8_t bit = static_cast<std::uint8_t>(
(rnd.p0.own ^ rnd.p1.own ^ rnd.p2.own) & 1u);
const std::uint8_t u = static_cast<std::uint8_t>(
dpf::uniform_sample<std::uint8_t>() & 1u);
const std::uint8_t v = static_cast<std::uint8_t>(
dpf::uniform_sample<std::uint8_t>() & 1u);
const std::uint8_t w = static_cast<std::uint8_t>(u ^ v ^ bit);
detail::set_bit(out.p0.bits_packed, i, u);
detail::set_bit(out.p0.bits_next, i, v);
detail::set_bit(out.p1.bits_packed, i, v);
detail::set_bit(out.p1.bits_next, i, w);
detail::set_bit(out.p2.bits_packed, i, w);
detail::set_bit(out.p2.bits_next, i, u);
// Arithmetic RSS of the same bit: own components sum to `bit`.
const Ring ra = dpf::uniform_sample<Ring>();
const Ring rb = dpf::uniform_sample<Ring>();
const Ring rc = static_cast<Ring>(Ring{bit} - ra - rb);
out.p0.arith = static_cast<Ring>(out.p0.arith + (ra << i));
out.p0.arith_next = static_cast<Ring>(out.p0.arith_next + (rb << i));
out.p1.arith = static_cast<Ring>(out.p1.arith + (rb << i));
out.p1.arith_next = static_cast<Ring>(out.p1.arith_next + (rc << i));
out.p2.arith = static_cast<Ring>(out.p2.arith + (rc << i));
out.p2.arith_next = static_cast<Ring>(out.p2.arith_next + (ra << i));
r = static_cast<Ring>(r + (Ring{bit} << i));
}
out.clear_r = r;
return out;
}
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
edabit_share<Ring> sample_rss(const rss::seed_bundle & bundle, unsigned me,
unsigned ell, std::uint64_t index)
{
auto all = sample_rss_all<Ring>(bundle, ell, index);
if (me == 0)
return all.p0;
if (me == 1)
return all.p1;
if (me == 2)
return all.p2;
throw std::invalid_argument("sample_rss party");
}
/// @brief Finish one GMW AND. `d_open` / `e_open` are the public `p⊕a` and `q⊕b`.
inline std::uint8_t and_finish(const ot::bit_triple & mine, std::uint8_t d_open,
std::uint8_t e_open, unsigned party)
{
std::uint8_t z = mine.c;
z = static_cast<std::uint8_t>(z ^ (d_open & mine.b));
z = static_cast<std::uint8_t>(z ^ (e_open & mine.a));
if (party == 0)
z = static_cast<std::uint8_t>(z ^ (d_open & e_open));
return static_cast<std::uint8_t>(z & 1u);
}
/// @brief Two-party AND: open the masked bits, then `and_finish` on each view.
HEDLEY_WARN_UNUSED_RESULT
inline std::pair<std::uint8_t, std::uint8_t> and_pair(std::uint8_t p0,
std::uint8_t p1, std::uint8_t q0, std::uint8_t q1)
{
auto tp = ot::sample_bit_triple_pair();
const std::uint8_t d = static_cast<std::uint8_t>(
(p0 ^ tp.p0.a) ^ (p1 ^ tp.p1.a));
const std::uint8_t e = static_cast<std::uint8_t>(
(q0 ^ tp.p0.b) ^ (q1 ^ tp.p1.b));
return {and_finish(tp.p0, d, e, 0), and_finish(tp.p1, d, e, 1)};
}
/// @brief A2B whose carry is a shared AND. Opens `x - r` and the AND masks only.
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
std::pair<std::vector<std::uint8_t>, std::vector<std::uint8_t>> a2b_gmw_pair(
const edabit_pair<Ring> & eda, Ring x0, Ring x1)
{
const unsigned ell = eda.p0.width;
const Ring delta = static_cast<Ring>(
(x0 - eda.p0.arith) + (x1 - eda.p1.arith));
std::vector<std::uint8_t> b0((ell + 7u) / 8u, 0);
std::vector<std::uint8_t> b1((ell + 7u) / 8u, 0);
std::uint8_t c0 = 0;
std::uint8_t c1 = 0;
for (unsigned i = 0; i < ell; ++i)
{
const std::uint8_t r0 = detail::get_bit(eda.p0.bits_packed, i);
const std::uint8_t r1 = detail::get_bit(eda.p1.bits_packed, i);
const std::uint8_t di = static_cast<std::uint8_t>(
(static_cast<std::uint64_t>(delta) >> i) & 1u);
detail::set_bit(b0, i, static_cast<std::uint8_t>(r0 ^ di ^ c0));
detail::set_bit(b1, i, static_cast<std::uint8_t>(r1 ^ c1));
const std::uint8_t rd0 = di ? r0 : 0;
const std::uint8_t rd1 = di ? r1 : 0;
const std::uint8_t dc0 = di ? c0 : 0;
const std::uint8_t dc1 = di ? c1 : 0;
auto rc = and_pair(r0, r1, c0, c1);
c0 = static_cast<std::uint8_t>(rd0 ^ dc0 ^ rc.first);
c1 = static_cast<std::uint8_t>(rd1 ^ dc1 ^ rc.second);
}
return {std::move(b0), std::move(b1)};
}
/// @brief Unsigned compare of XOR bit-shares. The predicate stays shared.
HEDLEY_WARN_UNUSED_RESULT
inline std::pair<std::uint8_t, std::uint8_t> gt_bits_pair(
const std::vector<std::uint8_t> & x0, const std::vector<std::uint8_t> & x1,
const std::vector<std::uint8_t> & y0, const std::vector<std::uint8_t> & y1,
unsigned n)
{
std::uint8_t gt0 = 0;
std::uint8_t gt1 = 0;
std::uint8_t eq0 = 1;
std::uint8_t eq1 = 0;
for (unsigned k = n; k-- > 0; )
{
const std::uint8_t xb0 = detail::get_bit(x0, k);
const std::uint8_t xb1 = detail::get_bit(x1, k);
const std::uint8_t yb0 = detail::get_bit(y0, k);
const std::uint8_t yb1 = detail::get_bit(y1, k);
auto xny = and_pair(xb0, xb1, static_cast<std::uint8_t>(yb0 ^ 1u), yb1);
auto bit = and_pair(xny.first, xny.second, eq0, eq1);
gt0 = static_cast<std::uint8_t>(gt0 ^ bit.first);
gt1 = static_cast<std::uint8_t>(gt1 ^ bit.second);
auto eq = and_pair(eq0, eq1, static_cast<std::uint8_t>(xb0 ^ yb0 ^ 1u),
static_cast<std::uint8_t>(xb1 ^ yb1));
eq0 = eq.first;
eq1 = eq.second;
}
return {gt0, gt1};
}
inline std::uint64_t reconstruct_bits(
const std::vector<std::uint8_t> & a,
const std::vector<std::uint8_t> & b, unsigned ell)
{
std::uint64_t v = 0;
for (unsigned i = 0; i < ell; ++i)
{
const std::uint8_t bit = static_cast<std::uint8_t>(
detail::get_bit(a, i) ^ detail::get_bit(b, i));
v |= (static_cast<std::uint64_t>(bit) << i);
}
return v;
}
/// @brief One party's B2A contribution after opening `mask = b ⊕ r` per bit.
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
Ring b2a_party_bit(std::uint8_t b_share, const ot::dabit<Ring> & r,
std::uint8_t mask_open, unsigned party, unsigned shift)
{
// b = r ⊕ mask. Arithmetic: r_arith + mask * (1 - 2*r_bit_as...)
// Standard: open c = b ⊕ r; then [b] = [r] + c - 2c[r] for arith r in {0,1}.
// With XOR r_bit matching r_arith:
// share = r.arith + (party==0 ? Ring{mask_open} : 0)
// - Ring{2} * Ring{mask_open} * Ring{r.bit}
// but r.bit is XOR-shared: use r.arith which equals the bit value shares.
Ring s = r.arith;
if (party == 0)
s = static_cast<Ring>(s + Ring{mask_open});
// Subtract 2 * mask * r: each party subtracts 2*mask*r.arith (additive r).
s = static_cast<Ring>(s - static_cast<Ring>(Ring{2} * Ring{mask_open} * r.arith));
(void)b_share;
return static_cast<Ring>(s << shift);
}
/// @brief Two-party B2A: each uses its bit share + dabits; open mask = b⊕r.
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
std::pair<Ring, Ring> b2a_pair(const std::vector<std::uint8_t> & bits0,
const std::vector<std::uint8_t> & bits1, unsigned ell,
ot::pack & pack0, ot::pack & pack1)
{
if (pack0.remaining_b2a() < ell || pack1.remaining_b2a() < ell)
throw std::runtime_error("b2a_pair: need dabits");
Ring s0{}, s1{};
for (unsigned i = 0; i < ell; ++i)
{
auto d0 = pack0.take_dabit<Ring>();
auto d1 = pack1.take_dabit<Ring>();
const std::uint8_t b0 = detail::get_bit(bits0, i);
const std::uint8_t b1 = detail::get_bit(bits1, i);
const std::uint8_t mask = static_cast<std::uint8_t>(
(b0 ^ d0.bit) ^ (b1 ^ d1.bit));
s0 = static_cast<Ring>(s0 + b2a_party_bit(b0, d0, mask, 0, i));
s1 = static_cast<Ring>(s1 + b2a_party_bit(b1, d1, mask, 1, i));
}
return {s0, s1};
}
/// @brief Oracle expected value (not a party protocol).
template <typename Ring = std::uint64_t>
HEDLEY_WARN_UNUSED_RESULT
Ring b2a_clear(const std::vector<std::uint8_t> & bits0,
const std::vector<std::uint8_t> & bits1, unsigned ell)
{
return static_cast<Ring>(reconstruct_bits(bits0, bits1, ell));
}
} // namespace edabit
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_EDABIT_HPP__