578 lines
22 KiB
C++
578 lines
22 KiB
C++
|
|
/// @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__
|