/// @file dpf/yao_stack.hpp /// @brief Stacked and one-hot garbling of half-gate netlists. /// @details A conditional sends one string of AND rows, as long as the heavier /// branch, not the sum. Each branch is garbled from a seed that is the /// hash of the control label for that semantic bit. The generator XORs /// the two materials. The evaluator holds one control label, so she /// can regenerate the active seed herself. A two-row ciphertext under /// that label delivers the inactive seed. She rebuilds the inactive /// material, XORs it out of the stack, and evaluates the branch she /// actually holds. Output labels are translated into one common /// free-XOR encoding, four rows per output bit, so the garbler's share /// does not depend on which branch ran. /// /// A k-way switch is the same stack over k seeds. The selector is a /// short bit string. The demux row for the selector color carries the /// inactive seeds. The mux has two rows per selector color. /// @note Stacked garbling: David Heath and Vladimir Kolesnikov, "Stacked /// Garbling: Garbled Circuit Proportional to Longest Execution Path," /// CRYPTO 2020 (ePrint 2020/973). One-hot garbling: David Heath and /// Vladimir Kolesnikov, CCS 2021. Gate rows remain half-gates (Zahur, /// Rosulek, and Evans, EUROCRYPT 2015). /// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors) /// @license Released under a GNU General Public v2.0 (GPLv2) license; /// see [LICENSE.md](@ref license) for details. #ifndef LIBDPF_INCLUDE_DPF_YAO_STACK_HPP__ #define LIBDPF_INCLUDE_DPF_YAO_STACK_HPP__ #include #include #include #include #include #include #include "hedley/hedley.h" #include "dpf/prg_aes.hpp" #include "dpf/random.hpp" #include "dpf/yao.hpp" namespace dpf { namespace yao { /// @brief What a stacked conditional or a one-hot switch sent, and the shares. struct stack_result { std::vector share0; std::vector share1; /// @brief AND-row blocks on the stack (two per AND of the heaviest branch). std::size_t stack_blocks = 0; /// @brief AND-row blocks if every branch were sent on its own. std::size_t naive_blocks = 0; /// @brief Mux / demux blocks on top of the stack. std::size_t extra_blocks = 0; }; namespace stack_detail { using detail::and_eval; using detail::and_garble; using detail::lsb; using detail::mask_bit; using detail::require_bit; using detail::with_lsb; using detail::xor_b; using detail::zero_b; struct rng { block seed{}; std::uint32_t n = 0; block next() { return prg::aes128::eval(seed, n++); } }; struct seeded { block delta{}; std::vector z; std::vector tables; std::vector lg; std::vector le; }; inline block hash_at(block x, std::uint32_t pos) { const prg::purpose_scope counted(prg::purpose::hash); return prg::aes128::eval(x, pos); } inline seeded garble_seeded(const netlist & nl, block seed) { rng r; r.seed = seed; seeded g; g.delta = with_lsb(r.next(), 1); g.z.assign(nl.nwire(), zero_b()); g.lg.resize(nl.n_in()); g.le.resize(nl.n_in()); for (std::uint32_t i = 0; i < nl.n_in(); ++i) { const block lg = r.next(); g.lg[i] = lg; switch (nl.kind_at(i)) { case input::shared: { const block le = r.next(); g.le[i] = le; g.z[i] = xor_b(lg, le); break; } case input::priv0: g.z[i] = lg; break; case input::priv1: g.z[i] = lg; break; } } g.tables.reserve(static_cast(nl.n_and()) * 2u); std::uint64_t gid = 0; const auto & gates = nl.gates(); const std::uint32_t base = nl.n_in(); for (std::uint32_t gi = 0; gi < gates.size(); ++gi) { const auto & gate = gates[gi]; const std::uint32_t dst = base + gi; switch (gate.code) { case netlist::op::xor_: g.z[dst] = xor_b(g.z[gate.a], g.z[gate.b]); break; case netlist::op::not_: g.z[dst] = xor_b(g.z[gate.a], g.delta); break; case netlist::op::and_: { block row[2]; g.z[dst] = and_garble(g.z[gate.a], g.z[gate.b], g.delta, gid, row); g.tables.push_back(row[0]); g.tables.push_back(row[1]); ++gid; break; } } } return g; } inline std::vector eval_labels(const netlist & nl, const seeded & g, const block * tables, const std::uint8_t * p0, const std::uint8_t * p1) { detail::check_party_bits(nl, 0, p0); detail::check_party_bits(nl, 1, p1); std::vector w(nl.nwire(), zero_b()); for (std::uint32_t i = 0; i < nl.n_in(); ++i) { switch (nl.kind_at(i)) { case input::shared: { const block direct = xor_b(g.lg[i], mask_bit(p0[i], g.delta)); const block chosen = (p1[i] & 1u) ? xor_b(g.le[i], g.delta) : g.le[i]; w[i] = xor_b(direct, chosen); break; } case input::priv0: w[i] = xor_b(g.lg[i], mask_bit(p0[i], g.delta)); break; case input::priv1: w[i] = (p1[i] & 1u) ? xor_b(g.lg[i], g.delta) : g.lg[i]; break; } } std::uint64_t gid = 0; std::size_t ti = 0; const auto & gates = nl.gates(); const std::uint32_t base = nl.n_in(); for (std::uint32_t gi = 0; gi < gates.size(); ++gi) { const auto & gate = gates[gi]; const std::uint32_t dst = base + gi; switch (gate.code) { case netlist::op::xor_: w[dst] = xor_b(w[gate.a], w[gate.b]); break; case netlist::op::not_: w[dst] = w[gate.a]; break; case netlist::op::and_: w[dst] = and_eval(w[gate.a], w[gate.b], gid, tables + ti); ti += 2u; ++gid; break; } } std::vector outs; outs.reserve(nl.n_out()); for (std::uint32_t id : nl.outputs()) outs.push_back(w[id]); return outs; } inline std::vector pad_to(std::vector m, std::size_t n, block seed) { std::uint32_t i = 0; while (m.size() < n) m.push_back(prg::aes128::eval(seed, i++)); return m; } inline std::vector xor_vec(const std::vector & a, const std::vector & b) { if (a.size() != b.size()) throw std::logic_error("yao stack: length"); std::vector out(a.size()); for (std::size_t i = 0; i < a.size(); ++i) out[i] = xor_b(a[i], b[i]); return out; } inline block fold_labels(const std::vector & labels) { const prg::purpose_scope counted(prg::purpose::hash); block acc = zero_b(); for (std::size_t i = 0; i < labels.size(); ++i) acc = prg::aes128::eval(xor_b(acc, labels[i]), static_cast(i + 1)); return acc; } inline void require_same_outs(const netlist & a, const netlist & b) { if (a.n_out() != b.n_out() || a.n_out() == 0) throw std::invalid_argument("yao stack: outputs"); } } // namespace stack_detail /// @brief Secret if/else. Branch 0 runs when the shared control bit is 0. /// @details `then_nl` and `else_nl` may differ in size and in their inputs. /// `*_p0[i]` / `*_p1[i]` are that party's bit on input `i` of that /// branch, with the same layout as `yao::eval_pair`. The control bit /// is `control_p0 XOR control_p1`. HEDLEY_WARN_UNUSED_RESULT inline stack_result eval_if(const netlist & then_nl, const netlist & else_nl, std::uint8_t control_p0, std::uint8_t control_p1, const std::uint8_t * then_p0, const std::uint8_t * then_p1, const std::uint8_t * else_p0, const std::uint8_t * else_p1) { stack_detail::require_same_outs(then_nl, else_nl); stack_detail::require_bit(control_p0, "yao stack: control"); stack_detail::require_bit(control_p1, "yao stack: control"); const std::uint8_t sem = static_cast(control_p0 ^ control_p1); const block delta = stack_detail::with_lsb(dpf::uniform_sample(), 1); const block lg = dpf::uniform_sample(); const block le = dpf::uniform_sample(); const block z = stack_detail::xor_b(lg, le); const block direct = stack_detail::xor_b(lg, stack_detail::mask_bit(control_p0, delta)); const block chosen = (control_p1 & 1u) ? stack_detail::xor_b(le, delta) : le; const block active_ctl = stack_detail::xor_b(direct, chosen); const block l0 = z; const block l1 = stack_detail::xor_b(z, delta); const block seed0 = stack_detail::hash_at(l0, 11); const block seed1 = stack_detail::hash_at(l1, 11); auto g0 = stack_detail::garble_seeded(then_nl, seed0); auto g1 = stack_detail::garble_seeded(else_nl, seed1); const std::size_t n = std::max(g0.tables.size(), g1.tables.size()); const block pad_seed = stack_detail::hash_at( stack_detail::xor_b(seed0, seed1), 19); auto m0 = stack_detail::pad_to(g0.tables, n, pad_seed); auto m1 = stack_detail::pad_to(g1.tables, n, pad_seed); const auto stacked = stack_detail::xor_vec(m0, m1); block ct[2]; ct[stack_detail::lsb(l0)] = stack_detail::xor_b( stack_detail::hash_at(l0, 23), seed1); ct[stack_detail::lsb(l1)] = stack_detail::xor_b( stack_detail::hash_at(l1, 23), seed0); const block inactive_seed = stack_detail::xor_b( stack_detail::hash_at(active_ctl, 23), ct[stack_detail::lsb(active_ctl)]); const block active_seed = stack_detail::hash_at(active_ctl, 11); const block expect_inactive = (sem == 0) ? seed1 : seed0; const block expect_active = (sem == 0) ? seed0 : seed1; if (std::memcmp(&inactive_seed, &expect_inactive, sizeof(block)) != 0 || std::memcmp(&active_seed, &expect_active, sizeof(block)) != 0) throw std::logic_error("yao stack: control seed"); const auto & active_nl = (sem == 0) ? then_nl : else_nl; const auto & inactive_nl = (sem == 0) ? else_nl : then_nl; const stack_detail::seeded & active_g = (sem == 0) ? g0 : g1; auto rebuilt = stack_detail::garble_seeded(inactive_nl, inactive_seed); auto inactive_m = stack_detail::pad_to(std::move(rebuilt.tables), n, pad_seed); auto opened = stack_detail::xor_vec(stacked, inactive_m); if (active_g.tables.size() > opened.size()) throw std::logic_error("yao stack: short stack"); const std::uint8_t * ap0 = (sem == 0) ? then_p0 : else_p0; const std::uint8_t * ap1 = (sem == 0) ? then_p1 : else_p1; auto outs = stack_detail::eval_labels(active_nl, active_g, opened.data(), ap0, ap1); const block common_delta = stack_detail::with_lsb(dpf::uniform_sample(), 1); stack_result result; result.stack_blocks = n; result.naive_blocks = g0.tables.size() + g1.tables.size(); result.extra_blocks = 2u + static_cast(then_nl.n_out()) * 4u; result.share0.resize(then_nl.n_out()); result.share1.resize(then_nl.n_out()); const auto & outs0 = then_nl.outputs(); const auto & outs1 = else_nl.outputs(); for (std::uint32_t i = 0; i < then_nl.n_out(); ++i) { const block common_zero = dpf::uniform_sample(); block table[4]; for (unsigned cc = 0; cc < 2; ++cc) { for (unsigned oc = 0; oc < 2; ++oc) { const block ctl = (stack_detail::lsb(l0) == cc) ? l0 : l1; const std::uint8_t sem_c = static_cast( cc ^ stack_detail::lsb(l0)); const stack_detail::seeded & br = (sem_c == 0) ? g0 : g1; const std::uint32_t oid = (sem_c == 0) ? outs0[i] : outs1[i]; const std::uint8_t sem_o = static_cast( oc ^ stack_detail::lsb(br.z[oid])); const block ol = stack_detail::xor_b(br.z[oid], stack_detail::mask_bit(sem_o, br.delta)); const block common = stack_detail::xor_b(common_zero, stack_detail::mask_bit(sem_o, common_delta)); const unsigned idx = (cc << 1u) | oc; table[idx] = stack_detail::xor_b(common, stack_detail::hash_at( stack_detail::xor_b(ctl, ol), 40u + idx)); } } const unsigned idx = (static_cast(stack_detail::lsb(active_ctl)) << 1u) | stack_detail::lsb(outs[i]); const block got = stack_detail::xor_b(table[idx], stack_detail::hash_at( stack_detail::xor_b(active_ctl, outs[i]), 40u + idx)); result.share0[i] = stack_detail::lsb(common_zero); result.share1[i] = stack_detail::lsb(got); } return result; } /// @brief Secret choice of branch `index_p0 XOR index_p1`. /// @details `k` is `branches.size()`, from 2 to 8. The index uses the low /// `ceil(log2(k))` bits. Every branch has the same output count. /// `p0[b]` / `p1[b]` are that party's input bits for branch `b`. HEDLEY_WARN_UNUSED_RESULT inline stack_result eval_one_hot(const std::vector & branches, std::uint16_t index_p0, std::uint16_t index_p1, const std::vector> & p0, const std::vector> & p1) { const std::size_t k = branches.size(); if (k < 2 || k > 8) throw std::invalid_argument("yao stack: one-hot width"); if (p0.size() != k || p1.size() != k) throw std::invalid_argument("yao stack: one-hot inputs"); for (std::size_t i = 1; i < k; ++i) stack_detail::require_same_outs(branches[0], branches[i]); unsigned width = 1; while ((1u << width) < k) ++width; const std::uint16_t index = static_cast(index_p0 ^ index_p1); if (index >= k) throw std::invalid_argument("yao stack: index"); const block delta = stack_detail::with_lsb(dpf::uniform_sample(), 1); std::vector zbit(width); std::vector active_bit(width); for (unsigned b = 0; b < width; ++b) { const block lg = dpf::uniform_sample(); const block le = dpf::uniform_sample(); zbit[b] = stack_detail::xor_b(lg, le); const std::uint8_t bit0 = static_cast((index_p0 >> b) & 1u); const std::uint8_t bit1 = static_cast((index_p1 >> b) & 1u); const block direct = stack_detail::xor_b(lg, stack_detail::mask_bit(bit0, delta)); const block chosen = (bit1 != 0) ? stack_detail::xor_b(le, delta) : le; active_bit[b] = stack_detail::xor_b(direct, chosen); } auto pattern_labels = [&](std::uint16_t pattern) { std::vector labels(width); for (unsigned b = 0; b < width; ++b) { const std::uint8_t bit = static_cast((pattern >> b) & 1u); labels[b] = stack_detail::xor_b(zbit[b], stack_detail::mask_bit(bit, delta)); } return labels; }; const unsigned npat = 1u << width; std::vector seeds(k); std::vector garbled(k); std::size_t naive = 0; std::size_t heaviest = 0; for (std::uint16_t i = 0; i < k; ++i) { seeds[i] = stack_detail::hash_at( stack_detail::fold_labels(pattern_labels(i)), 11); garbled[i] = stack_detail::garble_seeded(branches[i], seeds[i]); naive += garbled[i].tables.size(); heaviest = std::max(heaviest, garbled[i].tables.size()); } const block pad_seed = stack_detail::hash_at(seeds[0], 19); std::vector stacked(heaviest, stack_detail::zero_b()); for (std::size_t i = 0; i < k; ++i) { auto padded = stack_detail::pad_to(garbled[i].tables, heaviest, pad_seed); stacked = stack_detail::xor_vec(stacked, padded); } // Demux: one row per selector color, payload is the XOR of inactive seeds // so the evaluator can subtract them after regenerating each inactive // material. Each inactive seed is delivered explicitly (k blocks per row). std::vector> demux(npat); for (unsigned color = 0; color < npat; ++color) { std::uint16_t sem_idx = 0; std::vector color_labels(width); for (unsigned b = 0; b < width; ++b) { const std::uint8_t cbit = static_cast((color >> b) & 1u); const std::uint8_t sbit = static_cast( cbit ^ stack_detail::lsb(zbit[b])); sem_idx = static_cast(sem_idx | (sbit << b)); color_labels[b] = stack_detail::xor_b(zbit[b], stack_detail::mask_bit(sbit, delta)); } const block row_key = stack_detail::fold_labels(color_labels); demux[color].assign(k, stack_detail::zero_b()); if (sem_idx >= k) continue; for (std::uint16_t j = 0; j < k; ++j) { if (j == sem_idx) continue; demux[color][j] = stack_detail::xor_b(seeds[j], stack_detail::hash_at(row_key, static_cast(100u + j))); } } std::uint16_t active_color = 0; for (unsigned b = 0; b < width; ++b) active_color = static_cast( active_color | (stack_detail::lsb(active_bit[b]) << b)); std::vector active_labels(width); for (unsigned b = 0; b < width; ++b) active_labels[b] = active_bit[b]; const block active_key = stack_detail::fold_labels(active_labels); std::vector inactive_seeds(k); for (std::uint16_t j = 0; j < k; ++j) { if (j == index) continue; inactive_seeds[j] = stack_detail::xor_b(demux[active_color][j], stack_detail::hash_at(active_key, static_cast(100u + j))); } std::vector opened = stacked; for (std::uint16_t j = 0; j < k; ++j) { if (j == index) continue; auto rebuilt = stack_detail::garble_seeded(branches[j], inactive_seeds[j]); auto padded = stack_detail::pad_to(std::move(rebuilt.tables), heaviest, pad_seed); opened = stack_detail::xor_vec(opened, padded); } const std::size_t live = static_cast(branches[index].n_and()) * 2u; if (opened.size() < live) throw std::logic_error("yao stack: one-hot short"); // Re-garble the active branch from the seed the evaluator can derive, // and check the unstacked prefix matches. Evaluation uses `opened`. const block expect_seed = stack_detail::hash_at( stack_detail::fold_labels(pattern_labels(index)), 11); auto active_g = stack_detail::garble_seeded(branches[index], expect_seed); if (active_g.tables.size() != live) throw std::logic_error("yao stack: one-hot rows"); for (std::size_t t = 0; t < live; ++t) { const block diff = stack_detail::xor_b(opened[t], active_g.tables[t]); unsigned char bytes[sizeof(block)]; std::memcpy(bytes, &diff, sizeof(block)); for (unsigned char c : bytes) if (c != 0) throw std::logic_error("yao stack: unstack mismatch"); } if (p0[index].size() < branches[index].n_in() || p1[index].size() < branches[index].n_in()) throw std::invalid_argument("yao stack: input length"); auto outs = stack_detail::eval_labels(branches[index], active_g, opened.data(), p0[index].data(), p1[index].data()); const block common_delta = stack_detail::with_lsb(dpf::uniform_sample(), 1); stack_result result; result.stack_blocks = heaviest; result.naive_blocks = naive; result.extra_blocks = static_cast(npat) * k + static_cast(branches[0].n_out()) * npat * 2u; result.share0.resize(branches[0].n_out()); result.share1.resize(branches[0].n_out()); for (std::uint32_t oi = 0; oi < branches[0].n_out(); ++oi) { const block common_zero = dpf::uniform_sample(); std::vector table(static_cast(npat) * 2u); for (unsigned color = 0; color < npat; ++color) { std::uint16_t sem_idx = 0; std::vector color_labels(width); for (unsigned b = 0; b < width; ++b) { const std::uint8_t cbit = static_cast((color >> b) & 1u); const std::uint8_t sbit = static_cast( cbit ^ stack_detail::lsb(zbit[b])); sem_idx = static_cast(sem_idx | (sbit << b)); color_labels[b] = stack_detail::xor_b(zbit[b], stack_detail::mask_bit(sbit, delta)); } for (unsigned oc = 0; oc < 2; ++oc) { const std::size_t idx = static_cast(color) * 2u + oc; if (sem_idx >= k) { table[idx] = dpf::uniform_sample(); continue; } auto bg = stack_detail::garble_seeded(branches[sem_idx], seeds[sem_idx]); const std::uint32_t oid = branches[sem_idx].outputs()[oi]; const std::uint8_t sem_o = static_cast( oc ^ stack_detail::lsb(bg.z[oid])); const block ol = stack_detail::xor_b(bg.z[oid], stack_detail::mask_bit(sem_o, bg.delta)); const block common = stack_detail::xor_b(common_zero, stack_detail::mask_bit(sem_o, common_delta)); const block key = stack_detail::xor_b( stack_detail::fold_labels(color_labels), ol); table[idx] = stack_detail::xor_b(common, stack_detail::hash_at(key, static_cast(200u + idx))); } } const std::size_t idx = static_cast(active_color) * 2u + stack_detail::lsb(outs[oi]); const block key = stack_detail::xor_b(active_key, outs[oi]); const block got = stack_detail::xor_b(table[idx], stack_detail::hash_at(key, static_cast(200u + idx))); result.share0[oi] = stack_detail::lsb(common_zero); result.share1[oi] = stack_detail::lsb(got); } return result; } } // namespace yao } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_YAO_STACK_HPP__