libdpf/include/dpf/trunc.hpp

229 lines
7.2 KiB
C++
Raw Permalink Normal View History

/// @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 <cstddef>
#include <cstdint>
#include <stdexcept>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "dpf/edabit.hpp"
#include "dpf/random.hpp"
namespace dpf
{
namespace trunc
{
template <typename Ring>
HEDLEY_NO_THROW
constexpr Ring mask_bits(unsigned bits) noexcept
{
if (bits == 0)
return Ring{};
if (bits >= 8u * sizeof(Ring))
return static_cast<Ring>(~Ring{});
return static_cast<Ring>((Ring{1} << bits) - Ring{1});
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring trunc_prob(Ring share, unsigned s) noexcept
{
return static_cast<Ring>(share >> s);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring trunc_prob_clear(Ring x0, Ring x1, unsigned s) noexcept
{
return trunc_prob(static_cast<Ring>(x0 + x1), s);
}
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
Ring trunc_msb(Ring share, unsigned s, Ring /*msb_share*/) noexcept
{
return static_cast<Ring>(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 <typename Ring>
struct trunc_exact_prep
{
edabit::edabit_share<Ring> r0;
edabit::edabit_share<Ring> r1;
unsigned n = 0;
unsigned s = 0;
Ring clear_r{};
};
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
trunc_exact_prep<Ring> 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<Ring> p;
p.n = n;
p.s = s;
auto pair = edabit::sample_edabit_pair<Ring>(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 <typename Ring>
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<Ring>(out + static_cast<Ring>(delta >> s));
return out;
}
/// @brief Two-party exact trunc. Opens `x - r` only; the wrap carry is a GMW AND.
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
std::pair<Ring, Ring> trunc_exact_pair(Ring x0, Ring x1,
const trunc_exact_prep<Ring> & prep)
{
const Ring delta = static_cast<Ring>(
(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<std::uint8_t>(
(static_cast<std::uint64_t>(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<std::uint8_t>(rd0 ^ dc0 ^ rc.first);
c1 = static_cast<std::uint8_t>(rd1 ^ dc1 ^ rc.second);
}
auto dab = ot::sample_dabit_pair<Ring>();
const std::uint8_t mask = static_cast<std::uint8_t>(
(c0 ^ dab.p0.bit) ^ (c1 ^ dab.p1.bit));
const Ring w0 = edabit::b2a_party_bit<Ring>(c0, dab.p0, mask, 0, 0);
const Ring w1 = edabit::b2a_party_bit<Ring>(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 <typename Ring>
struct mul_trunc_shares
{
Ring z0{};
Ring z1{};
};
template <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
mul_trunc_shares<Ring> mul_trunc_clear(Ring x0, Ring x1, Ring y0, Ring y1,
unsigned s)
{
const Ring prod = static_cast<Ring>((x0 + x1) * (y0 + y1));
const Ring z = static_cast<Ring>(prod >> s);
const Ring mask = dpf::uniform_sample<Ring>();
return mul_trunc_shares<Ring>{mask, static_cast<Ring>(z - mask)};
}
template <typename Ring>
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<Ring>(z + d_open * b);
z = static_cast<Ring>(z + e_open * a);
if (party == 0)
z = static_cast<Ring>(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 <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
mul_trunc_shares<Ring> mul_trunc_pair(Ring x0, Ring x1, Ring y0, Ring y1,
unsigned s)
{
const Ring a0 = dpf::uniform_sample<Ring>();
const Ring b0 = dpf::uniform_sample<Ring>();
const Ring a1 = dpf::uniform_sample<Ring>();
const Ring b1 = dpf::uniform_sample<Ring>();
const Ring a = static_cast<Ring>(a0 + a1);
const Ring b = static_cast<Ring>(b0 + b1);
const Ring c = static_cast<Ring>(a * b);
const Ring c0 = dpf::uniform_sample<Ring>();
const Ring c1 = static_cast<Ring>(c - c0);
const Ring d = static_cast<Ring>((x0 + x1) - a);
const Ring e = static_cast<Ring>((y0 + y1) - b);
return mul_trunc_shares<Ring>{
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 <typename Ring>
HEDLEY_WARN_UNUSED_RESULT
mul_trunc_shares<Ring> 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<unsigned>(8u * sizeof(Ring));
if (s >= n)
throw std::invalid_argument("mul_exact_trunc shift");
auto prep = make_trunc_exact_prep<Ring>(n, s);
auto [t0, t1] = trunc_exact_pair(prod.z0, prod.z1, prep);
return mul_trunc_shares<Ring>{t0, t1};
}
} // namespace trunc
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_TRUNC_HPP__