394 lines
12 KiB
C++
394 lines
12 KiB
C++
|
|
/// @file dpf/share_cmp.hpp
|
||
|
|
/// @brief Comparison of two arithmetic shares (mask, open, MSB / mux).
|
||
|
|
#ifndef LIBDPF_INCLUDE_DPF_SHARE_CMP_HPP__
|
||
|
|
#define LIBDPF_INCLUDE_DPF_SHARE_CMP_HPP__
|
||
|
|
|
||
|
|
#include <cstddef>
|
||
|
|
#include <cstdint>
|
||
|
|
#include <stdexcept>
|
||
|
|
#include <utility>
|
||
|
|
#include <vector>
|
||
|
|
|
||
|
|
#include "hedley/hedley.h"
|
||
|
|
|
||
|
|
#include "dpf/beaver.hpp"
|
||
|
|
#include "dpf/edabit.hpp"
|
||
|
|
#include "dpf/random.hpp"
|
||
|
|
#include "dpf/trunc.hpp"
|
||
|
|
|
||
|
|
namespace dpf
|
||
|
|
{
|
||
|
|
namespace share_cmp
|
||
|
|
{
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
constexpr std::uint8_t msb_clear(Ring x, unsigned n) noexcept
|
||
|
|
{
|
||
|
|
if (n == 0)
|
||
|
|
return 0;
|
||
|
|
return static_cast<std::uint8_t>((x >> (n - 1u)) & 1u);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
constexpr std::uint8_t gt_clear(Ring x, Ring y, unsigned n) noexcept
|
||
|
|
{
|
||
|
|
const Ring mask = trunc::mask_bits<Ring>(n);
|
||
|
|
return static_cast<std::uint8_t>(((x & mask) > (y & mask)) ? 1 : 0);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_NO_THROW
|
||
|
|
constexpr std::uint8_t eq_clear(Ring x, Ring y) noexcept
|
||
|
|
{
|
||
|
|
return static_cast<std::uint8_t>(x == y ? 1 : 0);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
struct msb_prep
|
||
|
|
{
|
||
|
|
Ring r0{};
|
||
|
|
Ring r1{};
|
||
|
|
};
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
msb_prep<Ring> sample_msb_prep(unsigned /*n*/)
|
||
|
|
{
|
||
|
|
msb_prep<Ring> p;
|
||
|
|
const Ring r = dpf::uniform_sample<Ring>();
|
||
|
|
p.r0 = dpf::uniform_sample<Ring>();
|
||
|
|
p.r1 = static_cast<Ring>(r - p.r0);
|
||
|
|
return p;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief This party's XOR share of MSB after A2B (the bit at `n-1`).
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::uint8_t msb_party(const std::vector<std::uint8_t> & my_bits, unsigned n)
|
||
|
|
{
|
||
|
|
if (n == 0)
|
||
|
|
return 0;
|
||
|
|
return edabit::detail::get_bit(my_bits, n - 1u);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<std::uint8_t, std::uint8_t> msb_party_pair(Ring x0, Ring x1,
|
||
|
|
const msb_prep<Ring> & prep, unsigned n)
|
||
|
|
{
|
||
|
|
(void)prep;
|
||
|
|
auto eda = edabit::sample_edabit_pair<Ring>(n);
|
||
|
|
auto [b0, b1] = edabit::a2b_gmw_pair(eda, x0, x1);
|
||
|
|
return {msb_party<Ring>(b0, n), msb_party<Ring>(b1, n)};
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief This party's XOR share of an unsigned greater-than bit.
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::uint8_t gt_party(std::uint8_t my_bit) noexcept
|
||
|
|
{
|
||
|
|
return static_cast<std::uint8_t>(my_bit & 1u);
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Full-width unsigned compare. Bit shares stay split; their XOR is `x > y`.
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<std::uint8_t, std::uint8_t> gt_party_pair(Ring x0, Ring x1, Ring y0,
|
||
|
|
Ring y1, const msb_prep<Ring> & prep, unsigned n)
|
||
|
|
{
|
||
|
|
(void)prep;
|
||
|
|
auto ex = edabit::sample_edabit_pair<Ring>(n);
|
||
|
|
auto ey = edabit::sample_edabit_pair<Ring>(n);
|
||
|
|
auto bx = edabit::a2b_gmw_pair(ex, x0, x1);
|
||
|
|
auto by = edabit::a2b_gmw_pair(ey, y0, y1);
|
||
|
|
auto g = edabit::gt_bits_pair(bx.first, bx.second, by.first, by.second, n);
|
||
|
|
return {gt_party<Ring>(g.first), gt_party<Ring>(g.second)};
|
||
|
|
}
|
||
|
|
|
||
|
|
namespace detail
|
||
|
|
{
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> b2a_bit(std::uint8_t b0, std::uint8_t b1)
|
||
|
|
{
|
||
|
|
auto dab = ot::sample_dabit_pair<Ring>();
|
||
|
|
const std::uint8_t mask = static_cast<std::uint8_t>(
|
||
|
|
(b0 ^ dab.p0.bit) ^ (b1 ^ dab.p1.bit));
|
||
|
|
return {edabit::b2a_party_bit<Ring>(b0, dab.p0, mask, 0, 0),
|
||
|
|
edabit::b2a_party_bit<Ring>(b1, dab.p1, mask, 1, 0)};
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> bit_mul_shares(Ring b0, Ring b1, Ring x0, Ring x1)
|
||
|
|
{
|
||
|
|
beavers::session<Ring> s;
|
||
|
|
auto bw = s.bit();
|
||
|
|
auto xw = s.input();
|
||
|
|
auto zw = s.bit_mul(bw, xw);
|
||
|
|
s.pin(zw);
|
||
|
|
s.sample();
|
||
|
|
s.bind_shares(bw, b0, b1);
|
||
|
|
s.bind_shares(xw, x0, x1);
|
||
|
|
s.evaluate();
|
||
|
|
const auto v = s.value(zw);
|
||
|
|
return {v.p0, v.p1};
|
||
|
|
}
|
||
|
|
|
||
|
|
} // namespace detail
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring relu_clear(Ring x0, Ring x1, unsigned n)
|
||
|
|
{
|
||
|
|
const Ring x = static_cast<Ring>(x0 + x1);
|
||
|
|
return msb_clear(x, n) ? Ring{} : x;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief ReLU party share: `(1 - msb) · x` via `bit_mul` on additive bit shares.
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring relu_party(Ring x_share, Ring keep_arith)
|
||
|
|
{
|
||
|
|
return static_cast<Ring>(keep_arith * x_share);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> relu_party_pair(Ring x0, Ring x1, unsigned n)
|
||
|
|
{
|
||
|
|
auto prep = sample_msb_prep<Ring>(n);
|
||
|
|
auto [b0, b1] = msb_party_pair(x0, x1, prep, n);
|
||
|
|
auto [s0, s1] = detail::b2a_bit<Ring>(b0, b1);
|
||
|
|
const Ring k0 = static_cast<Ring>(Ring{1} - s0);
|
||
|
|
const Ring k1 = static_cast<Ring>(Ring{} - s1);
|
||
|
|
return detail::bit_mul_shares(k0, k1, x0, x1);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring max_clear(Ring x0, Ring x1, Ring y0, Ring y1, unsigned n)
|
||
|
|
{
|
||
|
|
const Ring x = static_cast<Ring>(x0 + x1);
|
||
|
|
const Ring y = static_cast<Ring>(y0 + y1);
|
||
|
|
return gt_clear(x, y, n) ? x : y;
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief `max = y + b·(x-y)` with `b` an additive share of the comparison bit.
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring max_party(Ring y_share, Ring select_share)
|
||
|
|
{
|
||
|
|
return static_cast<Ring>(y_share + select_share);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> max_party_pair(Ring x0, Ring x1, Ring y0, Ring y1,
|
||
|
|
unsigned n)
|
||
|
|
{
|
||
|
|
auto prep = sample_msb_prep<Ring>(n);
|
||
|
|
auto [g0, g1] = gt_party_pair(x0, x1, y0, y1, prep, n);
|
||
|
|
auto [s0, s1] = detail::b2a_bit<Ring>(g0, g1);
|
||
|
|
const Ring d0 = static_cast<Ring>(x0 - y0);
|
||
|
|
const Ring d1 = static_cast<Ring>(x1 - y1);
|
||
|
|
auto sel = detail::bit_mul_shares(s0, s1, d0, d1);
|
||
|
|
return {max_party(y0, sel.first), max_party(y1, sel.second)};
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> mux_clear(std::uint8_t sel, Ring a0, Ring a1, Ring b0,
|
||
|
|
Ring b1)
|
||
|
|
{
|
||
|
|
const Ring a = static_cast<Ring>(a0 + a1);
|
||
|
|
const Ring b = static_cast<Ring>(b0 + b1);
|
||
|
|
const Ring z = sel ? a : b;
|
||
|
|
const Ring m = dpf::uniform_sample<Ring>();
|
||
|
|
return {m, static_cast<Ring>(z - m)};
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring zext(Ring share, unsigned /*from_bits*/, unsigned /*to_bits*/) noexcept
|
||
|
|
{
|
||
|
|
return share;
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring sext_clear(Ring x, unsigned from_bits, unsigned to_bits) noexcept
|
||
|
|
{
|
||
|
|
if (from_bits == 0 || to_bits < from_bits)
|
||
|
|
return x;
|
||
|
|
const Ring sign = static_cast<Ring>((x >> (from_bits - 1u)) & 1u);
|
||
|
|
const Ring low = x & trunc::mask_bits<Ring>(from_bits);
|
||
|
|
if (!sign)
|
||
|
|
return low;
|
||
|
|
const Ring ext = trunc::mask_bits<Ring>(to_bits)
|
||
|
|
^ trunc::mask_bits<Ring>(from_bits);
|
||
|
|
return static_cast<Ring>(low | ext);
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Sign-extend using a secret MSB bit share (public when p1 holds 0).
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring sext_party(Ring x_share, std::uint8_t msb_pub, unsigned from_bits,
|
||
|
|
unsigned to_bits, unsigned party)
|
||
|
|
{
|
||
|
|
const Ring low = x_share & trunc::mask_bits<Ring>(from_bits);
|
||
|
|
if (!(msb_pub & 1u))
|
||
|
|
return low;
|
||
|
|
const Ring ext = trunc::mask_bits<Ring>(to_bits)
|
||
|
|
^ trunc::mask_bits<Ring>(from_bits);
|
||
|
|
// Party 0 absorbs the public extension bits.
|
||
|
|
if (party == 0)
|
||
|
|
return static_cast<Ring>(low | ext);
|
||
|
|
return low;
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> sext_party_pair(Ring x0, Ring x1, unsigned from_bits,
|
||
|
|
unsigned to_bits)
|
||
|
|
{
|
||
|
|
if (from_bits == 0 || to_bits < from_bits
|
||
|
|
|| from_bits >= 8u * sizeof(Ring))
|
||
|
|
return {x0, x1};
|
||
|
|
auto prep = sample_msb_prep<Ring>(from_bits);
|
||
|
|
auto [b0, b1] = msb_party_pair(x0, x1, prep, from_bits);
|
||
|
|
auto [s0, s1] = detail::b2a_bit<Ring>(b0, b1);
|
||
|
|
const Ring ext = static_cast<Ring>(trunc::mask_bits<Ring>(to_bits)
|
||
|
|
^ trunc::mask_bits<Ring>(from_bits));
|
||
|
|
auto shifted = trunc::trunc_exact_pair(x0, x1,
|
||
|
|
trunc::make_trunc_exact_prep<Ring>(
|
||
|
|
static_cast<unsigned>(8u * sizeof(Ring)), from_bits));
|
||
|
|
const Ring low0 = static_cast<Ring>(x0 - (shifted.first << from_bits));
|
||
|
|
const Ring low1 = static_cast<Ring>(x1 - (shifted.second << from_bits));
|
||
|
|
auto fill = detail::bit_mul_shares(s0, s1, ext, Ring{});
|
||
|
|
return {static_cast<Ring>(low0 + fill.first),
|
||
|
|
static_cast<Ring>(low1 + fill.second)};
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
bool range_ok_clear(Ring x, unsigned ell) noexcept
|
||
|
|
{
|
||
|
|
if (ell >= 8u * sizeof(Ring))
|
||
|
|
return true;
|
||
|
|
return (x & ~trunc::mask_bits<Ring>(ell)) == Ring{};
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring recip_newton_clear(Ring y, Ring y0_guess, unsigned frac_bits)
|
||
|
|
{
|
||
|
|
const Ring two = static_cast<Ring>(Ring{2} << frac_bits);
|
||
|
|
const Ring yx = static_cast<Ring>((y * y0_guess) >> frac_bits);
|
||
|
|
const Ring t = static_cast<Ring>(two - yx);
|
||
|
|
return static_cast<Ring>((y0_guess * t) >> frac_bits);
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring div_clear(Ring num, Ring den)
|
||
|
|
{
|
||
|
|
if (den == Ring{})
|
||
|
|
throw std::invalid_argument("div by zero");
|
||
|
|
return static_cast<Ring>(num / den);
|
||
|
|
}
|
||
|
|
|
||
|
|
/// @brief Goldschmidt/Newton division on shares via `mul_exact_trunc`.
|
||
|
|
/// @details Flagged leak: `floor(log2(den))` is public. There is no drop-in
|
||
|
|
/// replacement that keeps the magnitude secret and still normalizes
|
||
|
|
/// the Newton seed. See `dpf/revealing.hpp`.
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> div_party_pair(Ring num0, Ring num1, Ring den0, Ring den1,
|
||
|
|
unsigned frac_bits)
|
||
|
|
{
|
||
|
|
if (frac_bits == 0 || frac_bits >= 8u * sizeof(Ring) - 1u)
|
||
|
|
throw std::invalid_argument("div frac_bits");
|
||
|
|
auto prep = sample_msb_prep<Ring>(64);
|
||
|
|
auto nz = gt_party_pair(den0, den1, Ring{}, Ring{}, prep, 64);
|
||
|
|
if ((nz.first ^ nz.second) == 0)
|
||
|
|
throw std::invalid_argument("div by zero");
|
||
|
|
unsigned log2 = 0;
|
||
|
|
for (int bit = 5; bit >= 0; --bit)
|
||
|
|
{
|
||
|
|
const unsigned cand = log2 | (1u << static_cast<unsigned>(bit));
|
||
|
|
if (cand >= 63u)
|
||
|
|
continue;
|
||
|
|
const Ring thr = static_cast<Ring>((Ring{1} << cand) - Ring{1});
|
||
|
|
auto g = gt_party_pair(den0, den1, thr, Ring{}, prep, 64);
|
||
|
|
if ((g.first ^ g.second) != 0)
|
||
|
|
log2 = cand;
|
||
|
|
}
|
||
|
|
Ring d0 = den0;
|
||
|
|
Ring d1 = den1;
|
||
|
|
if (log2 < frac_bits)
|
||
|
|
{
|
||
|
|
const unsigned sh = frac_bits - log2;
|
||
|
|
d0 = static_cast<Ring>(d0 << sh);
|
||
|
|
d1 = static_cast<Ring>(d1 << sh);
|
||
|
|
}
|
||
|
|
else if (log2 > frac_bits)
|
||
|
|
{
|
||
|
|
auto tr = trunc::trunc_exact_pair(d0, d1,
|
||
|
|
trunc::make_trunc_exact_prep<Ring>(
|
||
|
|
static_cast<unsigned>(8u * sizeof(Ring)), log2 - frac_bits));
|
||
|
|
d0 = tr.first;
|
||
|
|
d1 = tr.second;
|
||
|
|
}
|
||
|
|
Ring g0 = static_cast<Ring>(Ring{1} << frac_bits);
|
||
|
|
Ring g1{};
|
||
|
|
const Ring two = static_cast<Ring>(Ring{2} << frac_bits);
|
||
|
|
for (int it = 0; it < 6; ++it)
|
||
|
|
{
|
||
|
|
auto prod = trunc::mul_exact_trunc(d0, d1, g0, g1, frac_bits);
|
||
|
|
const Ring t0 = static_cast<Ring>(two - prod.z0);
|
||
|
|
const Ring t1 = static_cast<Ring>(Ring{} - prod.z1);
|
||
|
|
auto ng = trunc::mul_exact_trunc(g0, g1, t0, t1, frac_bits);
|
||
|
|
g0 = ng.z0;
|
||
|
|
g1 = ng.z1;
|
||
|
|
}
|
||
|
|
Ring inv0 = g0;
|
||
|
|
Ring inv1 = g1;
|
||
|
|
if (log2 != 0)
|
||
|
|
{
|
||
|
|
auto tr = trunc::trunc_exact_pair(g0, g1,
|
||
|
|
trunc::make_trunc_exact_prep<Ring>(
|
||
|
|
static_cast<unsigned>(8u * sizeof(Ring)), log2));
|
||
|
|
inv0 = tr.first;
|
||
|
|
inv1 = tr.second;
|
||
|
|
}
|
||
|
|
auto q = trunc::mul_exact_trunc(num0, num1, inv0, inv1, frac_bits);
|
||
|
|
return {q.z0, q.z1};
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
std::pair<Ring, Ring> share_input(Ring clear, unsigned owner)
|
||
|
|
{
|
||
|
|
if (owner == 0)
|
||
|
|
return {clear, Ring{}};
|
||
|
|
if (owner == 1)
|
||
|
|
return {Ring{}, clear};
|
||
|
|
throw std::invalid_argument("share_input owner");
|
||
|
|
}
|
||
|
|
|
||
|
|
template <typename Ring>
|
||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
||
|
|
Ring declassify(Ring s0, Ring s1) noexcept
|
||
|
|
{
|
||
|
|
return static_cast<Ring>(s0 + s1);
|
||
|
|
}
|
||
|
|
|
||
|
|
} // namespace share_cmp
|
||
|
|
} // namespace dpf
|
||
|
|
|
||
|
|
#endif // LIBDPF_INCLUDE_DPF_SHARE_CMP_HPP__
|