/// @file dpf/arith_garble.hpp /// @brief Constant-round arithmetic garbling gadgets. /// @details Mixed-modulus circuits: free addition, free multiplication by a /// public constant coprime to the modulus, and a unary projection. /// A projection of modulus `m` sends `m - 1` ciphertexts (row /// reduction). A symmetric boolean gate, including a high fan-in AND /// or a threshold, is a projection of a free sum. Multiplication in a /// small prime field is the discrete-log reduction: project to the /// exponent, add, project back, and suppress the zero cases. /// /// Labels are vectors in `(Z_m)^k`. Digit 0 is the point-and-permute /// color. The global offset `Δ_m` has color digit 1, so the color of /// semantic value `s` is `τ + s`. Addition and public scaling are /// componentwise in that group, which is what makes them free. /// /// This is not an ABY2.0 session. A session opens one masked wire and /// then multiplies interactively. These gadgets never open an /// intermediate wire: the evaluator finishes from the garbled rows. /// @note Marshall Ball, Tal Malkin, and Mike Rosulek, "Garbling Gadgets for /// Boolean and Arithmetic Circuits," CCS 2016 (ePrint 2016/969). The /// free-addition offset is the one they attribute to Malkin, Pastro, and /// shelat. The interactive product in `beaver.hpp` remains Patra, /// Schneider, Suresh, and Yalame, USENIX Security 2021. /// @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_ARITH_GARBLE_HPP__ #define LIBDPF_INCLUDE_DPF_ARITH_GARBLE_HPP__ #include #include #include #include #include #include #include #include "hedley/hedley.h" #include "simde/simde/x86/avx2.h" #include "dpf/prg_aes.hpp" #include "dpf/random.hpp" namespace dpf { namespace arith_garble { /// @brief Digit width of one label, including the color digit. /// @details Ball, Malkin, and Rosulek use `λ / log2(m)` payload digits so the /// label is `λ` bits. This instantiation fixes the width. Digit 0 is /// the color in every modulus. inline constexpr std::size_t k_digits = 16; inline constexpr std::uint16_t k_max_mod = 128; struct lab { std::uint16_t mod = 0; std::array d{}; }; /// @brief Open a masked color. Garbler holds `mask` (the color of semantic 0). /// Evaluator holds `color`. HEDLEY_WARN_UNUSED_RESULT inline std::uint16_t open_shares(std::uint16_t mod, std::uint16_t mask, std::uint16_t color) { if (mod == 0) throw std::invalid_argument("arith_garble: modulus"); return static_cast((color + mod - (mask % mod)) % mod); } namespace detail { inline void require_mod(std::uint16_t m) { if (m < 2 || m > k_max_mod) throw std::invalid_argument("arith_garble: modulus"); } HEDLEY_WARN_UNUSED_RESULT inline std::uint16_t gcd_u(std::uint16_t a, std::uint16_t b) { while (b != 0) { const std::uint16_t t = static_cast(a % b); a = b; b = t; } return a; } HEDLEY_WARN_UNUSED_RESULT inline lab sample_lab(std::uint16_t mod, simde__m128i & seed, std::uint32_t & n) { lab out; out.mod = mod; for (std::size_t i = 0; i < k_digits; ++i) { const auto block = prg::aes128::eval(seed, n++); std::uint64_t lo = 0; std::memcpy(&lo, &block, sizeof(lo)); out.d[i] = static_cast(lo % mod); } return out; } HEDLEY_WARN_UNUSED_RESULT inline lab make_delta(std::uint16_t mod, simde__m128i & seed, std::uint32_t & n) { lab d = sample_lab(mod, seed, n); d.d[0] = 1; return d; } HEDLEY_WARN_UNUSED_RESULT inline lab add_lab(const lab & a, const lab & b) { if (a.mod != b.mod) throw std::invalid_argument("arith_garble: modulus"); lab out; out.mod = a.mod; for (std::size_t i = 0; i < k_digits; ++i) out.d[i] = static_cast( (static_cast(a.d[i]) + b.d[i]) % a.mod); return out; } HEDLEY_WARN_UNUSED_RESULT inline lab sub_lab(const lab & a, const lab & b) { if (a.mod != b.mod) throw std::invalid_argument("arith_garble: modulus"); lab out; out.mod = a.mod; for (std::size_t i = 0; i < k_digits; ++i) out.d[i] = static_cast( (static_cast(a.d[i]) + a.mod - b.d[i]) % a.mod); return out; } HEDLEY_WARN_UNUSED_RESULT inline lab scale_lab(const lab & a, std::uint16_t c) { lab out; out.mod = a.mod; for (std::size_t i = 0; i < k_digits; ++i) out.d[i] = static_cast( (static_cast(a.d[i]) * c) % a.mod); return out; } HEDLEY_WARN_UNUSED_RESULT inline lab neg_lab(const lab & a) { lab out; out.mod = a.mod; for (std::size_t i = 0; i < k_digits; ++i) out.d[i] = static_cast((a.mod - a.d[i]) % a.mod); return out; } /// @brief Label of semantic `s` on a wire whose semantic-0 label is `zero`. HEDLEY_WARN_UNUSED_RESULT inline lab shift_lab(const lab & zero, const lab & delta, std::uint16_t s) { return add_lab(zero, scale_lab(delta, static_cast(s % zero.mod))); } HEDLEY_WARN_UNUSED_RESULT inline lab hash_lab(std::uint32_t gid, std::uint32_t which, const lab & in, std::uint16_t out_mod) { const prg::purpose_scope counted(prg::purpose::hash); simde__m128i acc = simde_mm_set_epi64x( static_cast(gid), static_cast(which)); for (std::size_t i = 0; i < k_digits; i += 8) { simde__m128i chunk; std::memcpy(&chunk, in.d.data() + i, sizeof(chunk)); acc = prg::aes128::eval(simde_mm_xor_si128(acc, chunk), static_cast(i + which)); } lab out; out.mod = out_mod; for (std::size_t i = 0; i < k_digits; ++i) { const auto block = prg::aes128::eval(acc, static_cast(1000u + i + gid)); std::uint64_t lo = 0; std::memcpy(&lo, &block, sizeof(lo)); out.d[i] = static_cast(lo % out_mod); } return out; } HEDLEY_WARN_UNUSED_RESULT inline std::uint16_t pow_mod(std::uint16_t base, std::uint16_t exp, std::uint16_t mod) { unsigned r = 1; unsigned b = base % mod; unsigned e = exp; while (e != 0) { if ((e & 1u) != 0) r = (r * b) % mod; b = (b * b) % mod; e >>= 1u; } return static_cast(r); } HEDLEY_WARN_UNUSED_RESULT inline bool is_prime(std::uint16_t p) { if (p < 2) return false; for (std::uint16_t i = 2; i * i <= p; ++i) if (p % i == 0) return false; return true; } HEDLEY_WARN_UNUSED_RESULT inline std::uint16_t primitive_root(std::uint16_t p) { std::vector factors; std::uint16_t n = static_cast(p - 1); for (std::uint16_t i = 2; i * i <= n; ++i) { if (n % i != 0) continue; factors.push_back(i); while (n % i == 0) n = static_cast(n / i); } if (n > 1) factors.push_back(n); for (std::uint16_t g = 2; g < p; ++g) { bool ok = true; for (std::uint16_t f : factors) { if (pow_mod(g, static_cast((p - 1) / f), p) == 1) { ok = false; break; } } if (ok) return g; } throw std::logic_error("arith_garble: primitive root"); } struct proj_rows { std::vector row; }; struct pass_rows { lab payload[2]{}; std::uint16_t flag_ct[2]{}; }; } // namespace detail /// @brief One wire. Valid only for the circuit that minted it. struct wire { std::uint32_t id = 0; }; /// @brief Shares of one evaluation. `opened[i] = color[i] - mask[i]` mod the /// output modulus. `mask` is the garbler's share. `color` is the /// evaluator's share. struct shares { std::vector mask; std::vector color; std::vector opened; std::vector modulus; /// @brief Projection rows on the wire (`m - 1` each) plus two per bit-scale. std::size_t ciphertext_rows = 0; }; /// @brief Straight-line mixed-modulus circuit. /// @details Inputs are declared first. Every later wire names earlier wires. class circuit { public: enum class op : unsigned char { in = 0, add = 1, addk = 2, scale = 3, proj = 4, pass = 5 }; struct node { op code = op::in; std::uint16_t mod = 0; std::uint32_t a = 0; std::uint32_t b = 0; std::uint16_t k = 0; std::vector phi; }; HEDLEY_WARN_UNUSED_RESULT wire input(std::uint16_t mod) { detail::require_mod(mod); node n; n.code = op::in; n.mod = mod; nodes_.push_back(std::move(n)); return wire{static_cast(nodes_.size() - 1)}; } HEDLEY_WARN_UNUSED_RESULT wire add(wire x, wire y) { const node & a = at(x); const node & b = at(y); if (a.mod != b.mod) throw std::invalid_argument("arith_garble: add modulus"); node n; n.code = op::add; n.mod = a.mod; n.a = x.id; n.b = y.id; nodes_.push_back(std::move(n)); return wire{static_cast(nodes_.size() - 1)}; } /// @brief Add a public constant. No ciphertext and no evaluator label change. HEDLEY_WARN_UNUSED_RESULT wire add_const(wire x, std::uint16_t k) { const node & a = at(x); node n; n.code = op::addk; n.mod = a.mod; n.a = x.id; n.k = static_cast(k % a.mod); nodes_.push_back(std::move(n)); return wire{static_cast(nodes_.size() - 1)}; } /// @brief Multiply by a public constant coprime to the modulus. HEDLEY_WARN_UNUSED_RESULT wire scale(wire x, std::uint16_t c) { const node & a = at(x); c = static_cast(c % a.mod); if (detail::gcd_u(c, a.mod) != 1) throw std::invalid_argument("arith_garble: scale not coprime"); node n; n.code = op::scale; n.mod = a.mod; n.a = x.id; n.k = c; nodes_.push_back(std::move(n)); return wire{static_cast(nodes_.size() - 1)}; } /// @brief Unary map `phi : Z_mod(x) → Z_out`. `phi.size()` is the input modulus. /// The garbled row count is `phi.size() - 1`. HEDLEY_WARN_UNUSED_RESULT wire project(wire x, std::uint16_t out_mod, std::vector phi) { detail::require_mod(out_mod); const node & a = at(x); if (phi.size() != a.mod) throw std::invalid_argument("arith_garble: projection table"); for (std::uint16_t v : phi) if (v >= out_mod) throw std::invalid_argument("arith_garble: projection image"); node n; n.code = op::proj; n.mod = out_mod; n.a = x.id; n.phi = std::move(phi); nodes_.push_back(std::move(n)); return wire{static_cast(nodes_.size() - 1)}; } /// @brief `bit ? word : 0`. `bit` is mod 2. Two ciphertext rows. HEDLEY_WARN_UNUSED_RESULT wire bit_scale(wire word, wire bit) { const node & w = at(word); const node & b = at(bit); if (b.mod != 2) throw std::invalid_argument("arith_garble: bit_scale bit"); node n; n.code = op::pass; n.mod = w.mod; n.a = word.id; n.b = bit.id; nodes_.push_back(std::move(n)); return wire{static_cast(nodes_.size() - 1)}; } /// @brief AND (or threshold `t`) of 0/1 wires that already live in `Z_{b+1}`. /// @details The sum is free. The only rows are the final projection, `b` /// ciphertexts, as in Section 5 of Ball, Malkin, and Rosulek. /// The wires must have modulus `bits.size() + 1` and semantic /// values in `{0, 1}`. HEDLEY_WARN_UNUSED_RESULT wire threshold(const std::vector & bits, std::uint16_t t) { if (bits.empty() || bits.size() > k_max_mod - 1) throw std::invalid_argument("arith_garble: threshold fan-in"); const auto mod = static_cast(bits.size() + 1); if (t > bits.size()) throw std::invalid_argument("arith_garble: threshold"); wire acc = bits[0]; if (at(acc).mod != mod) throw std::invalid_argument("arith_garble: threshold modulus"); for (std::size_t i = 1; i < bits.size(); ++i) { if (at(bits[i]).mod != mod) throw std::invalid_argument("arith_garble: threshold modulus"); acc = add(acc, bits[i]); } std::vector phi(mod, 0); phi[t] = 1; return project(acc, 2, std::move(phi)); } /// @brief Fan-in AND of mod-2 bits. Lifts into `Z_{b+1}`, then `threshold`. HEDLEY_WARN_UNUSED_RESULT wire fanin_and(const std::vector & bits) { if (bits.empty()) throw std::invalid_argument("arith_garble: and"); const auto mod = static_cast(bits.size() + 1); std::vector lifted; lifted.reserve(bits.size()); for (wire b : bits) { if (at(b).mod != 2) throw std::invalid_argument("arith_garble: and bit"); lifted.push_back(project(b, mod, {0, 1})); } return threshold(lifted, static_cast(bits.size())); } /// @brief Product in a prime field, via discrete log. Mod-2 product is AND. HEDLEY_WARN_UNUSED_RESULT wire mul(wire x, wire y) { const node & a = at(x); const node & b = at(y); if (a.mod != b.mod) throw std::invalid_argument("arith_garble: mul modulus"); const std::uint16_t p = a.mod; if (p == 2) { auto lx = project(x, 3, {0, 1}); auto ly = project(y, 3, {0, 1}); auto s = add(lx, ly); return project(s, 2, {0, 0, 1}); } if (!detail::is_prime(p)) throw std::invalid_argument("arith_garble: mul prime"); const std::uint16_t g = detail::primitive_root(p); std::vector dlog(p, 0); std::vector exp(static_cast(p - 1), 0); unsigned acc = 1; for (std::uint16_t e = 0; e < p - 1; ++e) { dlog[acc] = e; exp[e] = static_cast(acc); acc = (acc * g) % p; } auto zx = project(x, 2, zero_flag(p)); auto zy = project(y, 2, zero_flag(p)); auto dx = project(x, static_cast(p - 1), dlog); auto dy = project(y, static_cast(p - 1), std::move(dlog)); auto ds = add(dx, dy); auto gpow = project(ds, p, std::move(exp)); auto z1 = project(zx, 3, {0, 1}); auto z2 = project(zy, 3, {0, 1}); auto zsum = add(z1, z2); auto zor = project(zsum, 2, {0, 1, 1}); auto nz = project(zor, 2, {1, 0}); return bit_scale(gpow, nz); } void out(wire x) { at(x); outs_.push_back(x.id); } const std::vector & nodes() const noexcept { return nodes_; } const std::vector & outputs() const noexcept { return outs_; } std::uint16_t modulus_at(wire x) const { return at(x).mod; } private: const node & at(wire x) const { if (x.id >= nodes_.size()) throw std::invalid_argument("arith_garble: wire"); return nodes_[x.id]; } static std::vector zero_flag(std::uint16_t p) { std::vector z(p, 0); z[0] = 1; return z; } std::vector nodes_; std::vector outs_; }; namespace detail { struct garble_state { std::vector zero; std::vector delta; std::vector have_delta; std::vector proj; std::vector pass; simde__m128i seed{}; std::uint32_t n = 0; lab & delta_of(std::uint16_t mod) { if (!have_delta[mod]) { delta[mod] = make_delta(mod, seed, n); have_delta[mod] = 1; } return delta[mod]; } }; inline garble_state garble(const circuit & c) { garble_state st; st.seed = dpf::uniform_sample(); st.zero.resize(c.nodes().size()); st.delta.assign(static_cast(k_max_mod) + 1, lab{}); st.have_delta.assign(static_cast(k_max_mod) + 1, 0); st.proj.resize(c.nodes().size()); st.pass.resize(c.nodes().size()); const auto & nodes = c.nodes(); for (std::uint32_t i = 0; i < nodes.size(); ++i) { const auto & nd = nodes[i]; switch (nd.code) { case circuit::op::in: st.zero[i] = sample_lab(nd.mod, st.seed, st.n); (void)st.delta_of(nd.mod); break; case circuit::op::add: st.zero[i] = add_lab(st.zero[nd.a], st.zero[nd.b]); break; case circuit::op::addk: st.zero[i] = sub_lab(st.zero[nd.a], scale_lab(st.delta_of(nd.mod), nd.k)); break; case circuit::op::scale: st.zero[i] = scale_lab(st.zero[nd.a], nd.k); break; case circuit::op::proj: { const std::uint16_t m = nodes[nd.a].mod; const std::uint16_t nmod = nd.mod; const lab & zin = st.zero[nd.a]; const lab & din = st.delta_of(m); const lab & dout = st.delta_of(nmod); const std::uint16_t tau = zin.d[0]; const std::uint16_t s0 = static_cast((m - (tau % m)) % m); const lab label0 = shift_lab(zin, din, s0); const lab h0 = hash_lab(i, 0, label0, nmod); const lab decrypted = neg_lab(h0); const std::uint16_t p0 = nd.phi[s0]; st.zero[i] = sub_lab(decrypted, scale_lab(dout, p0)); st.proj[i].row.resize(static_cast(m - 1)); for (std::uint16_t color = 1; color < m; ++color) { const std::uint16_t s = static_cast( (static_cast(color) + m - tau) % m); const lab label = shift_lab(zin, din, s); const lab active = shift_lab(st.zero[i], dout, nd.phi[s]); const lab h = hash_lab(i, color, label, nmod); st.proj[i].row[static_cast(color - 1)] = add_lab(active, h); } break; } case circuit::op::pass: { const std::uint16_t p = nd.mod; const lab & zw = st.zero[nd.a]; const lab & zb = st.zero[nd.b]; (void)st.delta_of(p); (void)st.delta_of(2); st.zero[i] = sample_lab(p, st.seed, st.n); const std::uint16_t tau_b = zb.d[0]; const lab addend = sub_lab(st.zero[i], zw); for (std::uint16_t color = 0; color < 2; ++color) { const std::uint16_t sem = static_cast( (color + 2u - (tau_b % 2u)) % 2u); const lab bit_label = shift_lab(zb, st.delta_of(2), sem); const lab h = hash_lab(i, color, bit_label, p); const lab payload = (sem == 0) ? st.zero[i] : addend; st.pass[i].payload[color] = add_lab(payload, h); const std::uint16_t pad = hash_lab(i, 8u + color, bit_label, 2).d[0]; st.pass[i].flag_ct[color] = static_cast(sem ^ (pad & 1u)); } break; } } } return st; } inline std::vector evaluate(const circuit & c, const garble_state & st, const std::uint16_t * semantic) { const auto & nodes = c.nodes(); std::vector active(nodes.size()); std::uint32_t in_i = 0; for (std::uint32_t i = 0; i < nodes.size(); ++i) { const auto & nd = nodes[i]; switch (nd.code) { case circuit::op::in: { if (semantic == nullptr) throw std::invalid_argument("arith_garble: inputs"); if (semantic[in_i] >= nd.mod) throw std::invalid_argument("arith_garble: input range"); active[i] = shift_lab(st.zero[i], st.delta[nd.mod], semantic[in_i]); ++in_i; break; } case circuit::op::add: active[i] = add_lab(active[nd.a], active[nd.b]); break; case circuit::op::addk: active[i] = active[nd.a]; break; case circuit::op::scale: active[i] = scale_lab(active[nd.a], nd.k); break; case circuit::op::proj: { const std::uint16_t color = active[nd.a].d[0]; const lab h = hash_lab(i, color, active[nd.a], nd.mod); if (color == 0) active[i] = neg_lab(h); else active[i] = sub_lab( st.proj[i].row[static_cast(color - 1)], h); break; } case circuit::op::pass: { const std::uint16_t color = active[nd.b].d[0]; const lab h = hash_lab(i, color, active[nd.b], nd.mod); const lab payload = sub_lab(st.pass[i].payload[color], h); const std::uint16_t pad = hash_lab(i, 8u + color, active[nd.b], 2).d[0]; const std::uint16_t flag = static_cast( st.pass[i].flag_ct[color] ^ (pad & 1u)); if (flag == 0) active[i] = payload; else active[i] = add_lab(active[nd.a], payload); break; } } } return active; } } // namespace detail /// @brief Clear semantics, one value per wire, inputs in wire order. HEDLEY_WARN_UNUSED_RESULT inline std::vector eval_plain(const circuit & c, const std::uint16_t * semantic) { const auto & nodes = c.nodes(); std::vector s(nodes.size(), 0); std::uint32_t in_i = 0; for (std::uint32_t i = 0; i < nodes.size(); ++i) { const auto & nd = nodes[i]; switch (nd.code) { case circuit::op::in: if (semantic == nullptr || semantic[in_i] >= nd.mod) throw std::invalid_argument("arith_garble: input"); s[i] = semantic[in_i++]; break; case circuit::op::add: s[i] = static_cast( (static_cast(s[nd.a]) + s[nd.b]) % nd.mod); break; case circuit::op::addk: s[i] = static_cast( (static_cast(s[nd.a]) + nd.k) % nd.mod); break; case circuit::op::scale: s[i] = static_cast( (static_cast(s[nd.a]) * nd.k) % nd.mod); break; case circuit::op::proj: s[i] = nd.phi[s[nd.a]]; break; case circuit::op::pass: s[i] = (s[nd.b] != 0) ? s[nd.a] : static_cast(0); break; } } return s; } /// @brief Garble and evaluate in one process. /// @details `semantic` is one value per `input` call, in that order. HEDLEY_WARN_UNUSED_RESULT inline shares eval_pair(const circuit & c, const std::uint16_t * semantic) { if (c.outputs().empty()) throw std::invalid_argument("arith_garble: no outputs"); auto st = detail::garble(c); auto active = detail::evaluate(c, st, semantic); shares out; std::size_t rows = 0; for (std::uint32_t i = 0; i < c.nodes().size(); ++i) { if (c.nodes()[i].code == circuit::op::proj) rows += st.proj[i].row.size(); else if (c.nodes()[i].code == circuit::op::pass) rows += 2; } out.ciphertext_rows = rows; for (std::uint32_t id : c.outputs()) { const std::uint16_t mod = c.nodes()[id].mod; const std::uint16_t mask = st.zero[id].d[0]; const std::uint16_t color = active[id].d[0]; out.modulus.push_back(mod); out.mask.push_back(mask); out.color.push_back(color); out.opened.push_back(open_shares(mod, mask, color)); } return out; } } // namespace arith_garble } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_ARITH_GARBLE_HPP__