libdpf/include/dpf/yao_stack.hpp

578 lines
22 KiB
C++
Raw Permalink Normal View History

/// @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 <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <utility>
#include <vector>
#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<std::uint8_t> share0;
std::vector<std::uint8_t> 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<block> z;
std::vector<block> tables;
std::vector<block> lg;
std::vector<block> 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<std::size_t>(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<block> 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<block> 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<block> outs;
outs.reserve(nl.n_out());
for (std::uint32_t id : nl.outputs())
outs.push_back(w[id]);
return outs;
}
inline std::vector<block> pad_to(std::vector<block> 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<block> xor_vec(const std::vector<block> & a,
const std::vector<block> & b)
{
if (a.size() != b.size())
throw std::logic_error("yao stack: length");
std::vector<block> 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<block> & 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<psnip_uint32_t>(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<std::uint8_t>(control_p0 ^ control_p1);
const block delta = stack_detail::with_lsb(dpf::uniform_sample<block>(), 1);
const block lg = dpf::uniform_sample<block>();
const block le = dpf::uniform_sample<block>();
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<block>(), 1);
stack_result result;
result.stack_blocks = n;
result.naive_blocks = g0.tables.size() + g1.tables.size();
result.extra_blocks = 2u + static_cast<std::size_t>(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>();
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<std::uint8_t>(
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<std::uint8_t>(
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<unsigned>(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<netlist> & branches,
std::uint16_t index_p0, std::uint16_t index_p1,
const std::vector<std::vector<std::uint8_t>> & p0,
const std::vector<std::vector<std::uint8_t>> & 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<std::uint16_t>(index_p0 ^ index_p1);
if (index >= k)
throw std::invalid_argument("yao stack: index");
const block delta = stack_detail::with_lsb(dpf::uniform_sample<block>(), 1);
std::vector<block> zbit(width);
std::vector<block> active_bit(width);
for (unsigned b = 0; b < width; ++b)
{
const block lg = dpf::uniform_sample<block>();
const block le = dpf::uniform_sample<block>();
zbit[b] = stack_detail::xor_b(lg, le);
const std::uint8_t bit0 = static_cast<std::uint8_t>((index_p0 >> b) & 1u);
const std::uint8_t bit1 = static_cast<std::uint8_t>((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<block> labels(width);
for (unsigned b = 0; b < width; ++b)
{
const std::uint8_t bit = static_cast<std::uint8_t>((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<block> seeds(k);
std::vector<stack_detail::seeded> 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<block> 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<std::vector<block>> demux(npat);
for (unsigned color = 0; color < npat; ++color)
{
std::uint16_t sem_idx = 0;
std::vector<block> color_labels(width);
for (unsigned b = 0; b < width; ++b)
{
const std::uint8_t cbit = static_cast<std::uint8_t>((color >> b) & 1u);
const std::uint8_t sbit = static_cast<std::uint8_t>(
cbit ^ stack_detail::lsb(zbit[b]));
sem_idx = static_cast<std::uint16_t>(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<std::uint32_t>(100u + j)));
}
}
std::uint16_t active_color = 0;
for (unsigned b = 0; b < width; ++b)
active_color = static_cast<std::uint16_t>(
active_color | (stack_detail::lsb(active_bit[b]) << b));
std::vector<block> 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<block> 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<std::uint32_t>(100u + j)));
}
std::vector<block> 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<std::size_t>(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<block>(), 1);
stack_result result;
result.stack_blocks = heaviest;
result.naive_blocks = naive;
result.extra_blocks = static_cast<std::size_t>(npat) * k
+ static_cast<std::size_t>(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<block>();
std::vector<block> table(static_cast<std::size_t>(npat) * 2u);
for (unsigned color = 0; color < npat; ++color)
{
std::uint16_t sem_idx = 0;
std::vector<block> color_labels(width);
for (unsigned b = 0; b < width; ++b)
{
const std::uint8_t cbit = static_cast<std::uint8_t>((color >> b) & 1u);
const std::uint8_t sbit = static_cast<std::uint8_t>(
cbit ^ stack_detail::lsb(zbit[b]));
sem_idx = static_cast<std::uint16_t>(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<std::size_t>(color) * 2u + oc;
if (sem_idx >= k)
{
table[idx] = dpf::uniform_sample<block>();
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<std::uint8_t>(
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<std::uint32_t>(200u + idx)));
}
}
const std::size_t idx = static_cast<std::size_t>(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<std::uint32_t>(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__