/// @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 #include #include #include #include #include #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 blocks; std::vector bits; std::vector b2a; std::vector 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(b & 1u); } inline simde__m128i bit_block(std::uint8_t bit) { simde__m128i z = simde_mm_setzero_si128(); const unsigned char b = static_cast(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(index >> 32), static_cast(index))); return prg::aes128::eval(mixed, static_cast(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(); 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(kappa); ++i) { if (bits[i] & 1u) packed[static_cast(i) >> 3] |= static_cast(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(domain), 0)); const std::size_t nblocks = (nbytes + 15) / 16; std::vector buf(nblocks); if (nblocks > 0) prg::aes128::eval(tweaked, buf.data(), static_cast(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(1u << (j & 7)); for (int i = 0; i < static_cast(kappa); ++i) { if (cols[static_cast(i) * nbytes + byte] & mask) packed[static_cast(i) >> 3] |= static_cast(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(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(net::msg::bytes); const auto A = p256_detail::decode_strict(am.enc, 33); std::vector bs(kappa); for (std::size_t i = 0; i < kappa; ++i) { delta_bits[i] = static_cast( dpf::uniform_sample() & 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 & m0, std::vector & 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 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( cols[i * nbytes + b] ^ u[i * nbytes + b]); } } std::vector 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 & 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 xbytes(nbytes); for (std::size_t j = 0; j < rows; ++j) { if (choices[off + j] & 1u) xbytes[j >> 3] = static_cast( xbytes[j >> 3] | (1u << (j & 7))); } std::vector u(kappa * nbytes); std::vector t0cols(kappa * nbytes); for (std::size_t i = 0; i < kappa; ++i) { std::vector 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( t0cols[i * nbytes + b] ^ t1[b] ^ xbytes[b]); } ch.send_bytes(net::msg::bytes, u.data(), u.size()); std::vector 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 & m0, const std::vector & m1, const std::vector & payloads, const std::vector & masks, const std::vector & choices, std::vector & 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 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(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(dpf::uniform_sample() & 1u); }; std::vector block_a(nblock), bit_a(nbit), bit_b(nbit), cw_bit(ncw), b2a_r(nb2a); std::vector 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(); } 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(); } 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 choices(nxor); std::vector 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 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 d01_share; detail::correct_ot(link, me, me == 0, me == 0 ? send_m0 : std::vector{}, me == 0 ? send_m1 : std::vector{}, me == 0 ? payloads : std::vector{}, me == 0 ? std::vector{} : recv_masks, me == 0 ? std::vector{} : 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 d10_share; detail::correct_ot(link, me, me == 1, me == 1 ? send_m0 : std::vector{}, me == 1 ? send_m1 : std::vector{}, me == 1 ? payloads : std::vector{}, me == 1 ? std::vector{} : recv_masks, me == 1 ? std::vector{} : 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 gamma(ncw); if (ncw != 0) { const std::size_t cw_off = nblock + nbit; if (me == 0) { std::vector msg(ncw); for (std::size_t k = 0; k < ncw; ++k) { gamma[k] = dpf::uniform_sample(); 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(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 rho0(nb2a), arith_w(nb2a); if (nb2a != 0) { std::vector bm0, bm1, bmasks; if (me == 0) { detail::extend_send(link, send_state, nb2a, bm0, bm1); std::vector 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(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(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(bit_a[j] & bit_b[j]); const auto c = static_cast(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(b2a_r[i]) + (rho0[i] << 1); else add = static_cast(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 & m0, const std::vector & m1, const std::vector & choices, std::vector & 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 r0, r1; detail::extend_send(ch, st, n, r0, r1); std::vector 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 masks; detail::extend_recv(ch, st, choices.data(), n, masks); const auto corr = ch.recv_vec(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(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__