libdpf/include/dpf/arith_garble.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

777 lines
25 KiB
C++

/// @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 <array>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <utility>
#include <vector>
#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<std::uint16_t, k_digits> 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<std::uint16_t>((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<std::uint16_t>(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<std::uint16_t>(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<std::uint16_t>(
(static_cast<unsigned>(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<std::uint16_t>(
(static_cast<unsigned>(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<std::uint16_t>(
(static_cast<unsigned>(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<std::uint16_t>((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<std::uint16_t>(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<std::int64_t>(gid), static_cast<std::int64_t>(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<psnip_uint32_t>(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<psnip_uint32_t>(1000u + i + gid));
std::uint64_t lo = 0;
std::memcpy(&lo, &block, sizeof(lo));
out.d[i] = static_cast<std::uint16_t>(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<std::uint16_t>(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<std::uint16_t> factors;
std::uint16_t n = static_cast<std::uint16_t>(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<std::uint16_t>(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<std::uint16_t>((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<lab> 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<std::uint16_t> mask;
std::vector<std::uint16_t> color;
std::vector<std::uint16_t> opened;
std::vector<std::uint16_t> 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<std::uint16_t> 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<std::uint32_t>(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<std::uint32_t>(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<std::uint16_t>(k % a.mod);
nodes_.push_back(std::move(n));
return wire{static_cast<std::uint32_t>(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<std::uint16_t>(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<std::uint32_t>(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<std::uint16_t> 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<std::uint32_t>(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<std::uint32_t>(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<wire> & 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<std::uint16_t>(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<std::uint16_t> 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<wire> & bits)
{
if (bits.empty())
throw std::invalid_argument("arith_garble: and");
const auto mod = static_cast<std::uint16_t>(bits.size() + 1);
std::vector<wire> 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<std::uint16_t>(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<std::uint16_t> dlog(p, 0);
std::vector<std::uint16_t> exp(static_cast<std::size_t>(p - 1), 0);
unsigned acc = 1;
for (std::uint16_t e = 0; e < p - 1; ++e)
{
dlog[acc] = e;
exp[e] = static_cast<std::uint16_t>(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<std::uint16_t>(p - 1), dlog);
auto dy = project(y, static_cast<std::uint16_t>(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<node> & nodes() const noexcept { return nodes_; }
const std::vector<std::uint32_t> & 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<std::uint16_t> zero_flag(std::uint16_t p)
{
std::vector<std::uint16_t> z(p, 0);
z[0] = 1;
return z;
}
std::vector<node> nodes_;
std::vector<std::uint32_t> outs_;
};
namespace detail
{
struct garble_state
{
std::vector<lab> zero;
std::vector<lab> delta;
std::vector<char> have_delta;
std::vector<proj_rows> proj;
std::vector<pass_rows> 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<simde__m128i>();
st.zero.resize(c.nodes().size());
st.delta.assign(static_cast<std::size_t>(k_max_mod) + 1, lab{});
st.have_delta.assign(static_cast<std::size_t>(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<std::uint16_t>((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<std::size_t>(m - 1));
for (std::uint16_t color = 1; color < m; ++color)
{
const std::uint16_t s = static_cast<std::uint16_t>(
(static_cast<unsigned>(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<std::size_t>(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<std::uint16_t>(
(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<std::uint16_t>(sem ^ (pad & 1u));
}
break;
}
}
}
return st;
}
inline std::vector<lab> evaluate(const circuit & c, const garble_state & st,
const std::uint16_t * semantic)
{
const auto & nodes = c.nodes();
std::vector<lab> 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<std::size_t>(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<std::uint16_t>(
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<std::uint16_t> eval_plain(const circuit & c,
const std::uint16_t * semantic)
{
const auto & nodes = c.nodes();
std::vector<std::uint16_t> 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<std::uint16_t>(
(static_cast<unsigned>(s[nd.a]) + s[nd.b]) % nd.mod);
break;
case circuit::op::addk:
s[i] = static_cast<std::uint16_t>(
(static_cast<unsigned>(s[nd.a]) + nd.k) % nd.mod);
break;
case circuit::op::scale:
s[i] = static_cast<std::uint16_t>(
(static_cast<unsigned>(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<std::uint16_t>(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__