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>
777 lines
25 KiB
C++
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__
|