/// @file party/flows_gadget.cpp /// @brief Bench flows for word garbling, stacked Yao, FLUTE, and hidden shuffle. /// @details Amalgamated into run.cpp. Every cross-party byte goes through /// `net::trio` (the same local mesh `party_bench` already dials). /// Two-party gadgets use the p0–p1 link. FLUTE's dealer is p2. /// The hidden shuffle is the three ring passes, and the receiver /// feeds the arrived array into the next party step. #include "cases.hpp" #include #include #include #include #include "dpf/arith_garble.hpp" #include "dpf/flute.hpp" #include "dpf/shuffle.hpp" #include "dpf/yao.hpp" #include "dpf/yao_stack.hpp" namespace dpf { namespace party { namespace gadget { using net::role; using net::trio; void require(bool cond, const char * msg) { if (!cond) throw std::runtime_error(msg); } void put_lab(std::vector & blob, const arith_garble::lab & lab) { blob.insert(blob.end(), lab.d.begin(), lab.d.end()); } arith_garble::lab take_lab(const std::uint16_t *& p, const std::uint16_t * end, std::uint16_t mod) { if (static_cast(end - p) < arith_garble::k_digits) throw std::runtime_error("arith garble: short label"); arith_garble::lab lab; lab.mod = mod; for (std::size_t i = 0; i < arith_garble::k_digits; ++i) lab.d[i] = *p++; return lab; } std::uint16_t take_u16(const std::uint16_t *& p, const std::uint16_t * end) { if (p == end) throw std::runtime_error("arith garble: short frame"); return *p++; } /// @brief p0 garbles and ships the evaluator view. p1 evaluates that view. void arith_exchange(trio & net, role self, const arith_garble::circuit & c, const std::uint16_t * semantic) { const auto plain = arith_garble::eval_plain(c, semantic); if (self == role::p0) { auto st = arith_garble::detail::garble(c); auto active = arith_garble::detail::evaluate(c, st, semantic); std::size_t out_i = 0; for (std::uint32_t id : c.outputs()) { const std::uint16_t mod = c.nodes()[id].mod; const std::uint16_t got = arith_garble::open_shares( mod, st.zero[id].d[0], active[id].d[0]); require(got == plain[id], "arith garble: garbler open"); ++out_i; } require(out_i == c.outputs().size(), "arith garble: outputs"); std::vector blob; const auto & nodes = c.nodes(); for (std::uint32_t i = 0; i < nodes.size(); ++i) { const auto & nd = nodes[i]; if (nd.code == arith_garble::circuit::op::in) put_lab(blob, active[i]); else if (nd.code == arith_garble::circuit::op::proj) { for (const auto & row : st.proj[i].row) put_lab(blob, row); } else if (nd.code == arith_garble::circuit::op::pass) { put_lab(blob, st.pass[i].payload[0]); put_lab(blob, st.pass[i].payload[1]); blob.push_back(st.pass[i].flag_ct[0]); blob.push_back(st.pass[i].flag_ct[1]); } } for (std::uint32_t id : c.outputs()) blob.push_back(st.zero[id].d[0]); net.send_vec_to(role::p1, net::msg::bytes, blob); return; } require(self == role::p1, "arith garble: party"); auto blob = net.recv_vec_from(role::p0, net::msg::bytes); const std::uint16_t * p = blob.data(); const std::uint16_t * end = p + blob.size(); arith_garble::detail::garble_state st; const auto & nodes = c.nodes(); st.zero.resize(nodes.size()); st.delta.assign(static_cast(arith_garble::k_max_mod) + 1, arith_garble::lab{}); for (std::uint16_t m = 1; m <= arith_garble::k_max_mod; ++m) st.delta[m].mod = m; st.proj.resize(nodes.size()); st.pass.resize(nodes.size()); std::size_t n_in = 0; for (const auto & nd : nodes) if (nd.code == arith_garble::circuit::op::in) ++n_in; for (std::uint32_t i = 0; i < nodes.size(); ++i) { const auto & nd = nodes[i]; if (nd.code == arith_garble::circuit::op::in) st.zero[i] = take_lab(p, end, nd.mod); else if (nd.code == arith_garble::circuit::op::proj) { const std::uint16_t m = nodes[nd.a].mod; require(m > 0, "arith garble: projection modulus"); st.proj[i].row.resize(static_cast(m - 1)); for (std::uint16_t color = 1; color < m; ++color) st.proj[i].row[static_cast(color - 1)] = take_lab(p, end, nd.mod); } else if (nd.code == arith_garble::circuit::op::pass) { st.pass[i].payload[0] = take_lab(p, end, nd.mod); st.pass[i].payload[1] = take_lab(p, end, nd.mod); st.pass[i].flag_ct[0] = take_u16(p, end); st.pass[i].flag_ct[1] = take_u16(p, end); } } std::vector masks; masks.reserve(c.outputs().size()); for (std::size_t i = 0; i < c.outputs().size(); ++i) masks.push_back(take_u16(p, end)); require(p == end, "arith garble: trailing bytes"); std::vector zeros(n_in, 0); auto active = arith_garble::detail::evaluate(c, st, zeros.data()); std::size_t out_i = 0; for (std::uint32_t id : c.outputs()) { const std::uint16_t mod = nodes[id].mod; const std::uint16_t got = arith_garble::open_shares( mod, masks[out_i], active[id].d[0]); require(got == plain[id], "arith garble: evaluator open"); ++out_i; } } int run_proj(role self, trio & net, std::uint16_t mod) { if (self == role::p2) return 0; arith_garble::circuit c; auto x = c.input(mod); std::vector phi(mod); for (std::uint16_t i = 0; i < mod; ++i) phi[i] = static_cast((i * 3) % mod); c.out(c.project(x, mod, std::move(phi))); const std::uint16_t in = static_cast(mod / 2); arith_exchange(net, self, c, &in); return 0; } int run_mul(role self, trio & net, std::uint16_t p) { if (self == role::p2) return 0; arith_garble::circuit c; auto x = c.input(p); auto y = c.input(p); c.out(c.mul(x, y)); const std::uint16_t in[2] = {static_cast(p - 1), 2}; arith_exchange(net, self, c, in); return 0; } int run_thresh(role self, trio & net, unsigned b) { if (self == role::p2) return 0; arith_garble::circuit c; const auto mod = static_cast(b + 1); std::vector bits; bits.reserve(b); for (unsigned i = 0; i < b; ++i) bits.push_back(c.input(mod)); c.out(c.threshold(bits, static_cast(b / 2))); std::vector in(b, 0); for (unsigned i = 0; i < b; i += 2) in[i] = 1; arith_exchange(net, self, c, in.data()); return 0; } int run_chain(role self, trio & net) { if (self == role::p2) return 0; arith_garble::circuit c; auto x = c.input(7); auto y = c.input(7); auto z = c.mul(x, y); for (int i = 0; i < 3; ++i) z = c.mul(z, x); c.out(z); const std::uint16_t in[2] = {2, 3}; arith_exchange(net, self, c, in); return 0; } yao::netlist and_n(unsigned n) { yao::netlist nl; std::vector in; in.reserve(n); for (unsigned i = 0; i < n; ++i) in.push_back(nl.shared_in()); auto acc = in[0]; for (unsigned i = 1; i < n; ++i) acc = nl.and_(acc, in[i]); nl.out(acc); return nl; } yao::netlist xor2() { yao::netlist nl; auto a = nl.shared_in(); auto b = nl.shared_in(); nl.out(nl.xor_(a, b)); return nl; } /// @brief Garbled tables and the evaluator's labels on the p0–p1 channel. /// @details Party 1 runs `eval_party1` on the bytes that arrived. The choice /// labels are the OT output for party 1's bits; base OT is its own /// bench (`iknp`), not this garble row. void yao_ship(trio & net, role self, const yao::netlist & nl, const std::uint8_t * p0_bits, const std::uint8_t * p1_bits) { const role peer = self == role::p0 ? role::p1 : role::p0; auto & link = net.to(peer); std::size_t nchoice = 0; std::size_t ndirect = 0; for (std::uint32_t i = 0; i < nl.n_in(); ++i) { const auto kind = nl.kind_at(i); if (kind == yao::input::shared || kind == yao::input::priv1) ++nchoice; if (kind == yao::input::shared || kind == yao::input::priv0) ++ndirect; } const std::size_t ntable = static_cast(nl.n_and()) * 2u; std::vector share; if (self == role::p0) { auto g = yao::detail::garble_party0(nl, p0_bits); require(g.tables.size() == ntable, "yao: table size"); require(g.direct.size() == ndirect, "yao: direct labels"); require(g.ot0.size() == nchoice && g.ot1.size() == nchoice, "yao: ot labels"); std::vector chosen(nchoice); std::size_t oi = 0; for (std::uint32_t i = 0; i < nl.n_in(); ++i) { const auto kind = nl.kind_at(i); if (kind != yao::input::shared && kind != yao::input::priv1) continue; chosen[oi] = (p1_bits[i] & 1u) ? g.ot1[oi] : g.ot0[oi]; ++oi; } std::vector blob; blob.reserve(ntable + ndirect + nchoice); blob.insert(blob.end(), g.tables.begin(), g.tables.end()); blob.insert(blob.end(), g.direct.begin(), g.direct.end()); blob.insert(blob.end(), chosen.begin(), chosen.end()); link.send_vec(blob, net::msg::bytes); share = std::move(g.share); } else { const auto blob = link.recv_vec(net::msg::bytes); require(blob.size() == ntable + ndirect + nchoice, "yao: payload size"); std::vector tables(ntable); std::vector direct(ndirect); std::vector chosen(nchoice); if (ntable != 0) std::memcpy(tables.data(), blob.data(), ntable * sizeof(yao::block)); if (ndirect != 0) std::memcpy(direct.data(), blob.data() + ntable, ndirect * sizeof(yao::block)); if (nchoice != 0) std::memcpy(chosen.data(), blob.data() + ntable + ndirect, nchoice * sizeof(yao::block)); share = yao::detail::eval_party1(nl, tables, direct, chosen); } auto other = net.exchange_vec_with(peer, share, net::msg::delta); require(other.size() == share.size(), "yao: share length"); std::vector semantic(nl.n_in()); for (std::uint32_t i = 0; i < nl.n_in(); ++i) semantic[i] = static_cast(p0_bits[i] ^ p1_bits[i]); auto plain = yao::eval_plain(nl, semantic.data()); require(plain.size() == share.size(), "yao: output length"); for (std::size_t i = 0; i < plain.size(); ++i) { const auto opened = static_cast(share[i] ^ other[i]); require(opened == plain[i], "yao: open"); } } int run_if(role self, trio & net, unsigned heavy, unsigned light) { if (self == role::p2) return 0; auto then_nl = and_n(heavy); std::vector h0(heavy, 1); std::vector h1(heavy, 0); if (self == role::p0) { auto else_nl = and_n(light); std::vector l0(light, 1); std::vector l1(light, 1); auto stacked = yao::eval_if(then_nl, else_nl, 0, 0, h0.data(), h1.data(), l0.data(), l1.data()); std::vector semantic(heavy, 1); auto plain = yao::eval_plain(then_nl, semantic.data()); require(stacked.share0.size() == plain.size(), "yao stack: outputs"); for (std::size_t i = 0; i < plain.size(); ++i) { const auto opened = static_cast( stacked.share0[i] ^ stacked.share1[i]); require(opened == plain[i], "yao stack: if"); } require(stacked.stack_blocks > 0, "yao stack: empty"); } yao_ship(net, self, then_nl, h0.data(), h1.data()); return 0; } int run_hot(role self, trio & net, unsigned k) { if (self == role::p2) return 0; auto active = and_n(2); const std::uint8_t bits0[2] = {1, 0}; const std::uint8_t bits1[2] = {0, 1}; if (self == role::p0) { std::vector branches; std::vector> p0; std::vector> p1; branches.reserve(k); for (unsigned i = 0; i < k; ++i) { branches.push_back(i % 2 == 0 ? and_n(2) : xor2()); p0.push_back({1, static_cast(i & 1u)}); p1.push_back({0, 1}); } auto stacked = yao::eval_one_hot(branches, 0, 0, p0, p1); const std::uint8_t semantic[2] = {1, 1}; auto plain = yao::eval_plain(active, semantic); require(stacked.share0.size() == plain.size(), "yao stack: one-hot outputs"); for (std::size_t i = 0; i < plain.size(); ++i) { const auto opened = static_cast( stacked.share0[i] ^ stacked.share1[i]); require(opened == plain[i], "yao stack: one-hot"); } require(stacked.stack_blocks > 0, "yao stack: one-hot empty"); } yao_ship(net, self, active, bits0, bits1); return 0; } std::vector flute_columns(unsigned delta, unsigned n_out) { const unsigned rows = 1u << delta; std::vector columns(static_cast(n_out) * rows); for (unsigned w = 0; w < n_out; ++w) for (unsigned j = 0; j < rows; ++j) columns[static_cast(w) * rows + j] = static_cast(((j >> (w % delta)) ^ w) & 1u); return columns; } std::vector flute_bits(unsigned delta) { std::vector bits(delta, 0); bits[0] = 1; if (delta > 2) bits[2] = 1; return bits; } std::vector flute_pack(const flute::detail::setup & s, unsigned party) { std::vector out; out.insert(out.end(), s.m.begin(), s.m.end()); out.insert(out.end(), s.share[party].begin(), s.share[party].end()); out.insert(out.end(), s.lamz[party].begin(), s.lamz[party].end()); return out; } int run_flute(role self, trio & net, unsigned delta, unsigned n_out) { const auto columns = flute_columns(delta, n_out); const auto x = flute_bits(delta); const unsigned rows = 1u << delta; const std::uint32_t full = rows - 1u; if (self == role::p2) { auto setup = flute::detail::make_setup(delta, n_out, 2, x.data()); auto to0 = flute_pack(setup, 0); auto to1 = flute_pack(setup, 1); net.send_bytes_to(role::p0, net::msg::bytes, to0.data(), to0.size()); net.send_bytes_to(role::p1, net::msg::bytes, to1.data(), to1.size()); return 0; } const unsigned party = self == role::p0 ? 0u : 1u; auto bytes = net.recv_bytes_from(role::p2, net::msg::bytes); const std::size_t need = static_cast(delta) + rows + n_out; require(bytes.size() == need, "flute: deal size"); flute::detail::setup setup; setup.m.assign(bytes.begin(), bytes.begin() + delta); setup.share.assign(2, {}); setup.share[party].assign(bytes.begin() + delta, bytes.begin() + delta + rows); setup.lamz.assign(2, {}); setup.lamz[party].assign(bytes.end() - static_cast(n_out), bytes.end()); std::vector mine(static_cast(n_out) * 2u); for (unsigned w = 0; w < n_out; ++w) { const std::uint8_t * column = columns.data() + static_cast(w) * rows; mine[w] = flute::detail::party_v(setup, party, delta, full, column, setup.lamz[party][w]); mine[n_out + w] = setup.lamz[party][w]; } const role peer = self == role::p0 ? role::p1 : role::p0; auto theirs = net.exchange_vec_with(peer, mine, net::msg::bytes); auto expect = flute::eval_plain(delta, n_out, columns.data(), x.data()); for (unsigned w = 0; w < n_out; ++w) { const std::uint8_t * column = columns.data() + static_cast(w) * rows; const std::uint8_t t = flute::detail::dot_column( full, delta, setup.m.data(), column); const auto opened = static_cast( mine[w] ^ theirs[w] ^ t ^ mine[n_out + w] ^ theirs[n_out + w]); require(opened == expect[w], "flute: open"); } return 0; } rss::seed_bundle shuffle_bundle() { rss::seed_bundle bundle{}; auto fill = [](rss::seed_block & seed, std::uint8_t tag) { auto * p = reinterpret_cast(&seed); for (std::size_t i = 0; i < sizeof(seed); ++i) p[i] = static_cast(tag + i * 17u); }; fill(bundle.k01, 1); fill(bundle.k12, 2); fill(bundle.k20, 3); return bundle; } shuffle::shuffle_party_view shuffle_view(unsigned me, std::size_t n) { std::uint64_t rng = 0xA5A5A5A5A5A5A5A5ull; auto draw = [&] { rng = rng * 6364136223846793005ull + 1u; return rng; }; shuffle::shuffle_party_view view; view.own.resize(n); view.next.resize(n); for (std::size_t i = 0; i < n; ++i) { const auto a = draw(); const auto b = draw(); const auto c = static_cast(i) - a - b; if (me == 0) { view.own[i] = a; view.next[i] = b; } else if (me == 1) { view.own[i] = b; view.next[i] = c; } else { view.own[i] = c; view.next[i] = a; } } return view; } int run_shuffle(role self, trio & net, std::size_t n) { const unsigned me = static_cast(self); const auto bundle = shuffle_bundle(); const std::uint64_t index = 1; auto seeds = rss::party_seeds::from_bundle(bundle, me); auto view = shuffle_view(me, n); const unsigned order[3] = {2u, 0u, 1u}; for (unsigned left : order) { const unsigned u = shuffle::hidden_u_party(left); const unsigned side = shuffle::hidden_side_party(left); if (me == u) { auto step = shuffle::shuffle_hidden_pass( me, seeds, view, index, left, nullptr); require(step.out.sends, "shuffle: u sends"); require(step.out.to == side, "shuffle: u recipient"); net.send_vec_to(static_cast(step.out.to), net::msg::ring_vector, step.out.data); view = std::move(step.view); } else if (me == side) { auto inbound = net.recv_vec_from( static_cast(u), net::msg::ring_vector); auto step = shuffle::shuffle_hidden_pass( me, seeds, view, index, left, &inbound); require(step.out.sends, "shuffle: side sends"); require(step.out.to == left, "shuffle: side recipient"); net.send_vec_to(static_cast(step.out.to), net::msg::ring_vector, step.out.data); view = std::move(step.view); } else { auto inbound = net.recv_vec_from( static_cast(side), net::msg::ring_vector); auto step = shuffle::shuffle_hidden_pass( me, seeds, view, index, left, &inbound); require(!step.out.sends, "shuffle: left-out is silent"); view = std::move(step.view); } } if (self == role::p0) { auto from1 = net.recv_vec_from(role::p1, net::msg::ring_vector); auto from2 = net.recv_vec_from(role::p2, net::msg::ring_vector); require(from1.size() == n && from2.size() == n, "shuffle: open size"); std::vector clear(n); for (std::size_t i = 0; i < n; ++i) clear[i] = i; auto expect = shuffle::permute(clear, shuffle::permutation_from_seed(bundle.k01, n, index)); expect = shuffle::permute(expect, shuffle::permutation_from_seed(bundle.k12, n, index)); expect = shuffle::permute(expect, shuffle::permutation_from_seed(bundle.k20, n, index)); for (std::size_t i = 0; i < n; ++i) { const auto got = view.own[i] + from1[i] + from2[i]; require(got == expect[i], "shuffle: open"); } } else { net.send_vec_to(role::p0, net::msg::ring_vector, view.own); } return 0; } int arith_proj_m5(role self, trio & net) { return run_proj(self, net, 5); } int arith_proj_m17(role self, trio & net) { return run_proj(self, net, 17); } int arith_proj_m64(role self, trio & net) { return run_proj(self, net, 64); } int arith_mul_p5(role self, trio & net) { return run_mul(self, net, 5); } int arith_mul_p7(role self, trio & net) { return run_mul(self, net, 7); } int arith_mul_p11(role self, trio & net) { return run_mul(self, net, 11); } int arith_thresh_b8(role self, trio & net) { return run_thresh(self, net, 8); } int arith_thresh_b16(role self, trio & net) { return run_thresh(self, net, 16); } int arith_chain_mul4(role self, trio & net) { return run_chain(self, net); } int yao_if_4_2(role self, trio & net) { return run_if(self, net, 4, 2); } int yao_if_16_8(role self, trio & net) { return run_if(self, net, 16, 8); } int yao_onehot_k4(role self, trio & net) { return run_hot(self, net, 4); } int yao_onehot_k8(role self, trio & net) { return run_hot(self, net, 8); } int flute_d2(role self, trio & net) { return run_flute(self, net, 2, 1); } int flute_d4(role self, trio & net) { return run_flute(self, net, 4, 1); } int flute_d8(role self, trio & net) { return run_flute(self, net, 8, 1); } int flute_d4_o8(role self, trio & net) { return run_flute(self, net, 4, 8); } int shuffle_n16(role self, trio & net) { return run_shuffle(self, net, 16); } int shuffle_n64(role self, trio & net) { return run_shuffle(self, net, 64); } int shuffle_n256(role self, trio & net) { return run_shuffle(self, net, 256); } #define REG(name, tags, fn) \ register_flow(flow{#name, tags, fn, true}) } // namespace gadget void register_gadget_flows() { using namespace gadget; REG(arith_proj_m5, "bench arith garble", arith_proj_m5); REG(arith_proj_m17, "bench arith garble", arith_proj_m17); REG(arith_proj_m64, "bench arith garble", arith_proj_m64); REG(arith_mul_p5, "bench arith garble", arith_mul_p5); REG(arith_mul_p7, "bench arith garble", arith_mul_p7); REG(arith_mul_p11, "bench arith garble", arith_mul_p11); REG(arith_thresh_b8, "bench arith garble", arith_thresh_b8); REG(arith_thresh_b16, "bench arith garble", arith_thresh_b16); REG(arith_chain_mul4, "bench arith garble", arith_chain_mul4); REG(yao_if_4_2, "bench yao stack", yao_if_4_2); REG(yao_if_16_8, "bench yao stack", yao_if_16_8); REG(yao_onehot_k4, "bench yao stack", yao_onehot_k4); REG(yao_onehot_k8, "bench yao stack", yao_onehot_k8); REG(flute_d2, "bench flute", flute_d2); REG(flute_d4, "bench flute", flute_d4); REG(flute_d8, "bench flute", flute_d8); REG(flute_d4_o8, "bench flute", flute_d4_o8); REG(shuffle_n16, "bench shuffle rss", shuffle_n16); REG(shuffle_n64, "bench shuffle rss", shuffle_n64); REG(shuffle_n256, "bench shuffle rss", shuffle_n256); } } // namespace party } // namespace dpf