libdpf/include/dpf/yao.hpp

687 lines
21 KiB
C++
Raw Permalink Normal View History

/// @file dpf/yao.hpp
/// @brief Semi-honest garbled netlist for a boolean function of a DPF leaf.
/// @details The key still computes the point, the comparison, the interval, and
/// a public-offset polynomial. This header is the step after a leaf
/// share exists and the next function is a straight-line bit circuit
/// the key does not contain. The circuit this library is built around
/// is the zero-key AES-MMO block the PRG already uses, so a seed or a
/// leaf can enter that block without being opened.
///
/// Free-XOR and half-gates (Zahur, Rosulek, and Evans, ePrint 2014/756).
/// A secret if/else or a k-way switch is `dpf/yao_stack.hpp`: stacked
/// garbling (Heath and Kolesnikov, CRYPTO 2020) and one-hot garbling
/// (Heath and Kolesnikov, CCS 2021). The netlist below is unchanged.
/// Party 0 garbles. Party 1 evaluates. A shared input is the XOR of
/// the two parties' bits. A private input is known to one party.
/// Each output is an XOR share: the garbler's share is the permute
/// bit of the zero-label, and the evaluator's share is the color of
/// the label it holds. XOR of the two shares is the output bit.
///
/// The gate hash is fixed-key AES in the same Matyas–Meyer–Oseas
/// shape as `iknp::detail::ot_hash`, with a distinct tweak. Tables
/// are one-time. Gate ids increase by one per AND and are never
/// reused. There is no Yao domain on the composer. A ring share
/// becomes these bits through `dpf/yao_share.hpp`, then comes back
/// the same way. Comparisons and public-offset LUTs stay on the key.
///
/// `session` keeps the IKNP base OT. The first `eval` that needs a
/// choice label runs Chou–Orlandi; later evals on that session only
/// extend. Do not interleave `iknp::sample` on the same channel.
///
/// Ideal functionality (semi-honest): both parties input the bits named by
/// the netlist; each receives an XOR share of every output wire. Party 0
/// does not learn party 1's private bits or party 1's shares. Party 1 does
/// not learn party 0's private bits, party 0's shares, or `Δ`.
/// @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_HPP__
#define LIBDPF_INCLUDE_DPF_YAO_HPP__
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <stdexcept>
#include <utility>
#include <vector>
#include "hedley/hedley.h"
#include "simde/simde/x86/avx2.h"
#include "dpf/iknp.hpp"
#include "dpf/prg_aes.hpp"
#include "dpf/random.hpp"
namespace dpf
{
namespace yao
{
using block = simde__m128i;
/// @brief One wire in a `netlist`. Valid only for the netlist that minted it.
struct bit
{
std::uint32_t id = 0;
};
/// @brief Who knows an input wire.
enum class input : unsigned char
{
shared = 0, ///< XOR of the two parties' bits
priv0 = 1, ///< party 0's private bit
priv1 = 2 ///< party 1's private bit
};
namespace detail
{
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
block xor_b(block a, block b) noexcept
{
return simde_mm_xor_si128(a, b);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
block zero_b() noexcept
{
return simde_mm_setzero_si128();
}
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
std::uint8_t lsb(block x) noexcept
{
unsigned char b = 0;
std::memcpy(&b, &x, 1);
return static_cast<std::uint8_t>(b & 1u);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_PURE
HEDLEY_NO_THROW
block with_lsb(block x, std::uint8_t bit) noexcept
{
unsigned char b = 0;
std::memcpy(&b, &x, 1);
b = static_cast<unsigned char>((b & 0xfeu) | (bit & 1u));
std::memcpy(&x, &b, 1);
return x;
}
HEDLEY_ALWAYS_INLINE
HEDLEY_CONST
HEDLEY_NO_THROW
block mask_bit(std::uint8_t bit, block delta) noexcept
{
return (bit & 1u) ? delta : zero_b();
}
/// @brief Correlation-robust hash for one half-gate. Tweak `2·gid` is the
/// left wire and `2·gid+1` is the right wire. Both labels of one wire
/// use the same tweak.
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block cr_hash(block x, std::uint64_t tweak) noexcept
{
const prg::purpose_scope counted(prg::purpose::hash);
const auto mixed = xor_b(x, simde_mm_set_epi64x(
static_cast<std::int64_t>(tweak >> 32),
static_cast<std::int64_t>(static_cast<std::uint32_t>(tweak))));
return prg::aes128::eval(mixed,
static_cast<psnip_uint32_t>(tweak * 0x9E3779B9u) ^ 0xC3A5u);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block random_delta() noexcept
{
return with_lsb(dpf::uniform_sample<block>(), 1);
}
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block random_label() noexcept
{
return dpf::uniform_sample<block>();
}
/// @brief Half-gates AND. `la0` and `lb0` are zero-labels. `table[0]` and
/// `table[1]` are the two ciphertexts. Returns the output zero-label.
/// Matches `halfgates_garble` / `halfgates_eval` in EMP-toolkit.
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block and_garble(block la0, block lb0, block delta, std::uint64_t gid,
block table[2]) noexcept
{
const bool pa = lsb(la0) != 0;
const bool pb = lsb(lb0) != 0;
const block ha0 = cr_hash(la0, gid * 2u);
const block ha1 = cr_hash(xor_b(la0, delta), gid * 2u);
const block hb0 = cr_hash(lb0, gid * 2u + 1u);
const block hb1 = cr_hash(xor_b(lb0, delta), gid * 2u + 1u);
table[0] = xor_b(xor_b(ha0, ha1), pb ? delta : zero_b());
block w0 = xor_b(ha0, pa ? table[0] : zero_b());
const block tmp = xor_b(hb0, hb1);
table[1] = xor_b(tmp, la0);
w0 = xor_b(w0, hb0);
w0 = xor_b(w0, pb ? tmp : zero_b());
return w0;
}
HEDLEY_ALWAYS_INLINE
HEDLEY_NO_THROW
block and_eval(block a, block b, std::uint64_t gid, const block table[2]) noexcept
{
const bool sa = lsb(a) != 0;
const bool sb = lsb(b) != 0;
block w = xor_b(cr_hash(a, gid * 2u), cr_hash(b, gid * 2u + 1u));
if (sa)
w = xor_b(w, table[0]);
if (sb)
{
w = xor_b(w, table[1]);
w = xor_b(w, a);
}
return w;
}
inline void require_bit(std::uint8_t b, const char * what)
{
if (b > 1u)
throw std::invalid_argument(what);
}
} // namespace detail
/// @brief Straight-line circuit: XOR, AND, XNOR, NOT, and XOR with a public bit.
/// @details XOR, XNOR, and NOT are free. Each AND is one garbled row (32 bytes).
/// Wires are SSA. There is no secret shift, no memory, and no loop.
class netlist
{
public:
enum class op : unsigned char
{
xor_ = 0,
and_ = 1,
not_ = 2
};
struct gate
{
op code = op::xor_;
std::uint32_t a = 0;
std::uint32_t b = 0;
};
HEDLEY_WARN_UNUSED_RESULT
bit shared_in()
{
return push_in(input::shared);
}
/// @param owner 0 or 1
HEDLEY_WARN_UNUSED_RESULT
bit priv_in(unsigned owner)
{
if (owner > 1u)
throw std::invalid_argument("yao: input owner");
return push_in(owner == 0 ? input::priv0 : input::priv1);
}
HEDLEY_WARN_UNUSED_RESULT
bit xor_(bit a, bit b)
{
check(a);
check(b);
gates_.push_back(gate{op::xor_, a.id, b.id});
return bit{nwire() - 1u};
}
HEDLEY_WARN_UNUSED_RESULT
bit and_(bit a, bit b)
{
check(a);
check(b);
gates_.push_back(gate{op::and_, a.id, b.id});
++n_and_;
return bit{nwire() - 1u};
}
HEDLEY_WARN_UNUSED_RESULT
bit xnor_(bit a, bit b)
{
return not_(xor_(a, b));
}
HEDLEY_WARN_UNUSED_RESULT
bit not_(bit a)
{
check(a);
gates_.push_back(gate{op::not_, a.id, 0});
return bit{nwire() - 1u};
}
/// @brief XOR `x` with a public bit. A zero bit returns `x`.
HEDLEY_WARN_UNUSED_RESULT
bit xor_public(bit x, std::uint8_t public_bit)
{
check(x);
if ((public_bit & 1u) == 0)
return x;
return not_(x);
}
void out(bit x)
{
check(x);
outs_.push_back(x.id);
}
std::uint32_t n_in() const noexcept
{
return static_cast<std::uint32_t>(kinds_.size());
}
std::uint32_t n_out() const noexcept
{
return static_cast<std::uint32_t>(outs_.size());
}
std::uint32_t n_and() const noexcept { return n_and_; }
std::uint32_t nwire() const noexcept
{
return static_cast<std::uint32_t>(kinds_.size() + gates_.size());
}
/// @brief Bytes of garbled AND rows (`32 · n_and`).
std::size_t table_bytes() const noexcept
{
return static_cast<std::size_t>(n_and_) * 32u;
}
input kind_at(std::uint32_t i) const
{
return kinds_.at(i);
}
const std::vector<gate> & gates() const noexcept { return gates_; }
const std::vector<std::uint32_t> & outputs() const noexcept { return outs_; }
private:
bit push_in(input k)
{
if (!gates_.empty() || !outs_.empty())
throw std::logic_error("yao: inputs are declared before gates");
kinds_.push_back(k);
return bit{n_in() - 1u};
}
void check(bit x) const
{
if (x.id >= nwire())
throw std::invalid_argument("yao: wire");
}
std::vector<input> kinds_;
std::vector<gate> gates_;
std::vector<std::uint32_t> outs_;
std::uint32_t n_and_ = 0;
};
namespace detail
{
struct garbled
{
block delta{};
std::vector<block> z;
std::vector<block> tables;
std::vector<block> direct;
std::vector<block> ot0;
std::vector<block> ot1;
std::vector<std::uint8_t> share;
};
inline void garble_gates(const netlist & nl, block delta, std::vector<block> & z,
std::vector<block> & tables)
{
tables.clear();
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 & g = gates[gi];
const std::uint32_t dst = base + gi;
switch (g.code)
{
case netlist::op::xor_:
z[dst] = xor_b(z[g.a], z[g.b]);
break;
case netlist::op::not_:
z[dst] = xor_b(z[g.a], delta);
break;
case netlist::op::and_:
{
block row[2];
z[dst] = and_garble(z[g.a], z[g.b], delta, gid, row);
tables.push_back(row[0]);
tables.push_back(row[1]);
++gid;
break;
}
}
}
}
inline std::vector<std::uint8_t> eval_active(const netlist & nl,
const block * tables, std::vector<block> & w)
{
std::uint64_t gid = 0;
std::size_t ti = 0;
const auto & gates = nl.gates();
const std::uint32_t base = nl.n_in();
const std::size_t ntable = static_cast<std::size_t>(nl.n_and()) * 2u;
for (std::uint32_t gi = 0; gi < gates.size(); ++gi)
{
const auto & g = gates[gi];
const std::uint32_t dst = base + gi;
switch (g.code)
{
case netlist::op::xor_:
w[dst] = xor_b(w[g.a], w[g.b]);
break;
case netlist::op::not_:
w[dst] = w[g.a];
break;
case netlist::op::and_:
if (ti + 2u > ntable)
throw std::runtime_error("yao: truncated garbled tables");
w[dst] = and_eval(w[g.a], w[g.b], gid, tables + ti);
ti += 2u;
++gid;
break;
}
}
std::vector<std::uint8_t> share(nl.n_out());
const auto & outs = nl.outputs();
for (std::uint32_t i = 0; i < outs.size(); ++i)
share[i] = lsb(w[outs[i]]);
return share;
}
inline void check_party_bits(const netlist & nl, int me, const std::uint8_t * in)
{
if (me != 0 && me != 1)
throw std::invalid_argument("yao: party");
if (nl.n_in() != 0 && in == nullptr)
throw std::invalid_argument("yao: missing inputs");
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
const bool used = k == input::shared
|| (k == input::priv0 && me == 0)
|| (k == input::priv1 && me == 1);
if (used)
require_bit(in[i], "yao: input bit");
}
}
inline garbled garble_party0(const netlist & nl, const std::uint8_t * p0)
{
check_party_bits(nl, 0, p0);
garbled g;
g.delta = random_delta();
g.z.assign(nl.nwire(), zero_b());
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const block lg = random_label();
switch (nl.kind_at(i))
{
case input::shared:
{
const block le = random_label();
g.z[i] = xor_b(lg, le);
g.direct.push_back(xor_b(lg, mask_bit(p0[i], g.delta)));
g.ot0.push_back(le);
g.ot1.push_back(xor_b(le, g.delta));
break;
}
case input::priv0:
g.z[i] = lg;
g.direct.push_back(xor_b(lg, mask_bit(p0[i], g.delta)));
break;
case input::priv1:
g.z[i] = lg;
g.ot0.push_back(lg);
g.ot1.push_back(xor_b(lg, g.delta));
break;
}
}
garble_gates(nl, g.delta, g.z, g.tables);
g.share.resize(nl.n_out());
const auto & outs = nl.outputs();
for (std::uint32_t i = 0; i < outs.size(); ++i)
g.share[i] = lsb(g.z[outs[i]]);
return g;
}
inline std::vector<std::uint8_t> eval_party1(const netlist & nl,
const std::vector<block> & tables, const std::vector<block> & direct,
const std::vector<block> & chosen)
{
if (tables.size() != static_cast<std::size_t>(nl.n_and()) * 2u)
throw std::runtime_error("yao: table size");
std::vector<block> w(nl.nwire(), zero_b());
std::size_t di = 0;
std::size_t oi = 0;
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
switch (nl.kind_at(i))
{
case input::shared:
if (di >= direct.size() || oi >= chosen.size())
throw std::runtime_error("yao: input label count");
w[i] = xor_b(direct[di++], chosen[oi++]);
break;
case input::priv0:
if (di >= direct.size())
throw std::runtime_error("yao: input label count");
w[i] = direct[di++];
break;
case input::priv1:
if (oi >= chosen.size())
throw std::runtime_error("yao: input label count");
w[i] = chosen[oi++];
break;
}
}
if (di != direct.size() || oi != chosen.size())
throw std::runtime_error("yao: input label count");
return eval_active(nl, tables.data(), w);
}
} // namespace detail
/// @brief Cleartext walk. `semantic[i]` is the bit on input wire `i`.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> eval_plain(const netlist & nl,
const std::uint8_t * semantic)
{
if (nl.n_in() != 0 && semantic == nullptr)
throw std::invalid_argument("yao: missing inputs");
std::vector<std::uint8_t> w(nl.nwire(), 0);
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
detail::require_bit(semantic[i], "yao: input bit");
w[i] = semantic[i];
}
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 & g = gates[gi];
const std::uint32_t dst = base + gi;
switch (g.code)
{
case netlist::op::xor_:
w[dst] = static_cast<std::uint8_t>(w[g.a] ^ w[g.b]);
break;
case netlist::op::and_:
w[dst] = static_cast<std::uint8_t>(w[g.a] & w[g.b]);
break;
case netlist::op::not_:
w[dst] = static_cast<std::uint8_t>(w[g.a] ^ 1u);
break;
}
}
std::vector<std::uint8_t> out(nl.n_out());
const auto & outs = nl.outputs();
for (std::uint32_t i = 0; i < outs.size(); ++i)
out[i] = w[outs[i]];
return out;
}
/// @brief Garble and evaluate in one process. Returns both XOR shares.
/// @details `p0` and `p1` use the same layout as `session::eval`. The
/// reconstructed bit is `first[i] XOR second[i]`.
HEDLEY_WARN_UNUSED_RESULT
inline std::pair<std::vector<std::uint8_t>, std::vector<std::uint8_t>>
eval_pair(const netlist & nl, const std::uint8_t * p0, const std::uint8_t * p1)
{
detail::check_party_bits(nl, 0, p0);
detail::check_party_bits(nl, 1, p1);
auto g = detail::garble_party0(nl, p0);
std::vector<block> chosen;
chosen.reserve(g.ot0.size());
std::size_t oi = 0;
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
if (k != input::shared && k != input::priv1)
continue;
if (oi >= g.ot0.size())
throw std::logic_error("yao: ot count");
chosen.push_back((p1[i] & 1u) ? g.ot1[oi] : g.ot0[oi]);
++oi;
}
auto share1 = detail::eval_party1(nl, g.tables, g.direct, chosen);
return {std::move(g.share), std::move(share1)};
}
/// @brief One-process garble of a known input. Shared wires take the semantic
/// bit as party 0's share and zero as party 1's share.
HEDLEY_WARN_UNUSED_RESULT
inline std::vector<std::uint8_t> eval_local(const netlist & nl,
const std::uint8_t * semantic)
{
std::vector<std::uint8_t> p0(nl.n_in(), 0);
std::vector<std::uint8_t> p1(nl.n_in(), 0);
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
detail::require_bit(semantic[i], "yao: input bit");
if (nl.kind_at(i) == input::priv1)
p1[i] = semantic[i];
else
p0[i] = semantic[i];
}
auto [a, b] = eval_pair(nl, p0.data(), p1.data());
for (std::size_t i = 0; i < a.size(); ++i)
a[i] = static_cast<std::uint8_t>(a[i] ^ b[i]);
return a;
}
/// @brief One party's garbling session. Party 0 always garbles.
/// @details The first `eval` fixes `me`. A second call with the other party
/// id throws. Base OT state stays inside the session.
class session
{
public:
session() = default;
session(const session &) = delete;
session & operator=(const session &) = delete;
session(session &&) = default;
session & operator=(session &&) = default;
/// @brief XOR shares of each output wire.
/// @param me 0 (garbler) or 1 (evaluator)
/// @param nl circuit. Both parties pass equal netlists.
/// @param in_bits length `nl.n_in()`. Shared entries are this party's
/// XOR share. `priv0` is read on party 0. `priv1` is read on party 1.
/// @param link peer channel. No other traffic may run on it during `eval`.
HEDLEY_WARN_UNUSED_RESULT
std::vector<std::uint8_t> eval(int me, const netlist & nl,
const std::uint8_t * in_bits, net::channel & link)
{
if (!bound_)
{
if (me != 0 && me != 1)
throw std::invalid_argument("yao: party");
bound_ = true;
me_ = me;
}
else if (me != me_)
throw std::logic_error("yao: session party is fixed");
detail::check_party_bits(nl, me, in_bits);
if (me == 0)
{
auto g = detail::garble_party0(nl, in_bits);
std::vector<block> ignored;
iknp::transfer_labels(link, 0, true, ot_, g.ot0, g.ot1, {}, ignored);
std::vector<block> blob;
blob.reserve(g.tables.size() + g.direct.size());
blob.insert(blob.end(), g.tables.begin(), g.tables.end());
blob.insert(blob.end(), g.direct.begin(), g.direct.end());
link.send_vec(blob, net::msg::bytes);
return g.share;
}
std::vector<std::uint8_t> choices;
choices.reserve(nl.n_in());
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
if (k == input::shared || k == input::priv1)
choices.push_back(in_bits[i]);
}
std::vector<block> chosen;
iknp::transfer_labels(link, 1, false, ot_, {}, {}, choices, chosen);
const auto blob = link.recv_vec<block>(net::msg::bytes);
const std::size_t ntable = static_cast<std::size_t>(nl.n_and()) * 2u;
std::size_t ndirect = 0;
for (std::uint32_t i = 0; i < nl.n_in(); ++i)
{
const auto k = nl.kind_at(i);
if (k == input::shared || k == input::priv0)
++ndirect;
}
if (blob.size() != ntable + ndirect)
throw std::runtime_error("yao: garbled payload size");
std::vector<block> tables(ntable);
std::vector<block> direct(ndirect);
if (ntable != 0)
std::memcpy(tables.data(), blob.data(), ntable * sizeof(block));
if (ndirect != 0)
std::memcpy(direct.data(), blob.data() + ntable, ndirect * sizeof(block));
return detail::eval_party1(nl, tables, direct, chosen);
}
private:
bool bound_ = false;
int me_ = 0;
iknp::detail::role_state ot_{};
};
} // namespace yao
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_YAO_HPP__