libdpf/include/dpf/iknp.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

700 lines
24 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @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__