/// @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 #include #include #include #include #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 HEDLEY_NO_THROW constexpr std::uint8_t msb_clear(Ring x, unsigned n) noexcept { if (n == 0) return 0; return static_cast((x >> (n - 1u)) & 1u); } template HEDLEY_NO_THROW constexpr std::uint8_t gt_clear(Ring x, Ring y, unsigned n) noexcept { const Ring mask = trunc::mask_bits(n); return static_cast(((x & mask) > (y & mask)) ? 1 : 0); } template HEDLEY_NO_THROW constexpr std::uint8_t eq_clear(Ring x, Ring y) noexcept { return static_cast(x == y ? 1 : 0); } template struct msb_prep { Ring r0{}; Ring r1{}; }; template HEDLEY_WARN_UNUSED_RESULT msb_prep sample_msb_prep(unsigned /*n*/) { msb_prep p; const Ring r = dpf::uniform_sample(); p.r0 = dpf::uniform_sample(); p.r1 = static_cast(r - p.r0); return p; } /// @brief This party's XOR share of MSB after A2B (the bit at `n-1`). template HEDLEY_WARN_UNUSED_RESULT std::uint8_t msb_party(const std::vector & my_bits, unsigned n) { if (n == 0) return 0; return edabit::detail::get_bit(my_bits, n - 1u); } template HEDLEY_WARN_UNUSED_RESULT std::pair msb_party_pair(Ring x0, Ring x1, const msb_prep & prep, unsigned n) { (void)prep; auto eda = edabit::sample_edabit_pair(n); auto [b0, b1] = edabit::a2b_gmw_pair(eda, x0, x1); return {msb_party(b0, n), msb_party(b1, n)}; } /// @brief This party's XOR share of an unsigned greater-than bit. template HEDLEY_WARN_UNUSED_RESULT std::uint8_t gt_party(std::uint8_t my_bit) noexcept { return static_cast(my_bit & 1u); } /// @brief Full-width unsigned compare. Bit shares stay split; their XOR is `x > y`. template HEDLEY_WARN_UNUSED_RESULT std::pair gt_party_pair(Ring x0, Ring x1, Ring y0, Ring y1, const msb_prep & prep, unsigned n) { (void)prep; auto ex = edabit::sample_edabit_pair(n); auto ey = edabit::sample_edabit_pair(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(g.first), gt_party(g.second)}; } namespace detail { template HEDLEY_WARN_UNUSED_RESULT std::pair b2a_bit(std::uint8_t b0, std::uint8_t b1) { auto dab = ot::sample_dabit_pair(); const std::uint8_t mask = static_cast( (b0 ^ dab.p0.bit) ^ (b1 ^ dab.p1.bit)); return {edabit::b2a_party_bit(b0, dab.p0, mask, 0, 0), edabit::b2a_party_bit(b1, dab.p1, mask, 1, 0)}; } template HEDLEY_WARN_UNUSED_RESULT std::pair bit_mul_shares(Ring b0, Ring b1, Ring x0, Ring x1) { beavers::session 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 HEDLEY_WARN_UNUSED_RESULT Ring relu_clear(Ring x0, Ring x1, unsigned n) { const Ring x = static_cast(x0 + x1); return msb_clear(x, n) ? Ring{} : x; } /// @brief ReLU party share: `(1 - msb) · x` via `bit_mul` on additive bit shares. template HEDLEY_WARN_UNUSED_RESULT Ring relu_party(Ring x_share, Ring keep_arith) { return static_cast(keep_arith * x_share); } template HEDLEY_WARN_UNUSED_RESULT std::pair relu_party_pair(Ring x0, Ring x1, unsigned n) { auto prep = sample_msb_prep(n); auto [b0, b1] = msb_party_pair(x0, x1, prep, n); auto [s0, s1] = detail::b2a_bit(b0, b1); const Ring k0 = static_cast(Ring{1} - s0); const Ring k1 = static_cast(Ring{} - s1); return detail::bit_mul_shares(k0, k1, x0, x1); } template HEDLEY_WARN_UNUSED_RESULT Ring max_clear(Ring x0, Ring x1, Ring y0, Ring y1, unsigned n) { const Ring x = static_cast(x0 + x1); const Ring y = static_cast(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 HEDLEY_WARN_UNUSED_RESULT Ring max_party(Ring y_share, Ring select_share) { return static_cast(y_share + select_share); } template HEDLEY_WARN_UNUSED_RESULT std::pair max_party_pair(Ring x0, Ring x1, Ring y0, Ring y1, unsigned n) { auto prep = sample_msb_prep(n); auto [g0, g1] = gt_party_pair(x0, x1, y0, y1, prep, n); auto [s0, s1] = detail::b2a_bit(g0, g1); const Ring d0 = static_cast(x0 - y0); const Ring d1 = static_cast(x1 - y1); auto sel = detail::bit_mul_shares(s0, s1, d0, d1); return {max_party(y0, sel.first), max_party(y1, sel.second)}; } template HEDLEY_WARN_UNUSED_RESULT std::pair mux_clear(std::uint8_t sel, Ring a0, Ring a1, Ring b0, Ring b1) { const Ring a = static_cast(a0 + a1); const Ring b = static_cast(b0 + b1); const Ring z = sel ? a : b; const Ring m = dpf::uniform_sample(); return {m, static_cast(z - m)}; } template HEDLEY_WARN_UNUSED_RESULT Ring zext(Ring share, unsigned /*from_bits*/, unsigned /*to_bits*/) noexcept { return share; } template 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((x >> (from_bits - 1u)) & 1u); const Ring low = x & trunc::mask_bits(from_bits); if (!sign) return low; const Ring ext = trunc::mask_bits(to_bits) ^ trunc::mask_bits(from_bits); return static_cast(low | ext); } /// @brief Sign-extend using a secret MSB bit share (public when p1 holds 0). template 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(from_bits); if (!(msb_pub & 1u)) return low; const Ring ext = trunc::mask_bits(to_bits) ^ trunc::mask_bits(from_bits); // Party 0 absorbs the public extension bits. if (party == 0) return static_cast(low | ext); return low; } template HEDLEY_WARN_UNUSED_RESULT std::pair 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(from_bits); auto [b0, b1] = msb_party_pair(x0, x1, prep, from_bits); auto [s0, s1] = detail::b2a_bit(b0, b1); const Ring ext = static_cast(trunc::mask_bits(to_bits) ^ trunc::mask_bits(from_bits)); auto shifted = trunc::trunc_exact_pair(x0, x1, trunc::make_trunc_exact_prep( static_cast(8u * sizeof(Ring)), from_bits)); const Ring low0 = static_cast(x0 - (shifted.first << from_bits)); const Ring low1 = static_cast(x1 - (shifted.second << from_bits)); auto fill = detail::bit_mul_shares(s0, s1, ext, Ring{}); return {static_cast(low0 + fill.first), static_cast(low1 + fill.second)}; } template 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(ell)) == Ring{}; } template HEDLEY_WARN_UNUSED_RESULT Ring recip_newton_clear(Ring y, Ring y0_guess, unsigned frac_bits) { const Ring two = static_cast(Ring{2} << frac_bits); const Ring yx = static_cast((y * y0_guess) >> frac_bits); const Ring t = static_cast(two - yx); return static_cast((y0_guess * t) >> frac_bits); } template HEDLEY_WARN_UNUSED_RESULT Ring div_clear(Ring num, Ring den) { if (den == Ring{}) throw std::invalid_argument("div by zero"); return static_cast(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 HEDLEY_WARN_UNUSED_RESULT std::pair 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(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(bit)); if (cand >= 63u) continue; const Ring thr = static_cast((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(d0 << sh); d1 = static_cast(d1 << sh); } else if (log2 > frac_bits) { auto tr = trunc::trunc_exact_pair(d0, d1, trunc::make_trunc_exact_prep( static_cast(8u * sizeof(Ring)), log2 - frac_bits)); d0 = tr.first; d1 = tr.second; } Ring g0 = static_cast(Ring{1} << frac_bits); Ring g1{}; const Ring two = static_cast(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(two - prod.z0); const Ring t1 = static_cast(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( static_cast(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 HEDLEY_WARN_UNUSED_RESULT std::pair 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 HEDLEY_WARN_UNUSED_RESULT Ring declassify(Ring s0, Ring s1) noexcept { return static_cast(s0 + s1); } } // namespace share_cmp } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_SHARE_CMP_HPP__