/// @file dpf/trunc.hpp /// @brief Truncate-and-reduce on additive shares (probabilistic and exact). #ifndef LIBDPF_INCLUDE_DPF_TRUNC_HPP__ #define LIBDPF_INCLUDE_DPF_TRUNC_HPP__ #include #include #include #include #include #include "hedley/hedley.h" #include "dpf/edabit.hpp" #include "dpf/random.hpp" namespace dpf { namespace trunc { template HEDLEY_NO_THROW constexpr Ring mask_bits(unsigned bits) noexcept { if (bits == 0) return Ring{}; if (bits >= 8u * sizeof(Ring)) return static_cast(~Ring{}); return static_cast((Ring{1} << bits) - Ring{1}); } template HEDLEY_WARN_UNUSED_RESULT Ring trunc_prob(Ring share, unsigned s) noexcept { return static_cast(share >> s); } template HEDLEY_WARN_UNUSED_RESULT Ring trunc_prob_clear(Ring x0, Ring x1, unsigned s) noexcept { return trunc_prob(static_cast(x0 + x1), s); } template HEDLEY_WARN_UNUSED_RESULT Ring trunc_msb(Ring share, unsigned s, Ring /*msb_share*/) noexcept { return static_cast(share >> s); } inline std::uint64_t trunc_exact_clear(std::uint64_t x0, std::uint64_t x1, unsigned n, unsigned s) { if (s >= n || n > 64) throw std::invalid_argument("trunc_exact_clear"); const std::uint64_t low_m = (s >= 64) ? ~std::uint64_t{0} : ((std::uint64_t{1} << s) - 1u); const std::uint64_t high_m = ((n - s) >= 64) ? ~std::uint64_t{0} : ((std::uint64_t{1} << (n - s)) - 1u); const std::uint64_t v0 = x0 & low_m; const std::uint64_t v1 = x1 & low_m; const std::uint64_t u0 = (x0 >> s) & high_m; const std::uint64_t u1 = (x1 >> s) & high_m; const std::uint64_t cin = (v0 + v1) >> s; return (u0 + u1 + cin) & high_m; } template struct trunc_exact_prep { edabit::edabit_share r0; edabit::edabit_share r1; unsigned n = 0; unsigned s = 0; Ring clear_r{}; }; template HEDLEY_WARN_UNUSED_RESULT trunc_exact_prep make_trunc_exact_prep(unsigned n, unsigned s) { if (s >= n || n > 8u * sizeof(Ring)) throw std::invalid_argument("trunc_exact width"); trunc_exact_prep p; p.n = n; p.s = s; auto pair = edabit::sample_edabit_pair(s); p.r0 = pair.p0; p.r1 = pair.p1; p.clear_r = pair.clear_r; return p; } /// @brief Exact trunc party share after opening `delta = x - r`. /// @details `trunc(x) = (delta >> s) + wrap`. `wrap_share` is an additive share /// of the GMW carry. `r` is not opened. template HEDLEY_WARN_UNUSED_RESULT Ring trunc_exact_party(Ring x_share, Ring r_low_arith, Ring delta, Ring wrap_share, unsigned s, unsigned party, Ring r_public = Ring{}) { (void)x_share; (void)r_low_arith; (void)r_public; (void)s; Ring out = wrap_share; if (party == 0) out = static_cast(out + static_cast(delta >> s)); return out; } /// @brief Two-party exact trunc. Opens `x - r` only; the wrap carry is a GMW AND. template HEDLEY_WARN_UNUSED_RESULT std::pair trunc_exact_pair(Ring x0, Ring x1, const trunc_exact_prep & prep) { const Ring delta = static_cast( (x0 - prep.r0.arith) + (x1 - prep.r1.arith)); std::uint8_t c0 = 0; std::uint8_t c1 = 0; for (unsigned i = 0; i < prep.s; ++i) { const std::uint8_t r0 = edabit::detail::get_bit(prep.r0.bits_packed, i); const std::uint8_t r1 = edabit::detail::get_bit(prep.r1.bits_packed, i); const std::uint8_t di = static_cast( (static_cast(delta) >> i) & 1u); auto rc = edabit::and_pair(r0, r1, c0, 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; c0 = static_cast(rd0 ^ dc0 ^ rc.first); c1 = static_cast(rd1 ^ dc1 ^ rc.second); } auto dab = ot::sample_dabit_pair(); const std::uint8_t mask = static_cast( (c0 ^ dab.p0.bit) ^ (c1 ^ dab.p1.bit)); const Ring w0 = edabit::b2a_party_bit(c0, dab.p0, mask, 0, 0); const Ring w1 = edabit::b2a_party_bit(c1, dab.p1, mask, 1, 0); return {trunc_exact_party(x0, prep.r0.arith, delta, w0, prep.s, 0), trunc_exact_party(x1, prep.r1.arith, delta, w1, prep.s, 1)}; } template struct mul_trunc_shares { Ring z0{}; Ring z1{}; }; template HEDLEY_WARN_UNUSED_RESULT mul_trunc_shares mul_trunc_clear(Ring x0, Ring x1, Ring y0, Ring y1, unsigned s) { const Ring prod = static_cast((x0 + x1) * (y0 + y1)); const Ring z = static_cast(prod >> s); const Ring mask = dpf::uniform_sample(); return mul_trunc_shares{mask, static_cast(z - mask)}; } template HEDLEY_WARN_UNUSED_RESULT Ring mul_trunc_party(Ring x, Ring y, Ring a, Ring b, Ring c, Ring d_open, Ring e_open, unsigned s, unsigned party) { (void)x; (void)y; Ring z = c; z = static_cast(z + d_open * b); z = static_cast(z + e_open * a); if (party == 0) z = static_cast(z + d_open * e_open); return trunc_prob(z, s); } /// @brief Two-party mul_trunc via one Beaver product then local shift. /// @details Probabilistic trunc: each party shifts its product share. The /// clear product may need a wider intermediate when `s > 0` (fixed /// point); for `uint64` the ring multiply wraps — callers that need /// exact fixed-point should keep operands below `2^{64-s}`. template HEDLEY_WARN_UNUSED_RESULT mul_trunc_shares mul_trunc_pair(Ring x0, Ring x1, Ring y0, Ring y1, unsigned s) { const Ring a0 = dpf::uniform_sample(); const Ring b0 = dpf::uniform_sample(); const Ring a1 = dpf::uniform_sample(); const Ring b1 = dpf::uniform_sample(); const Ring a = static_cast(a0 + a1); const Ring b = static_cast(b0 + b1); const Ring c = static_cast(a * b); const Ring c0 = dpf::uniform_sample(); const Ring c1 = static_cast(c - c0); const Ring d = static_cast((x0 + x1) - a); const Ring e = static_cast((y0 + y1) - b); return mul_trunc_shares{ mul_trunc_party(x0, y0, a0, b0, c0, d, e, s, 0), mul_trunc_party(x1, y1, a1, b1, c1, d, e, s, 1)}; } /// @brief Beaver product, then exact trunc by `s` (edaBit wrap), no clear multiply. template HEDLEY_WARN_UNUSED_RESULT mul_trunc_shares mul_exact_trunc(Ring x0, Ring x1, Ring y0, Ring y1, unsigned s) { auto prod = mul_trunc_pair(x0, x1, y0, y1, 0); if (s == 0) return prod; const unsigned n = static_cast(8u * sizeof(Ring)); if (s >= n) throw std::invalid_argument("mul_exact_trunc shift"); auto prep = make_trunc_exact_prep(n, s); auto [t0, t1] = trunc_exact_pair(prod.z0, prod.z1, prep); return mul_trunc_shares{t0, t1}; } } // namespace trunc } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_TRUNC_HPP__