229 lines
7.2 KiB
C++
229 lines
7.2 KiB
C++
|
|
/// @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__
|