libdpf/include/grotto/lut_union.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

601 lines
23 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @file grotto/lut_union.hpp
/// @brief One comparison for several piecewise LUTs.
/// @details The breakpoints of every LUT are shifted by the public `eta` and
/// refined with the same carry cuts as offset Horner. Their union is
/// one sorted knot vector. One comparison, payload
/// `1, center, …, center^d` at the maximum degree, then one
/// prefix-parity walk of that union. Each LUT piece is a circular
/// span of those union knots (one polynomial, one public `kappa`),
/// so the functions are independent dots of the same segment table.
///
/// The geneval schedule is one full-depth `fss_cmp`. Its round count
/// is the input bitlength. It does not grow with the number of LUTs
/// or the number of union knots. Prefix parity stays local.
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license.
#ifndef LIBDPF_INCLUDE_GROTTO_LUT_UNION_HPP__
#define LIBDPF_INCLUDE_GROTTO_LUT_UNION_HPP__
#include "hedley/hedley.h"
#include "dpf/compose.hpp"
#include "grotto/constant_lut.hpp"
#include "grotto/easy_lut.hpp"
#include "grotto/offset_poly.hpp"
#include <algorithm>
#include <array>
#include <cstddef>
#include <cstdint>
#include <stdexcept>
#include <utility>
#include <vector>
namespace grotto
{
/// @brief Public piecewise polynomial. `knots` are strictly increasing left
/// endpoints. `coeff[i][k]` is the degree-`k` coefficient of that piece
/// in `Z/2^64`. Every row of one LUT has the same length.
/// @tparam InputT input domain type
template <typename InputT>
struct piecewise_lut
{
static_assert(std::is_integral_v<InputT>, "lut union domain must be an integer group");
std::vector<InputT> knots;
std::vector<std::vector<uint64_t>> coeff;
};
/// @brief One refined piece of one LUT, as a circular span of the union.
/// @details `coeff` is already the binomial shift by `kappa`. The dot with
/// segment shares of `center^m` is the piece's contribution.
/// Indices run over `lut_union_plan::knots`. When `begin <= end` the
/// span is `[begin, end)`. When `begin > end` it wraps:
/// `[begin, knots) ∪ [0, end)`.
/// @tparam InputT input domain type
template <typename InputT>
struct lut_span
{
std::size_t begin = 0;
std::size_t end = 0;
std::int64_t kappa = 0;
std::vector<uint64_t> coeff;
};
/// @brief Union of several LUTs at one public `eta`, plus the geneval shape.
/// @tparam InputT input domain type
template <typename InputT>
struct lut_union_plan
{
static_assert(std::is_integral_v<InputT>, "lut union domain must be an integer group");
/// @brief One comparison, however many LUTs and knots are in the union.
static constexpr std::size_t comparisons = 1;
/// @brief One memoized prefix-parity walk of `knots`. Local: no rounds.
static constexpr std::size_t prefix_walks = 1;
/// @brief Public offset the knots were shifted by.
InputT eta{};
/// @brief Maximum degree across the LUTs. The comparison payload has `lanes` words.
std::size_t degree = 0;
/// @brief `degree + 1`.
std::size_t lanes = 0;
/// @brief Geneval / comparison depth: bitlength of `InputT`. One correction per level.
std::size_t depth = 0;
/// @brief Union of the refined, `eta`-shifted knots. Sorted, unique.
std::vector<InputT> knots;
/// @brief `funcs[t]` is LUT `t`, in the order passed to `make_lut_union_plan`.
std::vector<std::vector<lut_span<InputT>>> funcs;
/// @brief Number of union breakpoints the prefix walk visits.
HEDLEY_NO_THROW
std::size_t endpoints() const noexcept { return knots.size(); }
/// @brief Interactive geneval rounds: one correction word per input bit.
HEDLEY_NO_THROW
std::size_t geneval_rounds() const noexcept { return depth; }
};
/// @brief Comparison slot for `schedule_lut_union`.
/// @details One AES block, or the power-vector width when that is wider.
/// @param lanes payload lanes (`degree + 1`)
/// @return Slot width in bytes
/// \complexity `Θ(1)`.
inline std::size_t lut_union_slot_bytes(std::size_t lanes)
{
const std::size_t value = lanes * sizeof(std::uint64_t);
return value < 16u ? 16u : value;
}
/// @brief Exact `easy_lut` as a piecewise program.
/// @details A denominator other than 1 is a rounding division. That is not a
/// dot in `Z/2^64`, so those tables are rejected. Trailing powers that
/// are zero on every piece are dropped.
/// @tparam Raw signed raw word
/// @param lut the cleartext table
/// @return Piecewise program on the same bounds
/// @throws std::invalid_argument if a denominator is not 1, or the table is empty
/// \complexity One pass over the pieces. `Θ(P)`.
/// @see grotto::easy_lut
template <typename Raw>
piecewise_lut<Raw> piecewise_from_easy(const easy_lut<Raw> & lut)
{
if (lut.bounds.empty() || lut.bounds.size() != lut.c0.size()
|| lut.c0.size() != lut.c1.size() || lut.c1.size() != lut.c2.size()
|| lut.c2.size() != lut.den.size())
throw std::invalid_argument("lut union: easy lut pieces are incomplete");
std::size_t width = 1;
for (std::size_t i = 0; i < lut.den.size(); ++i)
{
if (lut.den[i] != 1)
throw std::invalid_argument(
"lut union: easy lut denominator is not an exact Z/2^64 dot");
if (lut.c1[i] != 0)
width = std::max<std::size_t>(width, 2);
if (lut.c2[i] != 0)
width = std::max<std::size_t>(width, 3);
}
piecewise_lut<Raw> out;
out.knots = lut.bounds;
out.coeff.resize(lut.bounds.size());
for (std::size_t i = 0; i < lut.bounds.size(); ++i)
{
const uint64_t row[3] = {
static_cast<uint64_t>(lut.c0[i]),
static_cast<uint64_t>(lut.c1[i]),
static_cast<uint64_t>(lut.c2[i])};
out.coeff[i].assign(row, row + width);
}
return out;
}
/// @brief Degree-0 `constant_lut` as a piecewise program.
/// @tparam Raw signed raw word
/// @param lut the cleartext table
/// @return One constant coefficient per bound
/// @throws std::invalid_argument if the table is empty or mis-sized
/// \complexity One pass over the pieces. `Θ(P)`.
/// @see grotto::constant_lut
template <typename Raw>
piecewise_lut<Raw> piecewise_from_constant(const constant_lut<Raw> & lut)
{
if (lut.bounds.empty() || lut.bounds.size() != lut.values.size())
throw std::invalid_argument("lut union: constant lut pieces are incomplete");
piecewise_lut<Raw> out;
out.knots = lut.bounds;
out.coeff.resize(lut.values.size());
for (std::size_t i = 0; i < lut.values.size(); ++i)
out.coeff[i] = {static_cast<uint64_t>(lut.values[i])};
return out;
}
namespace lut_union_detail
{
template <typename InputT>
std::size_t knot_index(const std::vector<InputT> & knots, InputT knot)
{
const auto it = std::lower_bound(knots.begin(), knots.end(), knot);
if (it == knots.end() || *it != knot)
throw std::logic_error("lut union: refined knot missing from the union");
return static_cast<std::size_t>(it - knots.begin());
}
template <typename InputT>
void mark_span(std::vector<unsigned char> & seen, const lut_span<InputT> & span,
std::size_t n)
{
const auto mark = [&](std::size_t j) {
if (j >= n || seen[j] != 0)
throw std::logic_error("lut union: piece spans overlap or leave the union");
seen[j] = 1;
};
if (span.begin <= span.end)
{
for (std::size_t j = span.begin; j < span.end; ++j)
mark(j);
}
else
{
for (std::size_t j = span.begin; j < n; ++j)
mark(j);
for (std::size_t j = 0; j < span.end; ++j)
mark(j);
}
}
template <typename At>
uint64_t dot_span(std::size_t n, std::size_t begin, std::size_t end,
const std::vector<uint64_t> & coeff, At && at)
{
uint64_t value = 0;
const auto visit = [&](std::size_t j) {
for (std::size_t m = 0; m < coeff.size(); ++m)
value += at(j, m) * coeff[m];
};
if (begin <= end)
{
for (std::size_t j = begin; j < end; ++j)
visit(j);
}
else
{
for (std::size_t j = begin; j < n; ++j)
visit(j);
for (std::size_t j = 0; j < end; ++j)
visit(j);
}
return value;
}
template <typename InputT, typename Table>
std::vector<uint64_t> dot_plan(const lut_union_plan<InputT> & plan, const Table & table)
{
const std::size_t n = plan.knots.size();
std::vector<uint64_t> out(plan.funcs.size(), 0);
for (std::size_t f = 0; f < plan.funcs.size(); ++f)
{
for (const auto & span : plan.funcs[f])
{
if (span.coeff.size() > plan.lanes)
throw std::invalid_argument("lut union: piece wider than the plan");
out[f] += dot_span(n, span.begin, span.end, span.coeff,
[&](std::size_t j, std::size_t m) { return table[j][m]; });
}
}
return out;
}
} // namespace lut_union_detail
/// @brief Union the refined breakpoints and record each LUT's spans.
/// @tparam InputT input domain type
/// @param luts piecewise programs on the same domain
/// @param eta public offset `eta = x - r`
/// @return Plan whose geneval depth is `bitlength(InputT)` and whose
/// comparison count is 1
/// @throws std::invalid_argument if `luts` is empty, a degree exceeds 16, or a
/// LUT's knots and coefficient rows disagree
/// \complexity Each LUT is prepared once (`Θ(P log P)` per LUT, including the
/// carry cuts). The union sort is `Θ(U log U)` for `U` the total refined
/// knots. Building the spans is `Θ(U)`. No keys.
/// \rounds None.
/// \communication None.
/// @see grotto::lut_union_eval
/// @see grotto::schedule_lut_union
template <typename InputT>
lut_union_plan<InputT> make_lut_union_plan(
const std::vector<piecewise_lut<InputT>> & luts, InputT eta)
{
if (luts.empty())
throw std::invalid_argument("lut union: no luts");
using prepared = offset_poly_detail::prepared<InputT>;
std::vector<std::vector<prepared>> refined;
refined.reserve(luts.size());
std::size_t degree = 0;
std::vector<InputT> knots;
for (const auto & lut : luts)
{
if (lut.coeff.empty() || lut.coeff.front().empty())
throw std::invalid_argument("lut union: coefficient row is empty");
const std::size_t width = lut.coeff.front().size();
if (width - 1 > offset_poly_max_degree)
throw std::invalid_argument("lut union: degree exceeds 16");
degree = std::max(degree, width - 1);
auto rows = offset_poly_detail::prepare(lut.knots, lut.coeff, eta);
for (const auto & row : rows)
knots.push_back(row.knot);
refined.push_back(std::move(rows));
}
std::sort(knots.begin(), knots.end());
knots.erase(std::unique(knots.begin(), knots.end()), knots.end());
lut_union_plan<InputT> plan;
plan.eta = eta;
plan.degree = degree;
plan.lanes = degree + 1;
plan.depth = dpf::utils::bitlength_of_v<InputT>;
plan.knots = std::move(knots);
plan.funcs.resize(refined.size());
const std::size_t n = plan.knots.size();
for (std::size_t f = 0; f < refined.size(); ++f)
{
const auto & rows = refined[f];
std::vector<unsigned char> seen(n, 0);
plan.funcs[f].reserve(rows.size());
if (rows.size() == 1)
{
lut_span<InputT> span;
span.begin = 0;
span.end = n;
span.kappa = rows[0].kappa;
span.coeff = offset_poly_detail::binomial_shift(
rows[0].coeff, static_cast<uint64_t>(rows[0].kappa));
lut_union_detail::mark_span(seen, span, n);
plan.funcs[f].push_back(std::move(span));
}
else
{
for (std::size_t i = 0; i < rows.size(); ++i)
{
lut_span<InputT> span;
span.begin = lut_union_detail::knot_index(plan.knots, rows[i].knot);
const std::size_t next = (i + 1 == rows.size()) ? 0 : i + 1;
span.end = lut_union_detail::knot_index(plan.knots, rows[next].knot);
span.kappa = rows[i].kappa;
span.coeff = offset_poly_detail::binomial_shift(
rows[i].coeff, static_cast<uint64_t>(rows[i].kappa));
lut_union_detail::mark_span(seen, span, n);
plan.funcs[f].push_back(std::move(span));
}
}
for (unsigned char bit : seen)
{
if (bit == 0)
throw std::logic_error("lut union: a union knot is outside every piece");
}
}
return plan;
}
/// @brief Each LUT's share, from one segment walk of the union.
/// @tparam Party party index, `0` or `1`
/// @tparam InputT input domain type
/// @param mat one comparison whose degree is at least `plan.degree`
/// @param plan the union plan
/// @param tokens proof token folded by the segment walk, or null
/// @return `out[t]` is party `Party`'s share of LUT `t`
/// @throws std::invalid_argument if the key is narrower than the plan
/// \complexity One `lane_segments` walk of the union (`prefix_walks == 1`),
/// then a dot per LUT piece. The walk is `signed` prefix parity over `U`
/// endpoints of a `lanes`-wide payload, reusing the common prefix.
/// Arithmetic after the walk is `Θ(Σ pieces · degree)`.
/// \rounds None. `eta` is already in the plan.
/// \communication None.
/// \preprocessing None created here. Uses `make_offset_poly_keys` at `plan.degree`.
/// @see grotto::offset_poly_eval
/// @see grotto::geneval_lut_union
template <std::size_t Party, typename InputT>
std::vector<uint64_t> lut_union_eval(
const offset_poly_keys<InputT> & mat,
const lut_union_plan<InputT> & plan,
dpf::proof_token * tokens = nullptr)
{
static_assert(Party < 2, "lut union party is 0 or 1");
using namespace offset_horner_detail;
if (mat.degree < plan.degree)
throw std::invalid_argument("lut union: key degree is below the union");
const std::size_t lanes = mat.verifiable ? mat.keys_v.index() : mat.keys.index();
if (lanes != mat.degree + 1)
throw std::invalid_argument("lut union: comparison payload width differs from the degree");
dpf::proof_token * pi = (mat.verifiable && tokens != nullptr) ? &tokens[0] : nullptr;
const auto table = mat.verifiable
? lane_segments<Party>(mat.keys_v, plan.knots, mat.wrap_share, pi)
: lane_segments<Party>(mat.keys, plan.knots, mat.wrap_share, nullptr);
if (mat.verifiable)
replicate_proof(tokens, mat.degree + 1);
return lut_union_detail::dot_plan(plan, table);
}
/// @brief Both parties' LUT shares from one Doerner–Shelat comparison.
/// @tparam InputT input domain type
template <typename InputT>
struct geneval_lut_union_result
{
/// @brief Public offset copied from the plan.
InputT eta{};
std::vector<uint64_t> value0;
std::vector<uint64_t> value1;
};
namespace lut_union_detail
{
template <std::size_t N, typename Keys, typename InputT>
geneval_lut_union_result<InputT> geneval_from_keys(
const Keys & keys, const lut_union_plan<InputT> & plan, const uint64_t * payload)
{
std::array<uint64_t, N> wrap0{};
std::array<uint64_t, N> wrap1{};
for (std::size_t m = 0; m < N; ++m)
{
const uint64_t blind = dpf::uniform_sample<uint64_t>();
wrap0[m] = blind;
wrap1[m] = payload[m] - blind;
}
const auto seg0 = offset_horner_detail::segments_lanes<N>(
keys.first, plan.knots, wrap0, nullptr);
const auto seg1 = offset_horner_detail::segments_lanes<N>(
keys.second, plan.knots, wrap1, nullptr);
const std::size_t pieces = plan.knots.size();
std::vector<std::vector<uint64_t>> table0(pieces, std::vector<uint64_t>(N));
std::vector<std::vector<uint64_t>> table1(pieces, std::vector<uint64_t>(N));
for (std::size_t m = 0; m < N; ++m)
{
for (std::size_t j = 0; j < pieces; ++j)
{
table0[j][m] = seg0[m][j];
table1[j][m] = seg1[m][j];
}
}
geneval_lut_union_result<InputT> out;
out.eta = plan.eta;
out.value0 = dot_plan(plan, table0);
out.value1 = dot_plan(plan, table1);
return out;
}
template <std::size_t N, typename InputT, typename Rng>
geneval_lut_union_result<InputT> geneval_lut_union_n(
bool arith, InputT center0, InputT center1,
const lut_union_plan<InputT> & plan, const uint64_t * payload, Rng rng)
{
dpf::vec<uint64_t, N> beta;
for (std::size_t m = 0; m < N; ++m)
beta[m] = payload[m];
if (arith)
{
auto keys = dpf::make_dpf_doerner_shelat(dpf::arith_input, center0, center1,
std::move(rng), dpf::gt(beta), dpf::verifiable{});
return geneval_from_keys<N>(keys, plan, payload);
}
auto keys = dpf::make_dpf_doerner_shelat(center0, center1,
std::move(rng), dpf::gt(beta), dpf::verifiable{});
return geneval_from_keys<N>(keys, plan, payload);
}
template <std::size_t N, typename InputT, typename Rng>
geneval_lut_union_result<InputT> geneval_dispatch(
bool arith, InputT center0, InputT center1,
const lut_union_plan<InputT> & plan, const uint64_t * payload, Rng & rng)
{
if (plan.lanes == N)
return geneval_lut_union_n<N>(arith, center0, center1, plan, payload,
std::move(rng));
if constexpr (N > 1)
return geneval_dispatch<N - 1>(arith, center0, center1, plan, payload, rng);
throw std::invalid_argument("lut union: lane count out of range");
}
template <typename InputT, typename Rng>
geneval_lut_union_result<InputT> geneval_lut_union_at(
bool arith, InputT center0, InputT center1, InputT center,
const lut_union_plan<InputT> & plan, Rng rng)
{
if (plan.lanes == 0 || plan.lanes > offset_horner_detail::lane_key_max)
throw std::invalid_argument("lut union: lane count out of range");
if (plan.knots.empty())
throw std::invalid_argument("lut union: plan has no knots");
std::vector<uint64_t> payload(plan.lanes, 1);
const uint64_t base = offset_horner_detail::lift(center);
for (std::size_t m = 1; m < plan.lanes; ++m)
payload[m] = payload[m - 1] * base;
return geneval_dispatch<offset_horner_detail::lane_key_max>(
arith, center0, center1, plan, payload.data(), rng);
}
} // namespace lut_union_detail
/// @brief Geneval of every LUT. The center is XOR-shared as in `geneval_point`.
/// @details One Doerner–Shelat comparison of `plan.lanes` words, then the same
/// local span dots as `lut_union_eval`. Further LUTs do not add a
/// comparison or a prefix walk.
/// @tparam InputT input domain type
/// @tparam Rng Doerner–Shelat randomness
/// @param center0 party 0 share of the center
/// @param center1 party 1 share of the center
/// @param plan the union plan
/// @param rng the Doerner–Shelat randomness tapes
/// @return Both parties' shares of every LUT
/// \complexity One Doerner–Shelat generation of a `plan.lanes` comparison
/// (`degree ≤ 16`), then one prefix-parity walk of the `U` union knots.
/// The dots are `Θ(Σ pieces · degree)`.
/// @note Rounds of that comparison are `plan.geneval_rounds()` (`bitlength` of
/// the input). They are not multiplied by the LUT count or by `U`.
/// \preprocessing The randomness object the caller passes (`Rng`). This function
/// also samples one `uint64_t` blind per power.
/// @see grotto::geneval_offset_horner
/// @see grotto::schedule_lut_union
template <typename InputT, typename Rng>
geneval_lut_union_result<InputT> geneval_lut_union(
InputT center0, InputT center1,
const lut_union_plan<InputT> & plan, Rng rng)
{
const InputT center = geneval_offset_horner_center(center0, center1);
return lut_union_detail::geneval_lut_union_at(
false, center0, center1, center, plan, std::move(rng));
}
/// @brief Same as `geneval_lut_union`, with the library's urandom pad tape.
/// \complexity Same as the overload that takes `Rng`.
/// \preprocessing Samples a `dpf::ds_randomness` tape from `uniform_sample`.
/// @see grotto::geneval_lut_union
template <typename InputT>
geneval_lut_union_result<InputT> geneval_lut_union(
InputT center0, InputT center1, const lut_union_plan<InputT> & plan)
{
using block = typename dpf::prg::aes128::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
dpf::ds_randomness<block (*)(), dpf::detail::urandom_pad_rng> rng{
dpf::uniform_sample<block>, {}};
HEDLEY_PRAGMA(GCC diagnostic pop)
return geneval_lut_union(center0, center1, plan, std::move(rng));
}
/// @brief Geneval when `center0 + center1` is the center in the input group.
/// \complexity Same as the XOR-share overload: one comparison, one prefix walk.
/// @see grotto::geneval_lut_union
template <typename InputT, typename Rng>
geneval_lut_union_result<InputT> geneval_lut_union(
dpf::arith_input_t, InputT center0, InputT center1,
const lut_union_plan<InputT> & plan, Rng rng)
{
const InputT center = offset_horner_group_add(center0, center1);
return lut_union_detail::geneval_lut_union_at(
true, center0, center1, center, plan, std::move(rng));
}
/// @brief Additive-share geneval with the library's urandom pad tape.
/// \complexity Same as the overload that takes `Rng`.
/// @see grotto::geneval_lut_union
template <typename InputT>
geneval_lut_union_result<InputT> geneval_lut_union(
dpf::arith_input_t, InputT center0, InputT center1,
const lut_union_plan<InputT> & plan)
{
using block = typename dpf::prg::aes128::block_type;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
dpf::ds_randomness<block (*)(), dpf::detail::urandom_pad_rng> rng{
dpf::uniform_sample<block>, {}};
HEDLEY_PRAGMA(GCC diagnostic pop)
return geneval_lut_union(dpf::arith_input, center0, center1, plan, std::move(rng));
}
/// @brief Record one full-depth comparison on an existing seed.
/// @details Does not walk the union knots and does not emit a node per LUT.
/// Prefix parity of `plan.knots` is local after this comparison.
/// @tparam InputT input domain type
/// @param composer the schedule being built
/// @param seed FSS seed already recorded on `composer`
/// @param plan the union plan
/// @return The comparison leaf (`domain::a`)
/// @throws std::invalid_argument if the plan depth or lane count is zero
/// \complexity `depth` expand/step pairs. `depth` is `bitlength(InputT)`.
/// Extra composer nodes are `Θ(depth)`, independent of `plan.funcs.size()`
/// and of `plan.endpoints()`.
/// \rounds `plan.geneval_rounds()` once the composer is scheduled.
/// \communication One correction slot per level, `lut_union_slot_bytes(plan.lanes)` wide.
/// @see grotto::make_lut_union_plan
template <typename InputT>
dpf::protocol::node schedule_lut_union(
dpf::protocol::composer & composer, dpf::protocol::node seed,
const lut_union_plan<InputT> & plan)
{
if (plan.depth == 0 || plan.lanes == 0)
throw std::invalid_argument("lut union: plan has no comparison");
return composer.fss_cmp(seed, plan.depth, lut_union_slot_bytes(plan.lanes));
}
/// @brief Record a fresh seed and one full-depth comparison for `plan`.
/// \complexity One input node plus `schedule_lut_union` on that seed.
/// \rounds `plan.geneval_rounds()`.
/// @see grotto::schedule_lut_union
template <typename InputT>
dpf::protocol::node schedule_lut_union(
dpf::protocol::composer & composer, const lut_union_plan<InputT> & plan)
{
auto seed = composer.input(dpf::protocol::domain::fss, 16);
return schedule_lut_union(composer, seed, plan);
}
} // namespace grotto
#endif // LIBDPF_INCLUDE_GROTTO_LUT_UNION_HPP__