701 lines
24 KiB
C++
701 lines
24 KiB
C++
|
|
/// @file dpf/iknp.hpp
|
|||
|
|
/// @brief Semi-honest IKNP OT extension and the pads a DPF dealer would sample.
|
|||
|
|
/// @details Base OTs are Chou–Orlandi (LATINCRYPT 2015, ePrint 2015/267) on
|
|||
|
|
/// P-256. Extension follows Ishai, Kilian, Nissim, and Petrank,
|
|||
|
|
/// CRYPTO 2003, with fixed-key AES as the correlation-robust hash.
|
|||
|
|
/// `sample` returns this party's shares of random bit triples,
|
|||
|
|
/// bit×block triples, comparison B2A pads, and Doerner–Shelat
|
|||
|
|
/// correction-word pads.
|
|||
|
|
///
|
|||
|
|
/// Ideal functionality of `sample` (semi-honest, two parties):
|
|||
|
|
/// - **Inputs.** Both parties pass the same lengths `(nblock, nbit, nb2a, ncw)`
|
|||
|
|
/// and call on the peer link before any other walk traffic.
|
|||
|
|
/// - **Outputs.** Party `i` receives shares such that
|
|||
|
|
/// - blocks: `(a0⊕a1)·(b0⊕b1) = c0⊕c1` (bit × 128-bit block);
|
|||
|
|
/// - bits: `(a0⊕a1)∧(b0⊕b1) = c0⊕c1`;
|
|||
|
|
/// - b2a: `(add0+add1) mod 2^64 = r0⊕r1` (0 or 1);
|
|||
|
|
/// - cw: `gamma0⊕gamma1 = (bit1·rand0)⊕(bit0·rand1)` (XOR shares — neither
|
|||
|
|
/// party learns the peer pad bit; see `ds_sample_cw`).
|
|||
|
|
/// - **Hidden.** Peer seeds, peer choice bits used only as OT choice, and the
|
|||
|
|
/// peer's pad bit (so an opened blind `path⊕pad` does not open the path).
|
|||
|
|
///
|
|||
|
|
/// Cost of `sample` with tape `T = nblock + nbit + ncw` and security
|
|||
|
|
/// parameter `κ = 128`: two Chou–Orlandi base sessions of `κ` OTs (P-256
|
|||
|
|
/// points, `Θ(κ)` scalar muls), then two OT-extension directions each
|
|||
|
|
/// sending `κ ⌈T/8⌉` bytes of U and `T · 16` bytes of correction, plus
|
|||
|
|
/// `ncw · 16` bytes of `gamma` and an optional B2A extension of length
|
|||
|
|
/// `nb2a`. See [tour_iknp](@ref tour_iknp) for the comparison with a p2
|
|||
|
|
/// dealer tape and with Half-Tree §5.2.
|
|||
|
|
|
|||
|
|
#ifndef LIBDPF_INCLUDE_DPF_IKNP_HPP__
|
|||
|
|
#define LIBDPF_INCLUDE_DPF_IKNP_HPP__
|
|||
|
|
|
|||
|
|
#include <cstddef>
|
|||
|
|
#include <cstdint>
|
|||
|
|
#include <cstring>
|
|||
|
|
#include <stdexcept>
|
|||
|
|
#include <utility>
|
|||
|
|
#include <vector>
|
|||
|
|
|
|||
|
|
#include "hedley/hedley.h"
|
|||
|
|
#include "simde/simde/x86/avx2.h"
|
|||
|
|
|
|||
|
|
#include "dpf/net/channel.hpp"
|
|||
|
|
#include "dpf/p256.hpp"
|
|||
|
|
#include "dpf/prg_aes.hpp"
|
|||
|
|
#include "dpf/random.hpp"
|
|||
|
|
|
|||
|
|
namespace dpf
|
|||
|
|
{
|
|||
|
|
namespace iknp
|
|||
|
|
{
|
|||
|
|
|
|||
|
|
struct block_share
|
|||
|
|
{
|
|||
|
|
std::uint8_t a = 0;
|
|||
|
|
simde__m128i b{};
|
|||
|
|
simde__m128i c{};
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct bit_share
|
|||
|
|
{
|
|||
|
|
std::uint8_t a = 0;
|
|||
|
|
std::uint8_t b = 0;
|
|||
|
|
std::uint8_t c = 0;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct b2a_share
|
|||
|
|
{
|
|||
|
|
std::uint8_t r = 0;
|
|||
|
|
std::uint64_t add = 0;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct cw_share
|
|||
|
|
{
|
|||
|
|
simde__m128i rand{};
|
|||
|
|
simde__m128i gamma{};
|
|||
|
|
std::uint8_t bit = 0;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct material
|
|||
|
|
{
|
|||
|
|
std::vector<block_share> blocks;
|
|||
|
|
std::vector<bit_share> bits;
|
|||
|
|
std::vector<b2a_share> b2a;
|
|||
|
|
std::vector<cw_share> cws;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
namespace detail
|
|||
|
|
{
|
|||
|
|
|
|||
|
|
inline constexpr std::size_t kappa = 128;
|
|||
|
|
inline constexpr std::size_t chunk_rows = 8192;
|
|||
|
|
|
|||
|
|
struct point_msg
|
|||
|
|
{
|
|||
|
|
std::uint8_t enc[33];
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
struct role_state
|
|||
|
|
{
|
|||
|
|
bool ready = false;
|
|||
|
|
simde__m128i k0[kappa]{};
|
|||
|
|
simde__m128i k1[kappa]{};
|
|||
|
|
simde__m128i seed[kappa]{};
|
|||
|
|
std::uint8_t delta_bits[kappa]{};
|
|||
|
|
simde__m128i delta{};
|
|||
|
|
std::uint64_t rows = 0;
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
inline simde__m128i xor_block(simde__m128i a, simde__m128i b)
|
|||
|
|
{
|
|||
|
|
return simde_mm_xor_si128(a, b);
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline std::uint8_t lsb(simde__m128i x)
|
|||
|
|
{
|
|||
|
|
unsigned char b = 0;
|
|||
|
|
std::memcpy(&b, &x, 1);
|
|||
|
|
return static_cast<std::uint8_t>(b & 1u);
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline simde__m128i bit_block(std::uint8_t bit)
|
|||
|
|
{
|
|||
|
|
simde__m128i z = simde_mm_setzero_si128();
|
|||
|
|
const unsigned char b = static_cast<unsigned char>(bit & 1u);
|
|||
|
|
std::memcpy(&z, &b, 1);
|
|||
|
|
return z;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline simde__m128i gate(std::uint8_t bit, simde__m128i block)
|
|||
|
|
{
|
|||
|
|
return (bit & 1u) ? block : simde_mm_setzero_si128();
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline std::uint64_t low64(simde__m128i x)
|
|||
|
|
{
|
|||
|
|
std::uint64_t v = 0;
|
|||
|
|
std::memcpy(&v, &x, sizeof(v));
|
|||
|
|
return v;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline simde__m128i ot_hash(std::uint64_t index, simde__m128i row)
|
|||
|
|
{
|
|||
|
|
const prg::purpose_scope counted(prg::purpose::hash);
|
|||
|
|
const auto mixed = xor_block(row,
|
|||
|
|
simde_mm_set_epi64x(static_cast<std::int64_t>(index >> 32),
|
|||
|
|
static_cast<std::int64_t>(index)));
|
|||
|
|
return prg::aes128::eval(mixed,
|
|||
|
|
static_cast<psnip_uint32_t>(index * 0x9E3779B9u) ^ 0xA5A5u);
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void sample_scalar(std::uint64_t k[4])
|
|||
|
|
{
|
|||
|
|
for (;;)
|
|||
|
|
{
|
|||
|
|
for (int i = 0; i < 4; ++i)
|
|||
|
|
k[i] = dpf::uniform_sample<std::uint64_t>();
|
|||
|
|
const bool zero = (k[0] | k[1] | k[2] | k[3]) == 0;
|
|||
|
|
if (!zero && p256_detail::limbs_cmp(k, p256_detail::N, 4) < 0)
|
|||
|
|
return;
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline simde__m128i hash_point(const p256_detail::affine & p)
|
|||
|
|
{
|
|||
|
|
std::uint8_t enc[33]{};
|
|||
|
|
p256_detail::encode_point(enc, p);
|
|||
|
|
simde__m128i b0 = simde_mm_setzero_si128();
|
|||
|
|
simde__m128i b1 = simde_mm_setzero_si128();
|
|||
|
|
std::memcpy(&b0, enc, 16);
|
|||
|
|
std::memcpy(&b1, enc + 16, 16);
|
|||
|
|
const simde__m128i b2 = simde_mm_set_epi64x(0, enc[32]);
|
|||
|
|
const prg::purpose_scope counted(prg::purpose::hash);
|
|||
|
|
auto h = prg::aes128::eval(b0, 1);
|
|||
|
|
h = xor_block(h, prg::aes128::eval(b1, 2));
|
|||
|
|
return xor_block(h, prg::aes128::eval(b2, 3));
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline simde__m128i pack_bits(const std::uint8_t * bits)
|
|||
|
|
{
|
|||
|
|
std::uint8_t packed[16]{};
|
|||
|
|
for (int i = 0; i < static_cast<int>(kappa); ++i)
|
|||
|
|
{
|
|||
|
|
if (bits[i] & 1u)
|
|||
|
|
packed[static_cast<unsigned>(i) >> 3] |=
|
|||
|
|
static_cast<std::uint8_t>(1u << (i & 7));
|
|||
|
|
}
|
|||
|
|
simde__m128i out = simde_mm_setzero_si128();
|
|||
|
|
std::memcpy(&out, packed, 16);
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void expand_column(simde__m128i seed, std::uint64_t domain,
|
|||
|
|
std::uint8_t * dst, std::size_t nbytes)
|
|||
|
|
{
|
|||
|
|
const auto tweaked = xor_block(seed,
|
|||
|
|
simde_mm_set_epi64x(static_cast<std::int64_t>(domain), 0));
|
|||
|
|
const std::size_t nblocks = (nbytes + 15) / 16;
|
|||
|
|
std::vector<simde__m128i> buf(nblocks);
|
|||
|
|
if (nblocks > 0)
|
|||
|
|
prg::aes128::eval(tweaked, buf.data(),
|
|||
|
|
static_cast<psnip_uint32_t>(nblocks), 0);
|
|||
|
|
std::memcpy(dst, buf.data(), nbytes);
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief Transpose a `kappa × nrows` bit matrix packed by columns into rows.
|
|||
|
|
inline void transpose_rows(const std::uint8_t * cols, std::size_t nbytes,
|
|||
|
|
std::size_t nrows, simde__m128i * rows)
|
|||
|
|
{
|
|||
|
|
// Process 8 row-bits at a time when possible via byte gathers; fall back
|
|||
|
|
// per-row for the tail. Still O(kappa · nrows) but with tight inner loops.
|
|||
|
|
for (std::size_t j = 0; j < nrows; ++j)
|
|||
|
|
{
|
|||
|
|
std::uint8_t packed[16]{};
|
|||
|
|
const std::size_t byte = j >> 3;
|
|||
|
|
const auto mask = static_cast<std::uint8_t>(1u << (j & 7));
|
|||
|
|
for (int i = 0; i < static_cast<int>(kappa); ++i)
|
|||
|
|
{
|
|||
|
|
if (cols[static_cast<std::size_t>(i) * nbytes + byte] & mask)
|
|||
|
|
packed[static_cast<unsigned>(i) >> 3] |=
|
|||
|
|
static_cast<std::uint8_t>(1u << (i & 7));
|
|||
|
|
}
|
|||
|
|
std::memcpy(&rows[j], packed, 16);
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void base_sender(net::channel & ch, simde__m128i k0[kappa],
|
|||
|
|
simde__m128i k1[kappa])
|
|||
|
|
{
|
|||
|
|
std::uint64_t a[4];
|
|||
|
|
sample_scalar(a);
|
|||
|
|
const auto A = p256_detail::point_scalarmul_limbs(
|
|||
|
|
p256_detail::generator_point(), a);
|
|||
|
|
point_msg am{};
|
|||
|
|
p256_detail::encode_point(am.enc, A);
|
|||
|
|
ch.send(net::msg::bytes, am);
|
|||
|
|
const auto bs = ch.recv_vec<point_msg>(net::msg::bytes);
|
|||
|
|
if (bs.size() != kappa)
|
|||
|
|
throw std::runtime_error("iknp base OT count");
|
|||
|
|
for (std::size_t i = 0; i < kappa; ++i)
|
|||
|
|
{
|
|||
|
|
const auto B = p256_detail::decode_strict(bs[i].enc, 33);
|
|||
|
|
k0[i] = hash_point(p256_detail::point_scalarmul_limbs(B, a));
|
|||
|
|
k1[i] = hash_point(p256_detail::point_scalarmul_limbs(
|
|||
|
|
p256_detail::point_sub(B, A), a));
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void base_receiver(net::channel & ch, simde__m128i seed[kappa],
|
|||
|
|
std::uint8_t delta_bits[kappa], simde__m128i & delta)
|
|||
|
|
{
|
|||
|
|
const auto am = ch.recv<point_msg>(net::msg::bytes);
|
|||
|
|
const auto A = p256_detail::decode_strict(am.enc, 33);
|
|||
|
|
std::vector<point_msg> bs(kappa);
|
|||
|
|
for (std::size_t i = 0; i < kappa; ++i)
|
|||
|
|
{
|
|||
|
|
delta_bits[i] = static_cast<std::uint8_t>(
|
|||
|
|
dpf::uniform_sample<unsigned char>() & 1u);
|
|||
|
|
std::uint64_t r[4];
|
|||
|
|
sample_scalar(r);
|
|||
|
|
const auto R = p256_detail::point_scalarmul_limbs(
|
|||
|
|
p256_detail::generator_point(), r);
|
|||
|
|
const auto B = delta_bits[i]
|
|||
|
|
? p256_detail::point_add(A, R) : R;
|
|||
|
|
p256_detail::encode_point(bs[i].enc, B);
|
|||
|
|
seed[i] = hash_point(p256_detail::point_scalarmul_limbs(A, r));
|
|||
|
|
}
|
|||
|
|
delta = pack_bits(delta_bits);
|
|||
|
|
ch.send_vec(bs, net::msg::bytes);
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void ensure_sender(net::channel & ch, role_state & st)
|
|||
|
|
{
|
|||
|
|
if (st.ready)
|
|||
|
|
return;
|
|||
|
|
base_receiver(ch, st.seed, st.delta_bits, st.delta);
|
|||
|
|
st.ready = true;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void ensure_receiver(net::channel & ch, role_state & st)
|
|||
|
|
{
|
|||
|
|
if (st.ready)
|
|||
|
|
return;
|
|||
|
|
base_sender(ch, st.k0, st.k1);
|
|||
|
|
st.ready = true;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void extend_send(net::channel & ch, role_state & st, std::size_t n,
|
|||
|
|
std::vector<simde__m128i> & m0, std::vector<simde__m128i> & m1)
|
|||
|
|
{
|
|||
|
|
if (n == 0)
|
|||
|
|
return;
|
|||
|
|
ensure_sender(ch, st);
|
|||
|
|
m0.resize(n);
|
|||
|
|
m1.resize(n);
|
|||
|
|
std::size_t off = 0;
|
|||
|
|
while (off < n)
|
|||
|
|
{
|
|||
|
|
const std::size_t rows = std::min(chunk_rows, n - off);
|
|||
|
|
const std::size_t nbytes = (rows + 7) / 8;
|
|||
|
|
const auto u = ch.recv_bytes(net::msg::bytes);
|
|||
|
|
if (u.size() != kappa * nbytes)
|
|||
|
|
throw std::runtime_error("iknp extension length");
|
|||
|
|
std::vector<std::uint8_t> cols(kappa * nbytes);
|
|||
|
|
for (std::size_t i = 0; i < kappa; ++i)
|
|||
|
|
{
|
|||
|
|
expand_column(st.seed[i], st.rows, cols.data() + i * nbytes, nbytes);
|
|||
|
|
if (st.delta_bits[i])
|
|||
|
|
{
|
|||
|
|
for (std::size_t b = 0; b < nbytes; ++b)
|
|||
|
|
cols[i * nbytes + b] = static_cast<std::uint8_t>(
|
|||
|
|
cols[i * nbytes + b] ^ u[i * nbytes + b]);
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
std::vector<simde__m128i> q(rows);
|
|||
|
|
transpose_rows(cols.data(), nbytes, rows, q.data());
|
|||
|
|
for (std::size_t j = 0; j < rows; ++j)
|
|||
|
|
{
|
|||
|
|
const auto index = st.rows + j;
|
|||
|
|
m0[off + j] = ot_hash(index, q[j]);
|
|||
|
|
m1[off + j] = ot_hash(index, xor_block(q[j], st.delta));
|
|||
|
|
}
|
|||
|
|
st.rows += rows;
|
|||
|
|
off += rows;
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
inline void extend_recv(net::channel & ch, role_state & st,
|
|||
|
|
const std::uint8_t * choices, std::size_t n,
|
|||
|
|
std::vector<simde__m128i> & masks)
|
|||
|
|
{
|
|||
|
|
if (n == 0)
|
|||
|
|
return;
|
|||
|
|
ensure_receiver(ch, st);
|
|||
|
|
masks.resize(n);
|
|||
|
|
std::size_t off = 0;
|
|||
|
|
while (off < n)
|
|||
|
|
{
|
|||
|
|
const std::size_t rows = std::min(chunk_rows, n - off);
|
|||
|
|
const std::size_t nbytes = (rows + 7) / 8;
|
|||
|
|
std::vector<std::uint8_t> xbytes(nbytes);
|
|||
|
|
for (std::size_t j = 0; j < rows; ++j)
|
|||
|
|
{
|
|||
|
|
if (choices[off + j] & 1u)
|
|||
|
|
xbytes[j >> 3] = static_cast<std::uint8_t>(
|
|||
|
|
xbytes[j >> 3] | (1u << (j & 7)));
|
|||
|
|
}
|
|||
|
|
std::vector<std::uint8_t> u(kappa * nbytes);
|
|||
|
|
std::vector<std::uint8_t> t0cols(kappa * nbytes);
|
|||
|
|
for (std::size_t i = 0; i < kappa; ++i)
|
|||
|
|
{
|
|||
|
|
std::vector<std::uint8_t> t1(nbytes);
|
|||
|
|
expand_column(st.k0[i], st.rows, t0cols.data() + i * nbytes, nbytes);
|
|||
|
|
expand_column(st.k1[i], st.rows, t1.data(), nbytes);
|
|||
|
|
for (std::size_t b = 0; b < nbytes; ++b)
|
|||
|
|
u[i * nbytes + b] = static_cast<std::uint8_t>(
|
|||
|
|
t0cols[i * nbytes + b] ^ t1[b] ^ xbytes[b]);
|
|||
|
|
}
|
|||
|
|
ch.send_bytes(net::msg::bytes, u.data(), u.size());
|
|||
|
|
std::vector<simde__m128i> t(rows);
|
|||
|
|
transpose_rows(t0cols.data(), nbytes, rows, t.data());
|
|||
|
|
for (std::size_t j = 0; j < rows; ++j)
|
|||
|
|
masks[off + j] = ot_hash(st.rows + j, t[j]);
|
|||
|
|
st.rows += rows;
|
|||
|
|
off += rows;
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief Correlated OT correction: sender holds (m0,m1), payload Δ; receiver
|
|||
|
|
/// with choice χ gets m_χ ⊕ (χ·Δ). Parties get additive XOR shares of
|
|||
|
|
/// χ·Δ by taking sender share = m0 and receiver share = got ⊕ m0 path.
|
|||
|
|
/// @details Sender transmits `d = m0⊕m1⊕payload`. Receiver returns
|
|||
|
|
/// `choice ? mask⊕d : mask`. Sender's share is `m0`; receiver's is
|
|||
|
|
/// the returned value. Then `sender⊕receiver = choice·payload` when
|
|||
|
|
/// the OT is correct (`mask = m_choice`).
|
|||
|
|
inline void correct_ot(net::channel & ch, int me, bool i_am_sender,
|
|||
|
|
const std::vector<simde__m128i> & m0,
|
|||
|
|
const std::vector<simde__m128i> & m1,
|
|||
|
|
const std::vector<simde__m128i> & payloads,
|
|||
|
|
const std::vector<simde__m128i> & masks,
|
|||
|
|
const std::vector<std::uint8_t> & choices,
|
|||
|
|
std::vector<simde__m128i> & share)
|
|||
|
|
{
|
|||
|
|
const std::size_t n = i_am_sender ? m0.size() : masks.size();
|
|||
|
|
share.resize(n);
|
|||
|
|
if (n == 0)
|
|||
|
|
return;
|
|||
|
|
if (i_am_sender)
|
|||
|
|
{
|
|||
|
|
std::vector<simde__m128i> d(n);
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
d[i] = xor_block(xor_block(m0[i], m1[i]), payloads[i]);
|
|||
|
|
// Fixed order: party 0 always sends first when it is the sender;
|
|||
|
|
// when party 1 is the sender it sends (party 0 receives).
|
|||
|
|
if (me == 0)
|
|||
|
|
ch.send_vec(d, net::msg::bytes);
|
|||
|
|
else
|
|||
|
|
ch.send_vec(d, net::msg::bytes);
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
share[i] = m0[i];
|
|||
|
|
return;
|
|||
|
|
}
|
|||
|
|
auto d = ch.recv_vec<simde__m128i>(net::msg::bytes);
|
|||
|
|
if (d.size() != n)
|
|||
|
|
throw std::runtime_error("iknp correction length");
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
share[i] = (choices[i] & 1u) ? xor_block(masks[i], d[i]) : masks[i];
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
} // namespace detail
|
|||
|
|
|
|||
|
|
/// @brief Sample this party's dealer pads. `me` is 0 or 1. Both parties pass
|
|||
|
|
/// the same lengths and call this before any other traffic on `link`.
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
inline material sample(net::channel & link, int me, std::size_t nblock,
|
|||
|
|
std::size_t nbit, std::size_t nb2a, std::size_t ncw)
|
|||
|
|
{
|
|||
|
|
if (me != 0 && me != 1)
|
|||
|
|
throw std::invalid_argument("iknp party");
|
|||
|
|
|
|||
|
|
auto rand_bit = [] {
|
|||
|
|
return static_cast<std::uint8_t>(dpf::uniform_sample<unsigned char>() & 1u);
|
|||
|
|
};
|
|||
|
|
|
|||
|
|
std::vector<std::uint8_t> block_a(nblock), bit_a(nbit), bit_b(nbit), cw_bit(ncw), b2a_r(nb2a);
|
|||
|
|
std::vector<simde__m128i> block_b(nblock), cw_rand(ncw);
|
|||
|
|
for (std::size_t i = 0; i < nblock; ++i)
|
|||
|
|
{
|
|||
|
|
block_a[i] = rand_bit();
|
|||
|
|
block_b[i] = dpf::uniform_sample<simde__m128i>();
|
|||
|
|
}
|
|||
|
|
for (std::size_t i = 0; i < nbit; ++i)
|
|||
|
|
{
|
|||
|
|
bit_a[i] = rand_bit();
|
|||
|
|
bit_b[i] = rand_bit();
|
|||
|
|
}
|
|||
|
|
for (std::size_t i = 0; i < ncw; ++i)
|
|||
|
|
{
|
|||
|
|
cw_bit[i] = rand_bit();
|
|||
|
|
cw_rand[i] = dpf::uniform_sample<simde__m128i>();
|
|||
|
|
}
|
|||
|
|
for (std::size_t i = 0; i < nb2a; ++i)
|
|||
|
|
b2a_r[i] = rand_bit();
|
|||
|
|
|
|||
|
|
// Cross terms via two IKNP directions. Sender holds payload Δ, receiver
|
|||
|
|
// chooses χ; shares sum to χ·Δ.
|
|||
|
|
// D01 (P0 sends): χ=a1 / bit_a1 / cw_bit1, Δ=B0 / bit_b0 / cw_rand0.
|
|||
|
|
// D10 (P1 sends): χ=a0 / bit_a0 / cw_bit0, Δ=B1 / bit_b1 / cw_rand1.
|
|||
|
|
const std::size_t nxor = nblock + nbit + ncw;
|
|||
|
|
|
|||
|
|
std::vector<std::uint8_t> choices(nxor);
|
|||
|
|
std::vector<simde__m128i> payloads(nxor);
|
|||
|
|
for (std::size_t i = 0; i < nblock; ++i)
|
|||
|
|
{
|
|||
|
|
choices[i] = block_a[i];
|
|||
|
|
payloads[i] = block_b[i];
|
|||
|
|
}
|
|||
|
|
for (std::size_t j = 0; j < nbit; ++j)
|
|||
|
|
{
|
|||
|
|
choices[nblock + j] = bit_a[j];
|
|||
|
|
payloads[nblock + j] = detail::bit_block(bit_b[j]);
|
|||
|
|
}
|
|||
|
|
for (std::size_t k = 0; k < ncw; ++k)
|
|||
|
|
{
|
|||
|
|
choices[nblock + nbit + k] = cw_bit[k];
|
|||
|
|
payloads[nblock + nbit + k] = cw_rand[k];
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
detail::role_state send_state; // used when we are IKNP sender (produce m0/m1)
|
|||
|
|
detail::role_state recv_state; // used when we are IKNP receiver
|
|||
|
|
std::vector<simde__m128i> send_m0, send_m1, recv_masks;
|
|||
|
|
|
|||
|
|
// D01: party 0 sends, party 1 receives.
|
|||
|
|
if (me == 0)
|
|||
|
|
detail::extend_send(link, send_state, nxor, send_m0, send_m1);
|
|||
|
|
else
|
|||
|
|
detail::extend_recv(link, recv_state, choices.data(), nxor, recv_masks);
|
|||
|
|
|
|||
|
|
std::vector<simde__m128i> d01_share;
|
|||
|
|
detail::correct_ot(link, me, me == 0,
|
|||
|
|
me == 0 ? send_m0 : std::vector<simde__m128i>{},
|
|||
|
|
me == 0 ? send_m1 : std::vector<simde__m128i>{},
|
|||
|
|
me == 0 ? payloads : std::vector<simde__m128i>{},
|
|||
|
|
me == 0 ? std::vector<simde__m128i>{} : recv_masks,
|
|||
|
|
me == 0 ? std::vector<std::uint8_t>{} : choices,
|
|||
|
|
d01_share);
|
|||
|
|
// d01_share_0 ⊕ d01_share_1 = χ1 · Δ0 = a1·B0 (etc.)
|
|||
|
|
|
|||
|
|
// D10: party 1 sends, party 0 receives.
|
|||
|
|
if (me == 0)
|
|||
|
|
detail::extend_recv(link, recv_state, choices.data(), nxor, recv_masks);
|
|||
|
|
else
|
|||
|
|
detail::extend_send(link, send_state, nxor, send_m0, send_m1);
|
|||
|
|
|
|||
|
|
std::vector<simde__m128i> d10_share;
|
|||
|
|
detail::correct_ot(link, me, me == 1,
|
|||
|
|
me == 1 ? send_m0 : std::vector<simde__m128i>{},
|
|||
|
|
me == 1 ? send_m1 : std::vector<simde__m128i>{},
|
|||
|
|
me == 1 ? payloads : std::vector<simde__m128i>{},
|
|||
|
|
me == 1 ? std::vector<simde__m128i>{} : recv_masks,
|
|||
|
|
me == 1 ? std::vector<std::uint8_t>{} : choices,
|
|||
|
|
d10_share);
|
|||
|
|
// d10_share_0 ⊕ d10_share_1 = χ0 · Δ1 = a0·B1 (etc.)
|
|||
|
|
|
|||
|
|
// Free large OT pads.
|
|||
|
|
send_m0.clear();
|
|||
|
|
send_m0.shrink_to_fit();
|
|||
|
|
send_m1.clear();
|
|||
|
|
send_m1.shrink_to_fit();
|
|||
|
|
recv_masks.clear();
|
|||
|
|
recv_masks.shrink_to_fit();
|
|||
|
|
|
|||
|
|
// CW gamma shares: gamma0 ⊕ gamma1 = (bit1·rand0) ⊕ (bit0·rand1)
|
|||
|
|
// = (D01 cw) ⊕ (D10 cw). Party 0 samples gamma0 and masks with both of
|
|||
|
|
// its OT shares; party 1 unmasks with both of its shares.
|
|||
|
|
std::vector<simde__m128i> gamma(ncw);
|
|||
|
|
if (ncw != 0)
|
|||
|
|
{
|
|||
|
|
const std::size_t cw_off = nblock + nbit;
|
|||
|
|
if (me == 0)
|
|||
|
|
{
|
|||
|
|
std::vector<simde__m128i> msg(ncw);
|
|||
|
|
for (std::size_t k = 0; k < ncw; ++k)
|
|||
|
|
{
|
|||
|
|
gamma[k] = dpf::uniform_sample<simde__m128i>();
|
|||
|
|
msg[k] = detail::xor_block(gamma[k],
|
|||
|
|
detail::xor_block(d01_share[cw_off + k], d10_share[cw_off + k]));
|
|||
|
|
}
|
|||
|
|
link.send_vec(msg, net::msg::bytes);
|
|||
|
|
}
|
|||
|
|
else
|
|||
|
|
{
|
|||
|
|
auto msg = link.recv_vec<simde__m128i>(net::msg::bytes);
|
|||
|
|
if (msg.size() != ncw)
|
|||
|
|
throw std::runtime_error("iknp cw pad length");
|
|||
|
|
for (std::size_t k = 0; k < ncw; ++k)
|
|||
|
|
gamma[k] = detail::xor_block(msg[k],
|
|||
|
|
detail::xor_block(d01_share[cw_off + k], d10_share[cw_off + k]));
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// B2A: arithmetic shares of the XOR of the two random bits.
|
|||
|
|
std::vector<std::uint64_t> rho0(nb2a), arith_w(nb2a);
|
|||
|
|
if (nb2a != 0)
|
|||
|
|
{
|
|||
|
|
std::vector<simde__m128i> bm0, bm1, bmasks;
|
|||
|
|
if (me == 0)
|
|||
|
|
{
|
|||
|
|
detail::extend_send(link, send_state, nb2a, bm0, bm1);
|
|||
|
|
std::vector<std::uint64_t> corr(nb2a);
|
|||
|
|
for (std::size_t i = 0; i < nb2a; ++i)
|
|||
|
|
{
|
|||
|
|
rho0[i] = detail::low64(bm0[i]);
|
|||
|
|
const auto rho1 = detail::low64(bm1[i]);
|
|||
|
|
corr[i] = rho0[i] - rho1 + b2a_r[i];
|
|||
|
|
}
|
|||
|
|
link.send_vec(corr, net::msg::bytes);
|
|||
|
|
}
|
|||
|
|
else
|
|||
|
|
{
|
|||
|
|
detail::extend_recv(link, recv_state, b2a_r.data(), nb2a, bmasks);
|
|||
|
|
auto corr = link.recv_vec<std::uint64_t>(net::msg::bytes);
|
|||
|
|
if (corr.size() != nb2a)
|
|||
|
|
throw std::runtime_error("iknp b2a length");
|
|||
|
|
for (std::size_t i = 0; i < nb2a; ++i)
|
|||
|
|
{
|
|||
|
|
const auto rho = detail::low64(bmasks[i]);
|
|||
|
|
arith_w[i] = rho + static_cast<std::uint64_t>(b2a_r[i]) * corr[i];
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
material out;
|
|||
|
|
out.blocks.resize(nblock);
|
|||
|
|
out.bits.resize(nbit);
|
|||
|
|
out.b2a.resize(nb2a);
|
|||
|
|
out.cws.resize(ncw);
|
|||
|
|
for (std::size_t i = 0; i < nblock; ++i)
|
|||
|
|
{
|
|||
|
|
// c0 ⊕ c1 = a0·B0 ⊕ a1·B1 ⊕ a1·B0 ⊕ a0·B1 = (a0⊕a1)·(B0⊕B1)
|
|||
|
|
const auto local = detail::gate(block_a[i], block_b[i]);
|
|||
|
|
// Party 0's share of a1·B0 is d01; of a0·B1 is d10. Same XOR for both
|
|||
|
|
// parties (their XOR shares already sum to the cross terms).
|
|||
|
|
const auto c = detail::xor_block(local,
|
|||
|
|
detail::xor_block(d01_share[i], d10_share[i]));
|
|||
|
|
out.blocks[i] = block_share{block_a[i], block_b[i], c};
|
|||
|
|
}
|
|||
|
|
for (std::size_t j = 0; j < nbit; ++j)
|
|||
|
|
{
|
|||
|
|
const auto off = nblock + j;
|
|||
|
|
const auto local = static_cast<std::uint8_t>(bit_a[j] & bit_b[j]);
|
|||
|
|
const auto c = static_cast<std::uint8_t>(local
|
|||
|
|
^ detail::lsb(d01_share[off]) ^ detail::lsb(d10_share[off]));
|
|||
|
|
out.bits[j] = bit_share{bit_a[j], bit_b[j], c};
|
|||
|
|
}
|
|||
|
|
for (std::size_t k = 0; k < ncw; ++k)
|
|||
|
|
out.cws[k] = cw_share{cw_rand[k], gamma[k], cw_bit[k]};
|
|||
|
|
for (std::size_t i = 0; i < nb2a; ++i)
|
|||
|
|
{
|
|||
|
|
std::uint64_t add;
|
|||
|
|
if (me == 0)
|
|||
|
|
add = static_cast<std::uint64_t>(b2a_r[i]) + (rho0[i] << 1);
|
|||
|
|
else
|
|||
|
|
add = static_cast<std::uint64_t>(b2a_r[i]) - (arith_w[i] << 1);
|
|||
|
|
out.b2a[i] = b2a_share{b2a_r[i], add};
|
|||
|
|
}
|
|||
|
|
return out;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief 1-out-of-2 OT of 128-bit strings. The sender holds `(m0, m1)`. The
|
|||
|
|
/// receiver holds a choice bit per row and receives `m_choice`.
|
|||
|
|
/// @details Chou–Orlandi base OT runs on the first call for this `st`. Later
|
|||
|
|
/// calls only extend. `st` is one direction: the sender's state is
|
|||
|
|
/// not the receiver's state. `n == 0` sends nothing. Both parties
|
|||
|
|
/// pass the same `n`.
|
|||
|
|
/// @param ch peer channel
|
|||
|
|
/// @param me 0 or 1
|
|||
|
|
/// @param i_am_sender this party holds `m0` and `m1`
|
|||
|
|
/// @param st extension state for this direction
|
|||
|
|
/// @param m0 sender's first message, length `n` (ignored by the receiver)
|
|||
|
|
/// @param m1 sender's second message, length `n` (ignored by the receiver)
|
|||
|
|
/// @param choices receiver's choice bits, length `n` (ignored by the sender)
|
|||
|
|
/// @param got receiver's output `m_choice` (cleared for the sender)
|
|||
|
|
inline void transfer_labels(net::channel & ch, int me, bool i_am_sender,
|
|||
|
|
detail::role_state & st, const std::vector<simde__m128i> & m0,
|
|||
|
|
const std::vector<simde__m128i> & m1,
|
|||
|
|
const std::vector<std::uint8_t> & choices, std::vector<simde__m128i> & got)
|
|||
|
|
{
|
|||
|
|
if (me != 0 && me != 1)
|
|||
|
|
throw std::invalid_argument("iknp party");
|
|||
|
|
const std::size_t n = i_am_sender ? m0.size() : choices.size();
|
|||
|
|
if (i_am_sender)
|
|||
|
|
{
|
|||
|
|
if (m1.size() != n)
|
|||
|
|
throw std::invalid_argument("iknp transfer length");
|
|||
|
|
}
|
|||
|
|
else if (choices.size() != n)
|
|||
|
|
throw std::invalid_argument("iknp transfer length");
|
|||
|
|
got.clear();
|
|||
|
|
if (n == 0)
|
|||
|
|
return;
|
|||
|
|
if (i_am_sender)
|
|||
|
|
{
|
|||
|
|
std::vector<simde__m128i> r0, r1;
|
|||
|
|
detail::extend_send(ch, st, n, r0, r1);
|
|||
|
|
std::vector<simde__m128i> corr(2u * n);
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
{
|
|||
|
|
corr[2u * i] = detail::xor_block(m0[i], r0[i]);
|
|||
|
|
corr[2u * i + 1u] = detail::xor_block(m1[i], r1[i]);
|
|||
|
|
}
|
|||
|
|
ch.send_vec(corr, net::msg::bytes);
|
|||
|
|
return;
|
|||
|
|
}
|
|||
|
|
std::vector<simde__m128i> masks;
|
|||
|
|
detail::extend_recv(ch, st, choices.data(), n, masks);
|
|||
|
|
const auto corr = ch.recv_vec<simde__m128i>(net::msg::bytes);
|
|||
|
|
if (corr.size() != 2u * n)
|
|||
|
|
throw std::runtime_error("iknp transfer correction length");
|
|||
|
|
got.resize(n);
|
|||
|
|
for (std::size_t i = 0; i < n; ++i)
|
|||
|
|
{
|
|||
|
|
const std::size_t slot = 2u * i + static_cast<std::size_t>(choices[i] & 1u);
|
|||
|
|
got[i] = detail::xor_block(masks[i], corr[slot]);
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
/// @brief Channel rounds inside `sample` (base OT, extension chunks, corrections).
|
|||
|
|
/// @details Two directions when `nblock+nbit+ncw > 0`: each is a 2-message
|
|||
|
|
/// Chou–Orlandi base, one U-matrix per `chunk_rows`, and one
|
|||
|
|
/// correction. CW gamma is one more round. B2A reuses the base OT
|
|||
|
|
/// and adds an extension plus a correction. Add this to
|
|||
|
|
/// `plan::rounds()` via `rounds_including`.
|
|||
|
|
HEDLEY_WARN_UNUSED_RESULT
|
|||
|
|
HEDLEY_PURE
|
|||
|
|
inline std::size_t setup_rounds(std::size_t nblock, std::size_t nbit,
|
|||
|
|
std::size_t nb2a, std::size_t ncw) noexcept
|
|||
|
|
{
|
|||
|
|
const auto chunks = [](std::size_t n) {
|
|||
|
|
return n == 0 ? std::size_t{0}
|
|||
|
|
: (n + detail::chunk_rows - 1) / detail::chunk_rows;
|
|||
|
|
};
|
|||
|
|
const std::size_t nxor = nblock + nbit + ncw;
|
|||
|
|
std::size_t rounds = 0;
|
|||
|
|
if (nxor != 0)
|
|||
|
|
{
|
|||
|
|
rounds += 2 + chunks(nxor) + 1;
|
|||
|
|
rounds += 2 + chunks(nxor) + 1;
|
|||
|
|
}
|
|||
|
|
if (ncw != 0)
|
|||
|
|
++rounds;
|
|||
|
|
if (nb2a != 0)
|
|||
|
|
rounds += chunks(nb2a) + 1;
|
|||
|
|
return rounds;
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
} // namespace iknp
|
|||
|
|
} // namespace dpf
|
|||
|
|
|
|||
|
|
#endif // LIBDPF_INCLUDE_DPF_IKNP_HPP__
|